Skip to main content

teeny_kernels/nn/activation/
misc.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// ── LeakyReLU ────────────────────────────────────────────────────────────────
27
28/// Forward: y = x if x > 0 else negative_slope * x
29#[kernel(backward = LeakyReluBackward)]
30pub fn leaky_relu_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    negative_slope: f32,
35) where
36    T::I32Tensor: types::Tensor<i32, 1>,
37    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
38    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
39{
40    let pid = T::program_id(Axis::X);
41    let block_start = pid * BLOCK_SIZE;
42    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
43    let in_bounds = offsets.lt(n_elements);
44
45    let x = T::load(
46        x_ptr.add_offsets(offsets),
47        Some(in_bounds),
48        None,
49        &[],
50        None,
51        None,
52        None,
53        false,
54    );
55    let slope = T::full(&[BLOCK_SIZE], D::from_f64(negative_slope as f64));
56    let x_pos = T::gt(x, T::zeros_like(x));
57    let y = T::where_(x_pos, x, slope * x);
58    T::store(
59        y_ptr.add_offsets(offsets),
60        y,
61        Some(in_bounds),
62        &[],
63        None,
64        None,
65    );
66}
67
68/// Backward: dx = dy if x > 0 else negative_slope * dy
69#[kernel]
70pub fn leaky_relu_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
71    dy_ptr: T::Pointer<D>,
72    x_ptr: T::Pointer<D>,
73    dx_ptr: T::Pointer<D>,
74    n_elements: i32,
75    negative_slope: f32,
76) where
77    T::I32Tensor: types::Tensor<i32, 1>,
78    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
79    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
80{
81    let pid = T::program_id(Axis::X);
82    let block_start = pid * BLOCK_SIZE;
83    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
84    let in_bounds = offsets.lt(n_elements);
85
86    let dy = T::load(
87        dy_ptr.add_offsets(offsets),
88        Some(in_bounds),
89        None,
90        &[],
91        None,
92        None,
93        None,
94        false,
95    );
96    let x = T::load(
97        x_ptr.add_offsets(offsets),
98        Some(in_bounds),
99        None,
100        &[],
101        None,
102        None,
103        None,
104        false,
105    );
106    let slope = T::full(&[BLOCK_SIZE], D::from_f64(negative_slope as f64));
107    let x_pos = T::gt(x, T::zeros_like(x));
108    let dx = T::where_(x_pos, dy, slope * dy);
109    T::store(
110        dx_ptr.add_offsets(offsets),
111        dx,
112        Some(in_bounds),
113        &[],
114        None,
115        None,
116    );
117}
118
119// ── Threshold ────────────────────────────────────────────────────────────────
120
121/// Forward: y = x if x > threshold else value
122#[kernel(backward = ThresholdBackward)]
123pub fn threshold_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
124    x_ptr: T::Pointer<D>,
125    y_ptr: T::Pointer<D>,
126    n_elements: i32,
127    threshold: f32,
128    value: f32,
129) where
130    T::I32Tensor: types::Tensor<i32, 1>,
131    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
132    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
133{
134    let pid = T::program_id(Axis::X);
135    let block_start = pid * BLOCK_SIZE;
136    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
137    let in_bounds = offsets.lt(n_elements);
138
139    let x = T::load(
140        x_ptr.add_offsets(offsets),
141        Some(in_bounds),
142        None,
143        &[],
144        None,
145        None,
146        None,
147        false,
148    );
149    let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
150    let val = T::full(&[BLOCK_SIZE], D::from_f64(value as f64));
151    let above = T::gt(x, thr);
152    let y = T::where_(above, x, val);
153    T::store(
154        y_ptr.add_offsets(offsets),
155        y,
156        Some(in_bounds),
157        &[],
158        None,
159        None,
160    );
161}
162
163/// Backward: dx = dy if x > threshold else 0
164#[kernel]
165pub fn threshold_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
166    dy_ptr: T::Pointer<D>,
167    x_ptr: T::Pointer<D>,
168    dx_ptr: T::Pointer<D>,
169    n_elements: i32,
170    threshold: f32,
171) where
172    T::I32Tensor: types::Tensor<i32, 1>,
173    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
174    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
175{
176    let pid = T::program_id(Axis::X);
177    let block_start = pid * BLOCK_SIZE;
178    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
179    let in_bounds = offsets.lt(n_elements);
180
181    let dy = T::load(
182        dy_ptr.add_offsets(offsets),
183        Some(in_bounds),
184        None,
185        &[],
186        None,
187        None,
188        None,
189        false,
190    );
191    let x = T::load(
192        x_ptr.add_offsets(offsets),
193        Some(in_bounds),
194        None,
195        &[],
196        None,
197        None,
198        None,
199        false,
200    );
201    let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
202    let above = T::gt(x, thr);
203    let dx = T::where_(above, dy, T::zeros_like(dy));
204    T::store(
205        dx_ptr.add_offsets(offsets),
206        dx,
207        Some(in_bounds),
208        &[],
209        None,
210        None,
211    );
212}
213
214// ── Softsign ─────────────────────────────────────────────────────────────────
215
216/// Forward: y = x / (1 + |x|)
217#[kernel(backward = SoftsignBackward)]
218pub fn softsign_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
219    x_ptr: T::Pointer<D>,
220    y_ptr: T::Pointer<D>,
221    n_elements: i32,
222) where
223    T::I32Tensor: types::Tensor<i32, 1>,
224    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
225    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
226{
227    let pid = T::program_id(Axis::X);
228    let block_start = pid * BLOCK_SIZE;
229    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
230    let in_bounds = offsets.lt(n_elements);
231
232    let x = T::load(
233        x_ptr.add_offsets(offsets),
234        Some(in_bounds),
235        None,
236        &[],
237        None,
238        None,
239        None,
240        false,
241    );
242    let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
243    let d = one + T::abs(x);
244    let y = x / d;
245    T::store(
246        y_ptr.add_offsets(offsets),
247        y,
248        Some(in_bounds),
249        &[],
250        None,
251        None,
252    );
253}
254
255/// Backward: dx = dy / (1 + |x|)²
256#[kernel]
257pub fn softsign_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
258    dy_ptr: T::Pointer<D>,
259    x_ptr: T::Pointer<D>,
260    dx_ptr: T::Pointer<D>,
261    n_elements: i32,
262) where
263    T::I32Tensor: types::Tensor<i32, 1>,
264    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
265    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
266{
267    let pid = T::program_id(Axis::X);
268    let block_start = pid * BLOCK_SIZE;
269    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
270    let in_bounds = offsets.lt(n_elements);
271
272    let dy = T::load(
273        dy_ptr.add_offsets(offsets),
274        Some(in_bounds),
275        None,
276        &[],
277        None,
278        None,
279        None,
280        false,
281    );
282    let x = T::load(
283        x_ptr.add_offsets(offsets),
284        Some(in_bounds),
285        None,
286        &[],
287        None,
288        None,
289        None,
290        false,
291    );
292    let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
293    let d = one + T::abs(x);
294    let dx = dy / (d * d);
295    T::store(
296        dx_ptr.add_offsets(offsets),
297        dx,
298        Some(in_bounds),
299        &[],
300        None,
301        None,
302    );
303}
304
305// ── Softshrink ───────────────────────────────────────────────────────────────
306
307/// Forward: y = x - lambda if x > lambda, x + lambda if x < -lambda, else 0
308#[kernel(backward = SoftshrinkBackward)]
309pub fn softshrink_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
310    x_ptr: T::Pointer<D>,
311    y_ptr: T::Pointer<D>,
312    n_elements: i32,
313    lambda: f32,
314) where
315    T::I32Tensor: types::Tensor<i32, 1>,
316    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
317    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
318{
319    let pid = T::program_id(Axis::X);
320    let block_start = pid * BLOCK_SIZE;
321    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
322    let in_bounds = offsets.lt(n_elements);
323
324    let x = T::load(
325        x_ptr.add_offsets(offsets),
326        Some(in_bounds),
327        None,
328        &[],
329        None,
330        None,
331        None,
332        false,
333    );
334    let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
335    let neg_lam = T::full(&[BLOCK_SIZE], D::from_f64(-(lambda as f64)));
336    let x_gt_lam = T::gt(x, lam);
337    let x_lt_neg = T::lt(x, neg_lam);
338    let y_upper = x - lam;
339    let y_lower = x + lam;
340    let y_mid = T::where_(x_lt_neg, y_lower, T::zeros_like(x));
341    let y = T::where_(x_gt_lam, y_upper, y_mid);
342    T::store(
343        y_ptr.add_offsets(offsets),
344        y,
345        Some(in_bounds),
346        &[],
347        None,
348        None,
349    );
350}
351
352/// Backward: dx = dy if |x| > lambda else 0
353#[kernel]
354pub fn softshrink_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
355    dy_ptr: T::Pointer<D>,
356    x_ptr: T::Pointer<D>,
357    dx_ptr: T::Pointer<D>,
358    n_elements: i32,
359    lambda: f32,
360) where
361    T::I32Tensor: types::Tensor<i32, 1>,
362    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
363    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
364{
365    let pid = T::program_id(Axis::X);
366    let block_start = pid * BLOCK_SIZE;
367    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
368    let in_bounds = offsets.lt(n_elements);
369
370    let dy = T::load(
371        dy_ptr.add_offsets(offsets),
372        Some(in_bounds),
373        None,
374        &[],
375        None,
376        None,
377        None,
378        false,
379    );
380    let x = T::load(
381        x_ptr.add_offsets(offsets),
382        Some(in_bounds),
383        None,
384        &[],
385        None,
386        None,
387        None,
388        false,
389    );
390    let lam = T::full(&[BLOCK_SIZE], D::from_f64(lambda as f64));
391    let outside = T::gt(T::abs(x), lam);
392    let dx = T::where_(outside, dy, T::zeros_like(dy));
393    T::store(
394        dx_ptr.add_offsets(offsets),
395        dx,
396        Some(in_bounds),
397        &[],
398        None,
399        None,
400    );
401}
402
403// ── Softplus ─────────────────────────────────────────────────────────────────
404
405/// Forward: y = (1/beta) * log(1 + exp(beta*x))
406///   For beta*x > threshold: y ≈ x (numerically safe pass-through)
407#[kernel(backward = SoftplusBackward)]
408pub fn softplus_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
409    x_ptr: T::Pointer<D>,
410    y_ptr: T::Pointer<D>,
411    n_elements: i32,
412    beta: f32,
413    threshold: f32,
414) where
415    T::I32Tensor: types::Tensor<i32, 1>,
416    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
417    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
418{
419    let pid = T::program_id(Axis::X);
420    let block_start = pid * BLOCK_SIZE;
421    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
422    let in_bounds = offsets.lt(n_elements);
423
424    let x = T::load(
425        x_ptr.add_offsets(offsets),
426        Some(in_bounds),
427        None,
428        &[],
429        None,
430        None,
431        None,
432        false,
433    );
434    let beta_t = T::full(&[BLOCK_SIZE], D::from_f64(beta as f64));
435    let inv_beta = T::full(&[BLOCK_SIZE], D::from_f64(1.0 / beta as f64));
436    let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
437    let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
438    let bx = beta_t * x;
439    let above_thr = T::gt(bx, thr);
440    let y_safe = inv_beta * T::log(one + T::exp(bx));
441    let y = T::where_(above_thr, x, y_safe);
442    T::store(
443        y_ptr.add_offsets(offsets),
444        y,
445        Some(in_bounds),
446        &[],
447        None,
448        None,
449    );
450}
451
452/// Backward: dx = dy * sigmoid(beta*x)
453///   For beta*x > threshold: dx ≈ dy
454#[kernel]
455pub fn softplus_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
456    dy_ptr: T::Pointer<D>,
457    x_ptr: T::Pointer<D>,
458    dx_ptr: T::Pointer<D>,
459    n_elements: i32,
460    beta: f32,
461    threshold: f32,
462) where
463    T::I32Tensor: types::Tensor<i32, 1>,
464    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
465    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
466{
467    let pid = T::program_id(Axis::X);
468    let block_start = pid * BLOCK_SIZE;
469    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
470    let in_bounds = offsets.lt(n_elements);
471
472    let dy = T::load(
473        dy_ptr.add_offsets(offsets),
474        Some(in_bounds),
475        None,
476        &[],
477        None,
478        None,
479        None,
480        false,
481    );
482    let x = T::load(
483        x_ptr.add_offsets(offsets),
484        Some(in_bounds),
485        None,
486        &[],
487        None,
488        None,
489        None,
490        false,
491    );
492    let beta_t = T::full(&[BLOCK_SIZE], D::from_f64(beta as f64));
493    let neg_beta = T::full(&[BLOCK_SIZE], D::from_f64(-(beta as f64)));
494    let thr = T::full(&[BLOCK_SIZE], D::from_f64(threshold as f64));
495    let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
496    let bx = beta_t * x;
497    let neg_bx = neg_beta * x;
498    let above_thr = T::gt(bx, thr);
499    // sigmoid(bx) = 1 / (1 + exp(-bx))
500    let dx_safe = dy * (one / (one + T::exp(neg_bx)));
501    let dx = T::where_(above_thr, dy, dx_safe);
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 LeakyReluOp<D: Float> {
513    pub forward: LeakyReluForward<D>,
514    pub backward: LeakyReluBackward<D>,
515}
516
517pub struct ThresholdOp<D: Float> {
518    pub forward: ThresholdForward<D>,
519    pub backward: ThresholdBackward<D>,
520}
521
522pub struct SoftsignOp<D: Float> {
523    pub forward: SoftsignForward<D>,
524    pub backward: SoftsignBackward<D>,
525}
526
527pub struct SoftshrinkOp<D: Float> {
528    pub forward: SoftshrinkForward<D>,
529    pub backward: SoftshrinkBackward<D>,
530}
531
532pub struct SoftplusOp<D: Float> {
533    pub forward: SoftplusForward<D>,
534    pub backward: SoftplusBackward<D>,
535}