Skip to main content

teeny_kernels/nn/tensor/
reduction.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//! Reduction kernels — each CTA handles one output element (one "row" of the
18//! flattened [outer, inner] view).  The caller is responsible for reshaping
19//! the input to `[n_outer, n_inner]` before invoking these kernels.
20//!
21//! Grid: `[n_outer, 1, 1]`
22//! Block: `[BLOCK_INNER, 1, 1]`
23
24#![allow(non_snake_case)]
25
26use teeny_core::dtype::{Float, Num};
27use teeny_macros::kernel;
28use teeny_triton::triton::{
29    types::{AddOffsets, Comparison},
30    *,
31};
32
33// ── Helper macro for reduction RuntimeOp ─────────────────────────────────────
34
35/// Standard reduction RuntimeOp: input shape [outer * inner], output shape [outer].
36/// pack_args: x_ptr, y_ptr, n_inner, n_outer
37macro_rules! impl_reduce_num_runtime_op {
38    ($Fwd:ident) => {
39        impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
40            fn n_activation_inputs(&self) -> usize {
41                1
42            }
43            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
44                vec![]
45            }
46            fn pack_args(
47                &self,
48                inputs: &[(teeny_core::model::RawPtr, &[usize])],
49                _: &[teeny_core::model::RawPtr],
50                output: teeny_core::model::RawPtr,
51                output_shape: &[usize],
52                _: i32,
53                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
54            ) {
55                // output_shape has been reduced; we need input_shape for n_inner.
56                // n_outer = product of output dims
57                // n_inner = product of input dims / n_outer
58                let n_outer: usize = output_shape.iter().product::<usize>().max(1);
59                let n_total: usize = inputs[0].1.iter().product();
60                let n_inner: usize = if n_outer > 0 {
61                    n_total / n_outer
62                } else {
63                    n_total
64                };
65                visitor.visit_ptr(inputs[0].0);
66                visitor.visit_ptr(output);
67                visitor.visit_i32(n_inner as i32);
68                visitor.visit_i32(n_outer as i32);
69            }
70            fn block(&self) -> [u32; 3] {
71                [self.block_inner as u32, 1, 1]
72            }
73            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
74                let n_outer: usize = output_shape.iter().product::<usize>().max(1);
75                [n_outer as u32, 1, 1]
76            }
77        }
78    };
79}
80
81macro_rules! impl_reduce_float_runtime_op {
82    ($Fwd:ident) => {
83        impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
84            fn n_activation_inputs(&self) -> usize {
85                1
86            }
87            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
88                vec![]
89            }
90            fn pack_args(
91                &self,
92                inputs: &[(teeny_core::model::RawPtr, &[usize])],
93                _: &[teeny_core::model::RawPtr],
94                output: teeny_core::model::RawPtr,
95                output_shape: &[usize],
96                _: i32,
97                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
98            ) {
99                let n_outer: usize = output_shape.iter().product::<usize>().max(1);
100                let n_total: usize = inputs[0].1.iter().product();
101                let n_inner: usize = if n_outer > 0 {
102                    n_total / n_outer
103                } else {
104                    n_total
105                };
106                visitor.visit_ptr(inputs[0].0);
107                visitor.visit_ptr(output);
108                visitor.visit_i32(n_inner as i32);
109                visitor.visit_i32(n_outer as i32);
110            }
111            fn block(&self) -> [u32; 3] {
112                [self.block_inner as u32, 1, 1]
113            }
114            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
115                let n_outer: usize = output_shape.iter().product::<usize>().max(1);
116                [n_outer as u32, 1, 1]
117            }
118        }
119    };
120}
121
122// ── ReduceSum ─────────────────────────────────────────────────────────────────
123
124/// Forward: y[row] = sum(x[row, :])
125// ANCHOR: reduce_sum_forward
126#[kernel]
127pub fn reduce_sum_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
128    x_ptr: T::Pointer<D>,
129    y_ptr: T::Pointer<D>,
130    n_inner: i32,
131    n_outer: 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 row = T::program_id(Axis::X);
138    if row >= n_outer {
139        return;
140    }
141    let col_offsets = T::arange(0, BLOCK_INNER);
142    let offsets = col_offsets + row * n_inner;
143    let mask = col_offsets.lt(n_inner);
144    let x = T::load(
145        x_ptr.add_offsets(offsets),
146        Some(mask),
147        Some(T::zeros::<D>(&[BLOCK_INNER])),
148        &[],
149        None,
150        None,
151        None,
152        false,
153    );
154    let sum = T::sum(x, Some(0), true); // [1] or scalar
155    let row_offsets = T::arange(0, 1) + row;
156    T::store(y_ptr.add_offsets(row_offsets), sum, None, &[], None, None);
157}
158
159// ANCHOR_END: reduce_sum_forward
160
161impl_reduce_num_runtime_op!(ReduceSumForward);
162
163// ── ReduceMean ────────────────────────────────────────────────────────────────
164
165/// Forward: y[row] = mean(x[row, :])
166#[kernel]
167pub fn reduce_mean_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
168    x_ptr: T::Pointer<D>,
169    y_ptr: T::Pointer<D>,
170    n_inner: i32,
171    n_outer: 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 row = T::program_id(Axis::X);
178    if row >= n_outer {
179        return;
180    }
181    let col_offsets = T::arange(0, BLOCK_INNER);
182    let offsets = col_offsets + row * n_inner;
183    let mask = col_offsets.lt(n_inner);
184    let x = T::load(
185        x_ptr.add_offsets(offsets),
186        Some(mask),
187        Some(T::zeros::<D>(&[BLOCK_INNER])),
188        &[],
189        None,
190        None,
191        None,
192        false,
193    );
194    let sum = T::sum(x, Some(0), true);
195    let n_f = T::cast::<i32, D>(T::full::<i32>(&[1], n_inner), None, false);
196    let mean = sum / n_f;
197    let row_offsets = T::arange(0, 1) + row;
198    T::store(y_ptr.add_offsets(row_offsets), mean, None, &[], None, None);
199}
200
201impl_reduce_float_runtime_op!(ReduceMeanForward);
202
203// ── ReduceMax ─────────────────────────────────────────────────────────────────
204
205/// Forward: y[row] = max(x[row, :])
206#[kernel]
207pub fn reduce_max_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
208    x_ptr: T::Pointer<D>,
209    y_ptr: T::Pointer<D>,
210    n_inner: i32,
211    n_outer: i32,
212) where
213    T::I32Tensor: types::Tensor<i32, 1>,
214    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
215    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
216{
217    let row = T::program_id(Axis::X);
218    if row >= n_outer {
219        return;
220    }
221    let col_offsets = T::arange(0, BLOCK_INNER);
222    let offsets = col_offsets + row * n_inner;
223    let mask = col_offsets.lt(n_inner);
224    // Load with a very small fill value for masked lanes
225    let neg_inf = T::cast::<f32, D>(
226        T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
227        None,
228        false,
229    );
230    let x = T::load(
231        x_ptr.add_offsets(offsets),
232        Some(mask),
233        Some(neg_inf),
234        &[],
235        None,
236        None,
237        None,
238        false,
239    );
240    let val = T::max(x, Some(0), true);
241    let row_offsets = T::arange(0, 1) + row;
242    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
243}
244
245impl_reduce_num_runtime_op!(ReduceMaxForward);
246
247// ── ReduceMin ─────────────────────────────────────────────────────────────────
248
249/// Forward: y[row] = min(x[row, :])
250#[kernel]
251pub fn reduce_min_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
252    x_ptr: T::Pointer<D>,
253    y_ptr: T::Pointer<D>,
254    n_inner: i32,
255    n_outer: i32,
256) where
257    T::I32Tensor: types::Tensor<i32, 1>,
258    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
259    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
260{
261    let row = T::program_id(Axis::X);
262    if row >= n_outer {
263        return;
264    }
265    let col_offsets = T::arange(0, BLOCK_INNER);
266    let offsets = col_offsets + row * n_inner;
267    let mask = col_offsets.lt(n_inner);
268    let pos_inf = T::cast::<f32, D>(
269        T::full::<f32>(&[BLOCK_INNER], 3.4028235e38_f32),
270        None,
271        false,
272    );
273    let x = T::load(
274        x_ptr.add_offsets(offsets),
275        Some(mask),
276        Some(pos_inf),
277        &[],
278        None,
279        None,
280        None,
281        false,
282    );
283    let val = T::min(x, Some(0), true);
284    let row_offsets = T::arange(0, 1) + row;
285    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
286}
287
288impl_reduce_num_runtime_op!(ReduceMinForward);
289
290// ── ReduceL1 ──────────────────────────────────────────────────────────────────
291
292/// Forward: y[row] = sum(|x[row, :]|)
293#[kernel]
294pub fn reduce_l1_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
295    x_ptr: T::Pointer<D>,
296    y_ptr: T::Pointer<D>,
297    n_inner: i32,
298    n_outer: i32,
299) where
300    T::I32Tensor: types::Tensor<i32, 1>,
301    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
302    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
303{
304    let row = T::program_id(Axis::X);
305    if row >= n_outer {
306        return;
307    }
308    let col_offsets = T::arange(0, BLOCK_INNER);
309    let offsets = col_offsets + row * n_inner;
310    let mask = col_offsets.lt(n_inner);
311    let x = T::load(
312        x_ptr.add_offsets(offsets),
313        Some(mask),
314        Some(T::zeros::<D>(&[BLOCK_INNER])),
315        &[],
316        None,
317        None,
318        None,
319        false,
320    );
321    let val = T::sum(T::abs(x), Some(0), true);
322    let row_offsets = T::arange(0, 1) + row;
323    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
324}
325
326impl_reduce_num_runtime_op!(ReduceL1Forward);
327
328// ── ReduceL2 ──────────────────────────────────────────────────────────────────
329
330/// Forward: y[row] = sqrt(sum(x[row, :]^2))
331#[kernel]
332pub fn reduce_l2_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
333    x_ptr: T::Pointer<D>,
334    y_ptr: T::Pointer<D>,
335    n_inner: i32,
336    n_outer: i32,
337) where
338    T::I32Tensor: types::Tensor<i32, 1>,
339    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
340    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
341{
342    let row = T::program_id(Axis::X);
343    if row >= n_outer {
344        return;
345    }
346    let col_offsets = T::arange(0, BLOCK_INNER);
347    let offsets = col_offsets + row * n_inner;
348    let mask = col_offsets.lt(n_inner);
349    let x = T::load(
350        x_ptr.add_offsets(offsets),
351        Some(mask),
352        Some(T::zeros::<D>(&[BLOCK_INNER])),
353        &[],
354        None,
355        None,
356        None,
357        false,
358    );
359    let sum_sq = T::sum(x * x, Some(0), true);
360    let val = T::sqrt(sum_sq);
361    let row_offsets = T::arange(0, 1) + row;
362    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
363}
364
365impl_reduce_float_runtime_op!(ReduceL2Forward);
366
367// ── ReduceSumSquare ───────────────────────────────────────────────────────────
368
369/// Forward: y[row] = sum(x[row, :]^2)
370#[kernel]
371pub fn reduce_sum_square_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
372    x_ptr: T::Pointer<D>,
373    y_ptr: T::Pointer<D>,
374    n_inner: i32,
375    n_outer: i32,
376) where
377    T::I32Tensor: types::Tensor<i32, 1>,
378    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
379    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
380{
381    let row = T::program_id(Axis::X);
382    if row >= n_outer {
383        return;
384    }
385    let col_offsets = T::arange(0, BLOCK_INNER);
386    let offsets = col_offsets + row * n_inner;
387    let mask = col_offsets.lt(n_inner);
388    let x = T::load(
389        x_ptr.add_offsets(offsets),
390        Some(mask),
391        Some(T::zeros::<D>(&[BLOCK_INNER])),
392        &[],
393        None,
394        None,
395        None,
396        false,
397    );
398    let val = T::sum(x * x, Some(0), true);
399    let row_offsets = T::arange(0, 1) + row;
400    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
401}
402
403impl_reduce_num_runtime_op!(ReduceSumSquareForward);
404
405// ── ReduceLogSum ──────────────────────────────────────────────────────────────
406
407/// Forward: y[row] = log(sum(x[row, :]))  (numerically unsafe; use ReduceLogSumExp for stable)
408#[kernel]
409pub fn reduce_log_sum_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
410    x_ptr: T::Pointer<D>,
411    y_ptr: T::Pointer<D>,
412    n_inner: i32,
413    n_outer: i32,
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 row = T::program_id(Axis::X);
420    if row >= n_outer {
421        return;
422    }
423    let col_offsets = T::arange(0, BLOCK_INNER);
424    let offsets = col_offsets + row * n_inner;
425    let mask = col_offsets.lt(n_inner);
426    let x = T::load(
427        x_ptr.add_offsets(offsets),
428        Some(mask),
429        Some(T::zeros::<D>(&[BLOCK_INNER])),
430        &[],
431        None,
432        None,
433        None,
434        false,
435    );
436    let sum = T::sum(x, Some(0), true);
437    let val = T::log(sum);
438    let row_offsets = T::arange(0, 1) + row;
439    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
440}
441
442impl_reduce_float_runtime_op!(ReduceLogSumForward);
443
444// ── ReduceLogSumExp ───────────────────────────────────────────────────────────
445
446/// Forward: y[row] = log(sum(exp(x[row, :]))) — numerically stable via max subtraction
447#[kernel]
448pub fn reduce_log_sum_exp_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
449    x_ptr: T::Pointer<D>,
450    y_ptr: T::Pointer<D>,
451    n_inner: i32,
452    n_outer: i32,
453) where
454    T::I32Tensor: types::Tensor<i32, 1>,
455    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
456    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
457{
458    let row = T::program_id(Axis::X);
459    if row >= n_outer {
460        return;
461    }
462    let col_offsets = T::arange(0, BLOCK_INNER);
463    let offsets = col_offsets + row * n_inner;
464    let mask = col_offsets.lt(n_inner);
465    let neg_inf = T::cast::<f32, D>(
466        T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
467        None,
468        false,
469    );
470    let x = T::load(
471        x_ptr.add_offsets(offsets),
472        Some(mask),
473        Some(neg_inf),
474        &[],
475        None,
476        None,
477        None,
478        false,
479    );
480    // Numerically stable: log(sum(exp(x))) = m + log(sum(exp(x - m)))
481    // where m = max(x)
482    let m = T::max(x, Some(0), true); // [1]
483    let fill = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_INNER], 0.0_f32), None, false);
484    let x_adj = T::load(
485        x_ptr.add_offsets(offsets),
486        Some(mask),
487        Some(fill),
488        &[],
489        None,
490        None,
491        None,
492        false,
493    );
494    let sum_exp = T::sum(T::exp(x_adj - m), Some(0), true);
495    let val = m + T::log(sum_exp);
496    let row_offsets = T::arange(0, 1) + row;
497    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
498}
499
500impl_reduce_float_runtime_op!(ReduceLogSumExpForward);
501
502// ── ReduceProd ────────────────────────────────────────────────────────────────
503
504/// Forward: y[row] = prod(x[row, :])
505/// Note: implemented as exp(sum(log(x))) — only valid for positive x.
506/// For general use this is a placeholder.
507#[kernel]
508pub fn reduce_prod_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
509    x_ptr: T::Pointer<D>,
510    y_ptr: T::Pointer<D>,
511    n_inner: i32,
512    n_outer: i32,
513) where
514    T::I32Tensor: types::Tensor<i32, 1>,
515    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
516    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
517{
518    let row = T::program_id(Axis::X);
519    if row >= n_outer {
520        return;
521    }
522    let col_offsets = T::arange(0, BLOCK_INNER);
523    let offsets = col_offsets + row * n_inner;
524    let mask = col_offsets.lt(n_inner);
525    // Fill with 1.0 for masked-off lanes so they don't affect the product.
526    let one_fill = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_INNER], 1.0_f32), None, false);
527    let x = T::load(
528        x_ptr.add_offsets(offsets),
529        Some(mask),
530        Some(one_fill),
531        &[],
532        None,
533        None,
534        None,
535        false,
536    );
537    // exp(sum(log(x))) approximates product for positive x.
538    let val = T::exp(T::sum(T::log(x), Some(0), true));
539    let row_offsets = T::arange(0, 1) + row;
540    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
541}
542
543impl_reduce_float_runtime_op!(ReduceProdForward);
544
545// ── CumSum ────────────────────────────────────────────────────────────────────
546
547/// Forward: y = cumsum(x, axis=0) over a 1-D block
548/// Each CTA handles one complete row (n_inner elements).
549#[kernel]
550pub fn cum_sum_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
551    x_ptr: T::Pointer<D>,
552    y_ptr: T::Pointer<D>,
553    n_inner: i32,
554    n_outer: i32,
555) where
556    T::I32Tensor: types::Tensor<i32, 1>,
557    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
558    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
559{
560    let row = T::program_id(Axis::X);
561    if row >= n_outer {
562        return;
563    }
564    let col_offsets = T::arange(0, BLOCK_INNER);
565    let offsets = col_offsets + row * n_inner;
566    let mask = col_offsets.lt(n_inner);
567    let x = T::load(
568        x_ptr.add_offsets(offsets),
569        Some(mask),
570        Some(T::zeros::<D>(&[BLOCK_INNER])),
571        &[],
572        None,
573        None,
574        None,
575        false,
576    );
577    // Use Triton's cumsum: axis=0 over the 1-D block, not reversed.
578    let y = T::cumsum(x, 0, false);
579    T::store(y_ptr.add_offsets(offsets), y, Some(mask), &[], None, None);
580}
581
582impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for CumSumForward<D> {
583    fn n_activation_inputs(&self) -> usize {
584        1
585    }
586    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
587        vec![]
588    }
589    fn pack_args(
590        &self,
591        inputs: &[(teeny_core::model::RawPtr, &[usize])],
592        _: &[teeny_core::model::RawPtr],
593        output: teeny_core::model::RawPtr,
594        output_shape: &[usize],
595        _: i32,
596        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
597    ) {
598        // For cumsum: output_shape == input_shape; n_outer = all dims except last
599        let n_total: usize = output_shape.iter().product();
600        let n_inner = output_shape.last().copied().unwrap_or(1);
601        let n_outer = n_total / n_inner;
602        visitor.visit_ptr(inputs[0].0);
603        visitor.visit_ptr(output);
604        visitor.visit_i32(n_inner as i32);
605        visitor.visit_i32(n_outer as i32);
606    }
607    fn block(&self) -> [u32; 3] {
608        [self.block_inner as u32, 1, 1]
609    }
610    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
611        let n_total: usize = output_shape.iter().product();
612        let n_inner = output_shape.last().copied().unwrap_or(1);
613        let n_outer = n_total / n_inner;
614        [n_outer as u32, 1, 1]
615    }
616}
617
618// ── CumProd ───────────────────────────────────────────────────────────────────
619
620/// Forward: y = cumprod(x, axis=0) over a 1-D block
621#[kernel]
622pub fn cum_prod_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
623    x_ptr: T::Pointer<D>,
624    y_ptr: T::Pointer<D>,
625    n_inner: i32,
626    n_outer: i32,
627) where
628    T::I32Tensor: types::Tensor<i32, 1>,
629    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
630    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
631{
632    let row = T::program_id(Axis::X);
633    if row >= n_outer {
634        return;
635    }
636    let col_offsets = T::arange(0, BLOCK_INNER);
637    let offsets = col_offsets + row * n_inner;
638    let mask = col_offsets.lt(n_inner);
639    let x = T::load(
640        x_ptr.add_offsets(offsets),
641        Some(mask),
642        Some(T::zeros::<D>(&[BLOCK_INNER])),
643        &[],
644        None,
645        None,
646        None,
647        false,
648    );
649    let y = T::cumprod(x, 0, false);
650    T::store(y_ptr.add_offsets(offsets), y, Some(mask), &[], None, None);
651}
652
653impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for CumProdForward<D> {
654    fn n_activation_inputs(&self) -> usize {
655        1
656    }
657    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
658        vec![]
659    }
660    fn pack_args(
661        &self,
662        inputs: &[(teeny_core::model::RawPtr, &[usize])],
663        _: &[teeny_core::model::RawPtr],
664        output: teeny_core::model::RawPtr,
665        output_shape: &[usize],
666        _: i32,
667        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
668    ) {
669        let n_total: usize = output_shape.iter().product();
670        let n_inner = output_shape.last().copied().unwrap_or(1);
671        let n_outer = n_total / n_inner;
672        visitor.visit_ptr(inputs[0].0);
673        visitor.visit_ptr(output);
674        visitor.visit_i32(n_inner as i32);
675        visitor.visit_i32(n_outer as i32);
676    }
677    fn block(&self) -> [u32; 3] {
678        [self.block_inner as u32, 1, 1]
679    }
680    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
681        let n_total: usize = output_shape.iter().product();
682        let n_inner = output_shape.last().copied().unwrap_or(1);
683        let n_outer = n_total / n_inner;
684        [n_outer as u32, 1, 1]
685    }
686}
687
688// ArgMax and ArgMin kernels are deferred — the Triton type system requires
689// I32Tensor → Tensor<i32> coercion that isn't directly supported via #[kernel].
690// These are handled as TODO in the lowering match arm.
691
692// ── GlobalAvgPool ─────────────────────────────────────────────────────────────
693//
694// Treats input as [n_outer, n_inner] and averages over n_inner.
695// For a [N, C, H, W] input: n_outer = N * C, n_inner = H * W.
696
697/// Forward: y[row] = mean(x[row, :])  (same as ReduceMean)
698#[kernel]
699pub fn global_avg_pool_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
700    x_ptr: T::Pointer<D>,
701    y_ptr: T::Pointer<D>,
702    n_inner: i32,
703    n_outer: i32,
704) where
705    T::I32Tensor: types::Tensor<i32, 1>,
706    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
707    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
708{
709    let row = T::program_id(Axis::X);
710    if row >= n_outer {
711        return;
712    }
713    let col_offsets = T::arange(0, BLOCK_INNER);
714    let offsets = col_offsets + row * n_inner;
715    let mask = col_offsets.lt(n_inner);
716    let x = T::load(
717        x_ptr.add_offsets(offsets),
718        Some(mask),
719        Some(T::zeros::<D>(&[BLOCK_INNER])),
720        &[],
721        None,
722        None,
723        None,
724        false,
725    );
726    let sum = T::sum(x, Some(0), true);
727    let n_f = T::cast::<i32, D>(T::full::<i32>(&[1], n_inner), None, false);
728    let mean = sum / n_f;
729    let row_offsets = T::arange(0, 1) + row;
730    T::store(y_ptr.add_offsets(row_offsets), mean, None, &[], None, None);
731}
732
733impl_reduce_float_runtime_op!(GlobalAvgPoolForward);
734
735// ── GlobalMaxPool ─────────────────────────────────────────────────────────────
736
737/// Forward: y[row] = max(x[row, :])  (same as ReduceMax)
738#[kernel]
739pub fn global_max_pool_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
740    x_ptr: T::Pointer<D>,
741    y_ptr: T::Pointer<D>,
742    n_inner: i32,
743    n_outer: i32,
744) where
745    T::I32Tensor: types::Tensor<i32, 1>,
746    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
747    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
748{
749    let row = T::program_id(Axis::X);
750    if row >= n_outer {
751        return;
752    }
753    let col_offsets = T::arange(0, BLOCK_INNER);
754    let offsets = col_offsets + row * n_inner;
755    let mask = col_offsets.lt(n_inner);
756    let neg_inf = T::cast::<f32, D>(
757        T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
758        None,
759        false,
760    );
761    let x = T::load(
762        x_ptr.add_offsets(offsets),
763        Some(mask),
764        Some(neg_inf),
765        &[],
766        None,
767        None,
768        None,
769        false,
770    );
771    let val = T::max(x, Some(0), true);
772    let row_offsets = T::arange(0, 1) + row;
773    T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
774}
775
776impl_reduce_float_runtime_op!(GlobalMaxPoolForward);