Skip to main content

teeny_kernels/nn/activation/
hard.rs

1/*
2 * Copyright (c) 2026 Teenygrad.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *   http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17#![allow(non_snake_case)]
18
19use teeny_core::dtype::Float;
20use teeny_macros::kernel;
21use teeny_triton::triton::{
22    types::{AddOffsets, Comparison},
23    *,
24};
25
26// ── Hardtanh ─────────────────────────────────────────────────────────────────
27
28/// Forward: y = clamp(x, min_val, max_val)
29#[kernel(backward = HardtanhBackward)]
30pub fn hardtanh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
31    x_ptr: T::Pointer<D>,
32    y_ptr: T::Pointer<D>,
33    n_elements: i32,
34    min_val: f32,
35    max_val: f32,
36) where
37    T::I32Tensor: types::Tensor<i32, 1>,
38    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
39    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
40{
41    let pid = T::program_id(Axis::X);
42    let block_start = pid * BLOCK_SIZE;
43    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
44    let in_bounds = offsets.lt(n_elements);
45
46    let x = T::load(
47        x_ptr.add_offsets(offsets),
48        Some(in_bounds),
49        None,
50        &[],
51        None,
52        None,
53        None,
54        false,
55    );
56    let lo = T::full(&[BLOCK_SIZE], D::from_f64(min_val as f64));
57    let hi = T::full(&[BLOCK_SIZE], D::from_f64(max_val as f64));
58    let y = T::clamp(x, lo, hi);
59    T::store(
60        y_ptr.add_offsets(offsets),
61        y,
62        Some(in_bounds),
63        &[],
64        None,
65        None,
66    );
67}
68
69/// Backward: dx = dy if min_val < x < max_val else 0
70#[kernel]
71pub fn hardtanh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
72    dy_ptr: T::Pointer<D>,
73    x_ptr: T::Pointer<D>,
74    dx_ptr: T::Pointer<D>,
75    n_elements: i32,
76    min_val: f32,
77    max_val: f32,
78) where
79    T::I32Tensor: types::Tensor<i32, 1>,
80    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
81    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
82{
83    let pid = T::program_id(Axis::X);
84    let block_start = pid * BLOCK_SIZE;
85    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
86    let in_bounds = offsets.lt(n_elements);
87
88    let dy = T::load(
89        dy_ptr.add_offsets(offsets),
90        Some(in_bounds),
91        None,
92        &[],
93        None,
94        None,
95        None,
96        false,
97    );
98    let x = T::load(
99        x_ptr.add_offsets(offsets),
100        Some(in_bounds),
101        None,
102        &[],
103        None,
104        None,
105        None,
106        false,
107    );
108
109    // |x - midpoint| < half_range  ≡  min_val < x < max_val
110    let lo = T::full(&[BLOCK_SIZE], D::from_f64(min_val as f64));
111    let hi = T::full(&[BLOCK_SIZE], D::from_f64(max_val as f64));
112    let in_range = T::gt(T::minimum(x - lo, hi - x), T::zeros_like(x));
113    let dx = T::where_(in_range, dy, T::zeros_like(dy));
114    T::store(
115        dx_ptr.add_offsets(offsets),
116        dx,
117        Some(in_bounds),
118        &[],
119        None,
120        None,
121    );
122}
123
124// ── ReLU6 ────────────────────────────────────────────────────────────────────
125
126/// Forward: y = clamp(x, 0, 6)
127#[kernel(backward = Relu6Backward)]
128pub fn relu6_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
129    x_ptr: T::Pointer<D>,
130    y_ptr: T::Pointer<D>,
131    n_elements: i32,
132) where
133    T::I32Tensor: types::Tensor<i32, 1>,
134    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
135    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
136{
137    let pid = T::program_id(Axis::X);
138    let block_start = pid * BLOCK_SIZE;
139    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
140    let in_bounds = offsets.lt(n_elements);
141
142    let x = T::load(
143        x_ptr.add_offsets(offsets),
144        Some(in_bounds),
145        None,
146        &[],
147        None,
148        None,
149        None,
150        false,
151    );
152    let lo = T::zeros_like(x);
153    let hi = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
154    let y = T::clamp(x, lo, hi);
155    T::store(
156        y_ptr.add_offsets(offsets),
157        y,
158        Some(in_bounds),
159        &[],
160        None,
161        None,
162    );
163}
164
165/// Backward: dx = dy if 0 < x < 6 else 0
166#[kernel]
167pub fn relu6_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
168    dy_ptr: T::Pointer<D>,
169    x_ptr: T::Pointer<D>,
170    dx_ptr: T::Pointer<D>,
171    n_elements: i32,
172) where
173    T::I32Tensor: types::Tensor<i32, 1>,
174    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
175    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
176{
177    let pid = T::program_id(Axis::X);
178    let block_start = pid * BLOCK_SIZE;
179    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
180    let in_bounds = offsets.lt(n_elements);
181
182    let dy = T::load(
183        dy_ptr.add_offsets(offsets),
184        Some(in_bounds),
185        None,
186        &[],
187        None,
188        None,
189        None,
190        false,
191    );
192    let x = T::load(
193        x_ptr.add_offsets(offsets),
194        Some(in_bounds),
195        None,
196        &[],
197        None,
198        None,
199        None,
200        false,
201    );
202
203    let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
204    // min(x, 6-x) > 0  ≡  0 < x < 6
205    let in_range = T::gt(T::minimum(x, six - x), T::zeros_like(x));
206    let dx = T::where_(in_range, dy, T::zeros_like(dy));
207    T::store(
208        dx_ptr.add_offsets(offsets),
209        dx,
210        Some(in_bounds),
211        &[],
212        None,
213        None,
214    );
215}
216
217// ── Hardsigmoid ──────────────────────────────────────────────────────────────
218
219/// Forward: y = clamp((x + 3) / 6, 0, 1)
220#[kernel(backward = HardsigmoidBackward)]
221pub fn hardsigmoid_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
222    x_ptr: T::Pointer<D>,
223    y_ptr: T::Pointer<D>,
224    n_elements: i32,
225) where
226    T::I32Tensor: types::Tensor<i32, 1>,
227    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
228    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
229{
230    let pid = T::program_id(Axis::X);
231    let block_start = pid * BLOCK_SIZE;
232    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
233    let in_bounds = offsets.lt(n_elements);
234
235    let x = T::load(
236        x_ptr.add_offsets(offsets),
237        Some(in_bounds),
238        None,
239        &[],
240        None,
241        None,
242        None,
243        false,
244    );
245    let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
246    let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
247    let lo = T::zeros_like(x);
248    let hi = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
249    let y = T::clamp((x + three) / six, lo, hi);
250    T::store(
251        y_ptr.add_offsets(offsets),
252        y,
253        Some(in_bounds),
254        &[],
255        None,
256        None,
257    );
258}
259
260/// Backward: dx = dy/6 if |x| < 3 else 0
261#[kernel]
262pub fn hardsigmoid_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
263    dy_ptr: T::Pointer<D>,
264    x_ptr: T::Pointer<D>,
265    dx_ptr: T::Pointer<D>,
266    n_elements: i32,
267) where
268    T::I32Tensor: types::Tensor<i32, 1>,
269    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
270    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
271{
272    let pid = T::program_id(Axis::X);
273    let block_start = pid * BLOCK_SIZE;
274    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
275    let in_bounds = offsets.lt(n_elements);
276
277    let dy = T::load(
278        dy_ptr.add_offsets(offsets),
279        Some(in_bounds),
280        None,
281        &[],
282        None,
283        None,
284        None,
285        false,
286    );
287    let x = T::load(
288        x_ptr.add_offsets(offsets),
289        Some(in_bounds),
290        None,
291        &[],
292        None,
293        None,
294        None,
295        false,
296    );
297
298    let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
299    let in_range = T::lt(T::abs(x), three); // |x| < 3
300    let sixth = T::full(&[BLOCK_SIZE], D::from_f64(1.0 / 6.0));
301    let dx = T::where_(in_range, dy * sixth, T::zeros_like(dy));
302    T::store(
303        dx_ptr.add_offsets(offsets),
304        dx,
305        Some(in_bounds),
306        &[],
307        None,
308        None,
309    );
310}
311
312// ── Hardswish ────────────────────────────────────────────────────────────────
313
314/// Forward: y = x * clamp((x + 3) / 6, 0, 1)
315#[kernel(backward = HardswishBackward)]
316pub fn hardswish_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
317    x_ptr: T::Pointer<D>,
318    y_ptr: T::Pointer<D>,
319    n_elements: i32,
320) where
321    T::I32Tensor: types::Tensor<i32, 1>,
322    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
323    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
324{
325    let pid = T::program_id(Axis::X);
326    let block_start = pid * BLOCK_SIZE;
327    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
328    let in_bounds = offsets.lt(n_elements);
329
330    let x = T::load(
331        x_ptr.add_offsets(offsets),
332        Some(in_bounds),
333        None,
334        &[],
335        None,
336        None,
337        None,
338        false,
339    );
340    let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
341    let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
342    let lo = T::zeros_like(x);
343    let hi = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
344    let hs = T::clamp((x + three) / six, lo, hi);
345    let y = x * hs;
346    T::store(
347        y_ptr.add_offsets(offsets),
348        y,
349        Some(in_bounds),
350        &[],
351        None,
352        None,
353    );
354}
355
356/// Backward:
357///   dx = 0           if x <= -3
358///   dx = dy          if x >= 3
359///   dx = dy*(2x+3)/6 otherwise
360#[kernel]
361pub fn hardswish_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
362    dy_ptr: T::Pointer<D>,
363    x_ptr: T::Pointer<D>,
364    dx_ptr: T::Pointer<D>,
365    n_elements: i32,
366) where
367    T::I32Tensor: types::Tensor<i32, 1>,
368    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
369    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
370{
371    let pid = T::program_id(Axis::X);
372    let block_start = pid * BLOCK_SIZE;
373    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
374    let in_bounds = offsets.lt(n_elements);
375
376    let dy = T::load(
377        dy_ptr.add_offsets(offsets),
378        Some(in_bounds),
379        None,
380        &[],
381        None,
382        None,
383        None,
384        false,
385    );
386    let x = T::load(
387        x_ptr.add_offsets(offsets),
388        Some(in_bounds),
389        None,
390        &[],
391        None,
392        None,
393        None,
394        false,
395    );
396
397    let three = T::full(&[BLOCK_SIZE], D::from_f64(3.0));
398    let neg_three = T::full(&[BLOCK_SIZE], D::from_f64(-3.0));
399    let six = T::full(&[BLOCK_SIZE], D::from_f64(6.0));
400    let two = T::full(&[BLOCK_SIZE], D::from_f64(2.0));
401
402    let x_le_neg3 = T::le(x, neg_three);
403    let x_ge_3 = T::ge(x, three);
404    let dx_mid = dy * (two * x + three) / six;
405
406    // Build from inner outward: start with mid, then override boundary regions.
407    let dx_not_lo = T::where_(x_ge_3, dy, dx_mid);
408    let dx = T::where_(x_le_neg3, T::zeros_like(dy), dx_not_lo);
409    T::store(
410        dx_ptr.add_offsets(offsets),
411        dx,
412        Some(in_bounds),
413        &[],
414        None,
415        None,
416    );
417}
418
419// ── Hardshrink ───────────────────────────────────────────────────────────────
420
421/// Forward: y = x if |x| > lambda else 0
422#[kernel(backward = HardshrinkBackward)]
423pub fn hardshrink_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
424    x_ptr: T::Pointer<D>,
425    y_ptr: T::Pointer<D>,
426    n_elements: i32,
427    lambda: f32,
428) where
429    T::I32Tensor: types::Tensor<i32, 1>,
430    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
431    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
432{
433    let pid = T::program_id(Axis::X);
434    let block_start = pid * BLOCK_SIZE;
435    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
436    let in_bounds = offsets.lt(n_elements);
437
438    let x = T::load(
439        x_ptr.add_offsets(offsets),
440        Some(in_bounds),
441        None,
442        &[],
443        None,
444        None,
445        None,
446        false,
447    );
448    let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
449    let outside = T::gt(T::abs(x), lam);
450    let y = T::where_(outside, x, T::zeros_like(x));
451    T::store(
452        y_ptr.add_offsets(offsets),
453        y,
454        Some(in_bounds),
455        &[],
456        None,
457        None,
458    );
459}
460
461/// Backward: dx = dy if |x| > lambda else 0
462#[kernel]
463pub fn hardshrink_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
464    dy_ptr: T::Pointer<D>,
465    x_ptr: T::Pointer<D>,
466    dx_ptr: T::Pointer<D>,
467    n_elements: i32,
468    lambda: f32,
469) where
470    T::I32Tensor: types::Tensor<i32, 1>,
471    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
472    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
473{
474    let pid = T::program_id(Axis::X);
475    let block_start = pid * BLOCK_SIZE;
476    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
477    let in_bounds = offsets.lt(n_elements);
478
479    let dy = T::load(
480        dy_ptr.add_offsets(offsets),
481        Some(in_bounds),
482        None,
483        &[],
484        None,
485        None,
486        None,
487        false,
488    );
489    let x = T::load(
490        x_ptr.add_offsets(offsets),
491        Some(in_bounds),
492        None,
493        &[],
494        None,
495        None,
496        None,
497        false,
498    );
499    let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
500    let outside = T::gt(T::abs(x), lam);
501    let dx = T::where_(outside, dy, T::zeros_like(dy));
502    T::store(
503        dx_ptr.add_offsets(offsets),
504        dx,
505        Some(in_bounds),
506        &[],
507        None,
508        None,
509    );
510}
511
512pub struct HardtanhOp<D: Float> {
513    pub forward: HardtanhForward<D>,
514    pub backward: HardtanhBackward<D>,
515}
516
517pub struct Relu6Op<D: Float> {
518    pub forward: Relu6Forward<D>,
519    pub backward: Relu6Backward<D>,
520}
521
522pub struct HardsigmoidOp<D: Float> {
523    pub forward: HardsigmoidForward<D>,
524    pub backward: HardsigmoidBackward<D>,
525}
526
527pub struct HardswishOp<D: Float> {
528    pub forward: HardswishForward<D>,
529    pub backward: HardswishBackward<D>,
530}
531
532pub struct HardshrinkOp<D: Float> {
533    pub forward: HardshrinkForward<D>,
534    pub backward: HardshrinkBackward<D>,
535}