Skip to main content

teeny_kernels/nn/activation/
extra.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//! Additional activation kernels: Swish, PRelu, LogSoftmax, Hardmax,
18//! ThresholdedRelu, Shrink.
19
20#![allow(non_snake_case)]
21
22use teeny_macros::kernel;
23use teeny_triton::triton::{
24    types::{AddOffsets, Comparison},
25    *,
26};
27
28// ── Swish (= SiLU: x * sigmoid(x)) ───────────────────────────────────────────
29
30/// Forward: y = x * sigmoid(x) = x / (1 + exp(-x))
31#[kernel]
32pub fn swish_forward<T: Triton, const BLOCK_SIZE: i32>(
33    x_ptr: T::Pointer<f32>,
34    y_ptr: T::Pointer<f32>,
35    n_elements: i32,
36) where
37    T::I32Tensor: types::Tensor<i32, 1>,
38    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
39    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
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    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 one = T::full::<f32>(&[BLOCK_SIZE], 1.0_f32);
56    let neg1 = T::full::<f32>(&[BLOCK_SIZE], -1.0_f32);
57    let sig = one / (one + T::exp(neg1 * x));
58    let y = x * sig;
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 = (sigmoid(x) + x * sigmoid(x) * (1 - sigmoid(x))) * dy
70///             = (sig + x * sig * (1 - sig)) * dy
71#[kernel]
72pub fn swish_backward<T: Triton, const BLOCK_SIZE: i32>(
73    dy_ptr: T::Pointer<f32>,
74    x_ptr: T::Pointer<f32>,
75    dx_ptr: T::Pointer<f32>,
76    n_elements: i32,
77) where
78    T::I32Tensor: types::Tensor<i32, 1>,
79    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
80    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
81{
82    let pid = T::program_id(Axis::X);
83    let block_start = pid * BLOCK_SIZE;
84    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
85    let in_bounds = offsets.lt(n_elements);
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 one = T::full::<f32>(&[BLOCK_SIZE], 1.0_f32);
107    let neg1 = T::full::<f32>(&[BLOCK_SIZE], -1.0_f32);
108    let sig = one / (one + T::exp(neg1 * x));
109    let dx = (sig + x * sig * (one - sig)) * dy;
110    T::store(
111        dx_ptr.add_offsets(offsets),
112        dx,
113        Some(in_bounds),
114        &[],
115        None,
116        None,
117    );
118}
119
120impl teeny_core::model::RuntimeOp for SwishForward {
121    fn n_activation_inputs(&self) -> usize {
122        1
123    }
124    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
125        vec![]
126    }
127    fn pack_args(
128        &self,
129        inputs: &[(teeny_core::model::RawPtr, &[usize])],
130        _: &[teeny_core::model::RawPtr],
131        output: teeny_core::model::RawPtr,
132        output_shape: &[usize],
133        _: i32,
134        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
135    ) {
136        let n: usize = output_shape.iter().product();
137        visitor.visit_ptr(inputs[0].0);
138        visitor.visit_ptr(output);
139        visitor.visit_i32(n as i32);
140    }
141    fn block(&self) -> [u32; 3] {
142        [self.block_size as u32, 1, 1]
143    }
144    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
145        let n: usize = output_shape.iter().product();
146        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
147    }
148    #[cfg(feature = "training")]
149    fn has_backward(&self) -> bool {
150        true
151    }
152    #[cfg(feature = "training")]
153    fn pack_backward_args(
154        &self,
155        inputs: &[(teeny_core::model::RawPtr, &[usize])],
156        _: &[teeny_core::model::RawPtr],
157        _: teeny_core::model::RawPtr,
158        output_shape: &[usize],
159        grad_output: teeny_core::model::RawPtr,
160        _: i32,
161        grad_inputs: &[teeny_core::model::RawPtr],
162        _: &[teeny_core::model::RawPtr],
163        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
164    ) {
165        let n: usize = output_shape.iter().product();
166        visitor.visit_ptr(grad_output);
167        visitor.visit_ptr(inputs[0].0);
168        visitor.visit_ptr(grad_inputs[0]);
169        visitor.visit_i32(n as i32);
170    }
171    #[cfg(feature = "training")]
172    fn backward_block(&self) -> [u32; 3] {
173        [self.block_size as u32, 1, 1]
174    }
175    #[cfg(feature = "training")]
176    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
177        let n: usize = output_shape.iter().product();
178        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
179    }
180}
181
182// ── PRelu (2-input: x, slope) ─────────────────────────────────────────────────
183
184/// Forward: y = max(0, x) + slope * min(0, x)
185/// The slope tensor has the same shape as x (or broadcastable; kernel assumes same shape here).
186#[kernel]
187pub fn prelu_forward<T: Triton, const BLOCK_SIZE: i32>(
188    x_ptr: T::Pointer<f32>,
189    slope_ptr: T::Pointer<f32>,
190    y_ptr: T::Pointer<f32>,
191    n_elements: i32,
192) where
193    T::I32Tensor: types::Tensor<i32, 1>,
194    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
195    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
196{
197    let pid = T::program_id(Axis::X);
198    let block_start = pid * BLOCK_SIZE;
199    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
200    let in_bounds = offsets.lt(n_elements);
201    let x = T::load(
202        x_ptr.add_offsets(offsets),
203        Some(in_bounds),
204        None,
205        &[],
206        None,
207        None,
208        None,
209        false,
210    );
211    let slope = T::load(
212        slope_ptr.add_offsets(offsets),
213        Some(in_bounds),
214        None,
215        &[],
216        None,
217        None,
218        None,
219        false,
220    );
221    let zero = T::zeros_like(x);
222    let pos = T::maximum(x, zero);
223    let neg = slope * T::minimum(x, zero);
224    let y = pos + neg;
225    T::store(
226        y_ptr.add_offsets(offsets),
227        y,
228        Some(in_bounds),
229        &[],
230        None,
231        None,
232    );
233}
234
235/// Backward: dx = dy if x >= 0 else slope * dy;
236///           dslope = dy * min(x, 0) = dy * x if x < 0 else 0
237#[kernel]
238pub fn prelu_backward<T: Triton, const BLOCK_SIZE: i32>(
239    dy_ptr: T::Pointer<f32>,
240    x_ptr: T::Pointer<f32>,
241    slope_ptr: T::Pointer<f32>,
242    dx_ptr: T::Pointer<f32>,
243    dslope_ptr: T::Pointer<f32>,
244    n_elements: i32,
245) where
246    T::I32Tensor: types::Tensor<i32, 1>,
247    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
248    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
249{
250    let pid = T::program_id(Axis::X);
251    let block_start = pid * BLOCK_SIZE;
252    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
253    let in_bounds = offsets.lt(n_elements);
254    let dy = T::load(
255        dy_ptr.add_offsets(offsets),
256        Some(in_bounds),
257        None,
258        &[],
259        None,
260        None,
261        None,
262        false,
263    );
264    let x = T::load(
265        x_ptr.add_offsets(offsets),
266        Some(in_bounds),
267        None,
268        &[],
269        None,
270        None,
271        None,
272        false,
273    );
274    let slope = T::load(
275        slope_ptr.add_offsets(offsets),
276        Some(in_bounds),
277        None,
278        &[],
279        None,
280        None,
281        None,
282        false,
283    );
284    let zero = T::zeros_like(x);
285    let x_pos = T::ge(x, zero);
286    let dx = T::where_(x_pos, dy, slope * dy);
287    let dslope = T::where_(x_pos, zero, x * dy);
288    T::store(
289        dx_ptr.add_offsets(offsets),
290        dx,
291        Some(in_bounds),
292        &[],
293        None,
294        None,
295    );
296    T::store(
297        dslope_ptr.add_offsets(offsets),
298        dslope,
299        Some(in_bounds),
300        &[],
301        None,
302        None,
303    );
304}
305
306impl teeny_core::model::RuntimeOp for PreluForward {
307    fn n_activation_inputs(&self) -> usize {
308        2
309    }
310    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
311        vec![]
312    }
313    fn pack_args(
314        &self,
315        inputs: &[(teeny_core::model::RawPtr, &[usize])],
316        _: &[teeny_core::model::RawPtr],
317        output: teeny_core::model::RawPtr,
318        output_shape: &[usize],
319        _: i32,
320        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
321    ) {
322        let n: usize = output_shape.iter().product();
323        visitor.visit_ptr(inputs[0].0); // x
324        visitor.visit_ptr(inputs[1].0); // slope
325        visitor.visit_ptr(output);
326        visitor.visit_i32(n as i32);
327    }
328    fn block(&self) -> [u32; 3] {
329        [self.block_size as u32, 1, 1]
330    }
331    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
332        let n: usize = output_shape.iter().product();
333        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
334    }
335    #[cfg(feature = "training")]
336    fn has_backward(&self) -> bool {
337        true
338    }
339    #[cfg(feature = "training")]
340    fn pack_backward_args(
341        &self,
342        inputs: &[(teeny_core::model::RawPtr, &[usize])],
343        _: &[teeny_core::model::RawPtr],
344        _: teeny_core::model::RawPtr,
345        output_shape: &[usize],
346        grad_output: teeny_core::model::RawPtr,
347        _: i32,
348        grad_inputs: &[teeny_core::model::RawPtr],
349        _: &[teeny_core::model::RawPtr],
350        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
351    ) {
352        let n: usize = output_shape.iter().product();
353        visitor.visit_ptr(grad_output);
354        visitor.visit_ptr(inputs[0].0); // x
355        visitor.visit_ptr(inputs[1].0); // slope
356        visitor.visit_ptr(grad_inputs[0]); // dx
357        visitor.visit_ptr(grad_inputs[1]); // dslope
358        visitor.visit_i32(n as i32);
359    }
360    #[cfg(feature = "training")]
361    fn backward_block(&self) -> [u32; 3] {
362        [self.block_size as u32, 1, 1]
363    }
364    #[cfg(feature = "training")]
365    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
366        let n: usize = output_shape.iter().product();
367        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
368    }
369}
370
371// ── ThresholdedRelu ───────────────────────────────────────────────────────────
372
373/// Forward: y = x if x > alpha else 0
374#[kernel]
375pub fn thresholded_relu_forward<T: Triton, const BLOCK_SIZE: i32>(
376    x_ptr: T::Pointer<f32>,
377    y_ptr: T::Pointer<f32>,
378    n_elements: i32,
379    alpha: f32,
380) where
381    T::I32Tensor: types::Tensor<i32, 1>,
382    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
383    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
384{
385    let pid = T::program_id(Axis::X);
386    let block_start = pid * BLOCK_SIZE;
387    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
388    let in_bounds = offsets.lt(n_elements);
389    let x = T::load(
390        x_ptr.add_offsets(offsets),
391        Some(in_bounds),
392        None,
393        &[],
394        None,
395        None,
396        None,
397        false,
398    );
399    let alpha_t = T::full::<f32>(&[BLOCK_SIZE], alpha);
400    let above = T::gt(x, alpha_t);
401    let y = T::where_(above, x, T::zeros_like(x));
402    T::store(
403        y_ptr.add_offsets(offsets),
404        y,
405        Some(in_bounds),
406        &[],
407        None,
408        None,
409    );
410}
411
412/// Backward: dx = dy if x > alpha else 0
413#[kernel]
414pub fn thresholded_relu_backward<T: Triton, const BLOCK_SIZE: i32>(
415    dy_ptr: T::Pointer<f32>,
416    x_ptr: T::Pointer<f32>,
417    dx_ptr: T::Pointer<f32>,
418    n_elements: i32,
419    alpha: f32,
420) where
421    T::I32Tensor: types::Tensor<i32, 1>,
422    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
423    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
424{
425    let pid = T::program_id(Axis::X);
426    let block_start = pid * BLOCK_SIZE;
427    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
428    let in_bounds = offsets.lt(n_elements);
429    let dy = T::load(
430        dy_ptr.add_offsets(offsets),
431        Some(in_bounds),
432        None,
433        &[],
434        None,
435        None,
436        None,
437        false,
438    );
439    let x = T::load(
440        x_ptr.add_offsets(offsets),
441        Some(in_bounds),
442        None,
443        &[],
444        None,
445        None,
446        None,
447        false,
448    );
449    let alpha_t = T::full::<f32>(&[BLOCK_SIZE], alpha);
450    let above = T::gt(x, alpha_t);
451    let dx = T::where_(above, dy, T::zeros_like(dy));
452    T::store(
453        dx_ptr.add_offsets(offsets),
454        dx,
455        Some(in_bounds),
456        &[],
457        None,
458        None,
459    );
460}
461
462/// RuntimeOp wrapper for ThresholdedRelu that stores the alpha scalar.
463pub struct ThresholdedReluRuntimeOp {
464    pub kernel: ThresholdedReluForward,
465    pub backward_kernel: ThresholdedReluBackward,
466    pub alpha: f32,
467}
468
469impl ThresholdedReluRuntimeOp {
470    pub fn new(block_size: i32, alpha: f32) -> Self {
471        Self {
472            kernel: ThresholdedReluForward::new(block_size),
473            backward_kernel: ThresholdedReluBackward::new(block_size),
474            alpha,
475        }
476    }
477    pub fn forward_source(&self) -> &str {
478        &self.kernel.source
479    }
480    pub fn backward_source(&self) -> &str {
481        &self.backward_kernel.source
482    }
483    pub fn kernel_name(&self) -> &str {
484        self.kernel.name
485    }
486}
487
488impl teeny_core::model::RuntimeOp for ThresholdedReluRuntimeOp {
489    fn n_activation_inputs(&self) -> usize {
490        1
491    }
492    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
493        vec![]
494    }
495    fn pack_args(
496        &self,
497        inputs: &[(teeny_core::model::RawPtr, &[usize])],
498        _: &[teeny_core::model::RawPtr],
499        output: teeny_core::model::RawPtr,
500        output_shape: &[usize],
501        _: i32,
502        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
503    ) {
504        let n: usize = output_shape.iter().product();
505        visitor.visit_ptr(inputs[0].0);
506        visitor.visit_ptr(output);
507        visitor.visit_i32(n as i32);
508        visitor.visit_f32(self.alpha);
509    }
510    fn block(&self) -> [u32; 3] {
511        [self.kernel.block_size as u32, 1, 1]
512    }
513    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
514        let n: usize = output_shape.iter().product();
515        [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
516    }
517    #[cfg(feature = "training")]
518    fn has_backward(&self) -> bool {
519        true
520    }
521    #[cfg(feature = "training")]
522    fn pack_backward_args(
523        &self,
524        inputs: &[(teeny_core::model::RawPtr, &[usize])],
525        _: &[teeny_core::model::RawPtr],
526        _: teeny_core::model::RawPtr,
527        output_shape: &[usize],
528        grad_output: teeny_core::model::RawPtr,
529        _: i32,
530        grad_inputs: &[teeny_core::model::RawPtr],
531        _: &[teeny_core::model::RawPtr],
532        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
533    ) {
534        let n: usize = output_shape.iter().product();
535        visitor.visit_ptr(grad_output);
536        visitor.visit_ptr(inputs[0].0);
537        visitor.visit_ptr(grad_inputs[0]);
538        visitor.visit_i32(n as i32);
539        visitor.visit_f32(self.alpha);
540    }
541    #[cfg(feature = "training")]
542    fn backward_block(&self) -> [u32; 3] {
543        [self.kernel.block_size as u32, 1, 1]
544    }
545    #[cfg(feature = "training")]
546    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
547        let n: usize = output_shape.iter().product();
548        [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
549    }
550}
551
552// ── Shrink ────────────────────────────────────────────────────────────────────
553
554/// Forward: y = x - bias if x > lambd, x + bias if x < -lambd, else 0
555#[kernel]
556pub fn shrink_forward<T: Triton, const BLOCK_SIZE: i32>(
557    x_ptr: T::Pointer<f32>,
558    y_ptr: T::Pointer<f32>,
559    n_elements: i32,
560    lambd: f32,
561    bias: f32,
562) where
563    T::I32Tensor: types::Tensor<i32, 1>,
564    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
565    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
566{
567    let pid = T::program_id(Axis::X);
568    let block_start = pid * BLOCK_SIZE;
569    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
570    let in_bounds = offsets.lt(n_elements);
571    let x = T::load(
572        x_ptr.add_offsets(offsets),
573        Some(in_bounds),
574        None,
575        &[],
576        None,
577        None,
578        None,
579        false,
580    );
581    let lam = T::full::<f32>(&[BLOCK_SIZE], lambd);
582    let neg_lam = T::full::<f32>(&[BLOCK_SIZE], -lambd);
583    let b = T::full::<f32>(&[BLOCK_SIZE], bias);
584    let x_gt = T::gt(x, lam);
585    let x_lt = T::lt(x, neg_lam);
586    let y_upper = x - b;
587    let y_lower = x + b;
588    let y_mid = T::where_(x_lt, y_lower, T::zeros_like(x));
589    let y = T::where_(x_gt, y_upper, y_mid);
590    T::store(
591        y_ptr.add_offsets(offsets),
592        y,
593        Some(in_bounds),
594        &[],
595        None,
596        None,
597    );
598}
599
600/// Backward: dx = dy if |x| > lambd else 0
601#[kernel]
602pub fn shrink_backward<T: Triton, const BLOCK_SIZE: i32>(
603    dy_ptr: T::Pointer<f32>,
604    x_ptr: T::Pointer<f32>,
605    dx_ptr: T::Pointer<f32>,
606    n_elements: i32,
607    lambd: f32,
608) where
609    T::I32Tensor: types::Tensor<i32, 1>,
610    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
611    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
612{
613    let pid = T::program_id(Axis::X);
614    let block_start = pid * BLOCK_SIZE;
615    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
616    let in_bounds = offsets.lt(n_elements);
617    let dy = T::load(
618        dy_ptr.add_offsets(offsets),
619        Some(in_bounds),
620        None,
621        &[],
622        None,
623        None,
624        None,
625        false,
626    );
627    let x = T::load(
628        x_ptr.add_offsets(offsets),
629        Some(in_bounds),
630        None,
631        &[],
632        None,
633        None,
634        None,
635        false,
636    );
637    let lam = T::full::<f32>(&[BLOCK_SIZE], lambd);
638    let outside = T::gt(T::abs(x), lam);
639    let dx = T::where_(outside, dy, T::zeros_like(dy));
640    T::store(
641        dx_ptr.add_offsets(offsets),
642        dx,
643        Some(in_bounds),
644        &[],
645        None,
646        None,
647    );
648}
649
650/// RuntimeOp wrapper for Shrink that stores lambd and bias.
651pub struct ShrinkRuntimeOp {
652    pub kernel: ShrinkForward,
653    pub backward_kernel: ShrinkBackward,
654    pub lambd: f32,
655    pub bias: f32,
656}
657
658impl ShrinkRuntimeOp {
659    pub fn new(block_size: i32, lambd: f32, bias: f32) -> Self {
660        Self {
661            kernel: ShrinkForward::new(block_size),
662            backward_kernel: ShrinkBackward::new(block_size),
663            lambd,
664            bias,
665        }
666    }
667    pub fn forward_source(&self) -> &str {
668        &self.kernel.source
669    }
670    pub fn backward_source(&self) -> &str {
671        &self.backward_kernel.source
672    }
673    pub fn kernel_name(&self) -> &str {
674        self.kernel.name
675    }
676}
677
678impl teeny_core::model::RuntimeOp for ShrinkRuntimeOp {
679    fn n_activation_inputs(&self) -> usize {
680        1
681    }
682    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
683        vec![]
684    }
685    fn pack_args(
686        &self,
687        inputs: &[(teeny_core::model::RawPtr, &[usize])],
688        _: &[teeny_core::model::RawPtr],
689        output: teeny_core::model::RawPtr,
690        output_shape: &[usize],
691        _: i32,
692        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
693    ) {
694        let n: usize = output_shape.iter().product();
695        visitor.visit_ptr(inputs[0].0);
696        visitor.visit_ptr(output);
697        visitor.visit_i32(n as i32);
698        visitor.visit_f32(self.lambd);
699        visitor.visit_f32(self.bias);
700    }
701    fn block(&self) -> [u32; 3] {
702        [self.kernel.block_size as u32, 1, 1]
703    }
704    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
705        let n: usize = output_shape.iter().product();
706        [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
707    }
708    #[cfg(feature = "training")]
709    fn has_backward(&self) -> bool {
710        true
711    }
712    #[cfg(feature = "training")]
713    fn pack_backward_args(
714        &self,
715        inputs: &[(teeny_core::model::RawPtr, &[usize])],
716        _: &[teeny_core::model::RawPtr],
717        _: teeny_core::model::RawPtr,
718        output_shape: &[usize],
719        grad_output: teeny_core::model::RawPtr,
720        _: i32,
721        grad_inputs: &[teeny_core::model::RawPtr],
722        _: &[teeny_core::model::RawPtr],
723        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
724    ) {
725        let n: usize = output_shape.iter().product();
726        visitor.visit_ptr(grad_output);
727        visitor.visit_ptr(inputs[0].0);
728        visitor.visit_ptr(grad_inputs[0]);
729        visitor.visit_i32(n as i32);
730        visitor.visit_f32(self.lambd);
731    }
732    #[cfg(feature = "training")]
733    fn backward_block(&self) -> [u32; 3] {
734        [self.kernel.block_size as u32, 1, 1]
735    }
736    #[cfg(feature = "training")]
737    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
738        let n: usize = output_shape.iter().product();
739        [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
740    }
741}
742
743// ── LogSoftmax (row-wise) ─────────────────────────────────────────────────────
744//
745// Grid: one CTA per row. BLOCK_SIZE must equal n_cols (power of 2).
746
747/// Forward: y = x - log(sum(exp(x)))  [numerically stable: subtract max first]
748#[kernel]
749pub fn log_softmax_forward<T: Triton, const BLOCK_SIZE: i32>(
750    x_ptr: T::Pointer<f32>,
751    y_ptr: T::Pointer<f32>,
752    _n_rows: i32,
753    n_cols: i32,
754) where
755    T::I32Tensor: types::Tensor<i32, 1>,
756    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
757    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
758{
759    let pid = T::program_id(Axis::X);
760    let row_offset = pid * n_cols;
761    let col_offsets = T::arange(0, BLOCK_SIZE);
762    let offsets = col_offsets + row_offset;
763    let x = T::load(
764        x_ptr.add_offsets(offsets),
765        None,
766        None,
767        &[],
768        None,
769        None,
770        None,
771        false,
772    );
773    let m = T::max(x, Some(0), true); // max for numerical stability
774    let x_m = x - m;
775    let log_sum = T::log(T::sum(T::exp(x_m), Some(0), true));
776    let y = x_m - log_sum;
777    T::store(y_ptr.add_offsets(offsets), y, None, &[], None, None);
778}
779
780/// Backward: dx = dy - softmax(x) * sum(dy)
781#[kernel]
782pub fn log_softmax_backward<T: Triton, const BLOCK_SIZE: i32>(
783    dy_ptr: T::Pointer<f32>,
784    y_ptr: T::Pointer<f32>,
785    dx_ptr: T::Pointer<f32>,
786    _n_rows: i32,
787    n_cols: i32,
788) where
789    T::I32Tensor: types::Tensor<i32, 1>,
790    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
791    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
792{
793    let pid = T::program_id(Axis::X);
794    let row_offset = pid * n_cols;
795    let col_offsets = T::arange(0, BLOCK_SIZE);
796    let offsets = col_offsets + row_offset;
797    let dy = T::load(
798        dy_ptr.add_offsets(offsets),
799        None,
800        None,
801        &[],
802        None,
803        None,
804        None,
805        false,
806    );
807    let y = T::load(
808        y_ptr.add_offsets(offsets),
809        None,
810        None,
811        &[],
812        None,
813        None,
814        None,
815        false,
816    );
817    // softmax(x) = exp(log_softmax(x))
818    let sm = T::exp(y);
819    let sum_dy = T::sum(dy, Some(0), true);
820    let dx = dy - sm * sum_dy;
821    T::store(dx_ptr.add_offsets(offsets), dx, None, &[], None, None);
822}
823
824impl teeny_core::model::RuntimeOp for LogSoftmaxForward {
825    fn n_activation_inputs(&self) -> usize {
826        1
827    }
828    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
829        vec![]
830    }
831    fn pack_args(
832        &self,
833        inputs: &[(teeny_core::model::RawPtr, &[usize])],
834        _: &[teeny_core::model::RawPtr],
835        output: teeny_core::model::RawPtr,
836        output_shape: &[usize],
837        _: i32,
838        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
839    ) {
840        let n_rows = output_shape.first().copied().unwrap_or(1) as i32;
841        let n_cols = output_shape.last().copied().unwrap_or(1) as i32;
842        visitor.visit_ptr(inputs[0].0);
843        visitor.visit_ptr(output);
844        visitor.visit_i32(n_rows);
845        visitor.visit_i32(n_cols);
846    }
847    fn block(&self) -> [u32; 3] {
848        [self.block_size as u32, 1, 1]
849    }
850    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
851        [output_shape.first().copied().unwrap_or(1) as u32, 1, 1]
852    }
853    #[cfg(feature = "training")]
854    fn has_backward(&self) -> bool {
855        true
856    }
857    #[cfg(feature = "training")]
858    fn pack_backward_args(
859        &self,
860        _: &[(teeny_core::model::RawPtr, &[usize])],
861        _: &[teeny_core::model::RawPtr],
862        output: teeny_core::model::RawPtr,
863        output_shape: &[usize],
864        grad_output: teeny_core::model::RawPtr,
865        _: i32,
866        grad_inputs: &[teeny_core::model::RawPtr],
867        _: &[teeny_core::model::RawPtr],
868        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
869    ) {
870        let n_rows = output_shape.first().copied().unwrap_or(1) as i32;
871        let n_cols = output_shape.last().copied().unwrap_or(1) as i32;
872        visitor.visit_ptr(grad_output);
873        visitor.visit_ptr(output);
874        visitor.visit_ptr(grad_inputs[0]);
875        visitor.visit_i32(n_rows);
876        visitor.visit_i32(n_cols);
877    }
878    #[cfg(feature = "training")]
879    fn backward_block(&self) -> [u32; 3] {
880        [self.block_size as u32, 1, 1]
881    }
882    #[cfg(feature = "training")]
883    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
884        [output_shape.first().copied().unwrap_or(1) as u32, 1, 1]
885    }
886}