Skip to main content

teeny_kernels/nn/tensor/
elemwise_binary.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, Num};
20use teeny_macros::kernel;
21use teeny_triton::triton::{
22    types::{AddOffsets, Comparison},
23    *,
24};
25
26// ── Helper macros for 2-input RuntimeOp ──────────────────────────────────────
27
28macro_rules! impl_binary_num_runtime_op_with_bwd {
29    ($Fwd:ident) => {
30        impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
31            fn n_activation_inputs(&self) -> usize {
32                2
33            }
34            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
35                vec![]
36            }
37            fn pack_args(
38                &self,
39                inputs: &[(teeny_core::model::RawPtr, &[usize])],
40                _: &[teeny_core::model::RawPtr],
41                output: teeny_core::model::RawPtr,
42                output_shape: &[usize],
43                _: i32,
44                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
45            ) {
46                let n: usize = output_shape.iter().product();
47                visitor.visit_ptr(inputs[0].0); // a_ptr
48                visitor.visit_ptr(inputs[1].0); // b_ptr
49                visitor.visit_ptr(output);
50                visitor.visit_i32(n as i32);
51            }
52            fn block(&self) -> [u32; 3] {
53                [self.block_size as u32, 1, 1]
54            }
55            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
56                let n: usize = output_shape.iter().product();
57                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
58            }
59            #[cfg(feature = "training")]
60            fn has_backward(&self) -> bool {
61                true
62            }
63            #[cfg(feature = "training")]
64            fn pack_backward_args(
65                &self,
66                inputs: &[(teeny_core::model::RawPtr, &[usize])],
67                _: &[teeny_core::model::RawPtr],
68                _: teeny_core::model::RawPtr,
69                output_shape: &[usize],
70                grad_output: teeny_core::model::RawPtr,
71                _: i32,
72                grad_inputs: &[teeny_core::model::RawPtr],
73                _: &[teeny_core::model::RawPtr],
74                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
75            ) {
76                let n: usize = output_shape.iter().product();
77                visitor.visit_ptr(grad_output);
78                visitor.visit_ptr(inputs[0].0);
79                visitor.visit_ptr(inputs[1].0);
80                visitor.visit_ptr(grad_inputs[0]);
81                visitor.visit_ptr(grad_inputs[1]);
82                visitor.visit_i32(n as i32);
83            }
84            #[cfg(feature = "training")]
85            fn backward_block(&self) -> [u32; 3] {
86                [self.block_size as u32, 1, 1]
87            }
88            #[cfg(feature = "training")]
89            fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
90                let n: usize = output_shape.iter().product();
91                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
92            }
93        }
94    };
95}
96
97macro_rules! impl_binary_float_runtime_op_with_bwd {
98    ($Fwd:ident) => {
99        impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
100            fn n_activation_inputs(&self) -> usize {
101                2
102            }
103            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
104                vec![]
105            }
106            fn pack_args(
107                &self,
108                inputs: &[(teeny_core::model::RawPtr, &[usize])],
109                _: &[teeny_core::model::RawPtr],
110                output: teeny_core::model::RawPtr,
111                output_shape: &[usize],
112                _: i32,
113                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
114            ) {
115                let n: usize = output_shape.iter().product();
116                visitor.visit_ptr(inputs[0].0);
117                visitor.visit_ptr(inputs[1].0);
118                visitor.visit_ptr(output);
119                visitor.visit_i32(n as i32);
120            }
121            fn block(&self) -> [u32; 3] {
122                [self.block_size as u32, 1, 1]
123            }
124            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
125                let n: usize = output_shape.iter().product();
126                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
127            }
128            #[cfg(feature = "training")]
129            fn has_backward(&self) -> bool {
130                true
131            }
132            #[cfg(feature = "training")]
133            fn pack_backward_args(
134                &self,
135                inputs: &[(teeny_core::model::RawPtr, &[usize])],
136                _: &[teeny_core::model::RawPtr],
137                _: teeny_core::model::RawPtr,
138                output_shape: &[usize],
139                grad_output: teeny_core::model::RawPtr,
140                _: i32,
141                grad_inputs: &[teeny_core::model::RawPtr],
142                _: &[teeny_core::model::RawPtr],
143                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
144            ) {
145                let n: usize = output_shape.iter().product();
146                visitor.visit_ptr(grad_output);
147                visitor.visit_ptr(inputs[0].0);
148                visitor.visit_ptr(inputs[1].0);
149                visitor.visit_ptr(grad_inputs[0]);
150                visitor.visit_ptr(grad_inputs[1]);
151                visitor.visit_i32(n as i32);
152            }
153            #[cfg(feature = "training")]
154            fn backward_block(&self) -> [u32; 3] {
155                [self.block_size as u32, 1, 1]
156            }
157            #[cfg(feature = "training")]
158            fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
159                let n: usize = output_shape.iter().product();
160                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
161            }
162        }
163    };
164}
165
166/// Binary num op with no backward (comparison / logical ops)
167macro_rules! impl_binary_num_runtime_op_no_bwd {
168    ($Fwd:ident) => {
169        impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
170            fn n_activation_inputs(&self) -> usize {
171                2
172            }
173            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
174                vec![]
175            }
176            fn pack_args(
177                &self,
178                inputs: &[(teeny_core::model::RawPtr, &[usize])],
179                _: &[teeny_core::model::RawPtr],
180                output: teeny_core::model::RawPtr,
181                output_shape: &[usize],
182                _: i32,
183                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
184            ) {
185                let n: usize = output_shape.iter().product();
186                visitor.visit_ptr(inputs[0].0);
187                visitor.visit_ptr(inputs[1].0);
188                visitor.visit_ptr(output);
189                visitor.visit_i32(n as i32);
190            }
191            fn block(&self) -> [u32; 3] {
192                [self.block_size as u32, 1, 1]
193            }
194            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
195                let n: usize = output_shape.iter().product();
196                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
197            }
198        }
199    };
200}
201
202// ── Mul ───────────────────────────────────────────────────────────────────────
203
204/// Forward: out = a * b
205#[kernel]
206pub fn elemwise_mul_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
207    a_ptr: T::Pointer<D>,
208    b_ptr: T::Pointer<D>,
209    out_ptr: T::Pointer<D>,
210    n_elements: i32,
211) where
212    T::I32Tensor: types::Tensor<i32, 1>,
213    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
214    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
215{
216    let pid = T::program_id(Axis::X);
217    let block_start = pid * BLOCK_SIZE;
218    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
219    let in_bounds = offsets.lt(n_elements);
220    let a = T::load(
221        a_ptr.add_offsets(offsets),
222        Some(in_bounds),
223        None,
224        &[],
225        None,
226        None,
227        None,
228        false,
229    );
230    let b = T::load(
231        b_ptr.add_offsets(offsets),
232        Some(in_bounds),
233        None,
234        &[],
235        None,
236        None,
237        None,
238        false,
239    );
240    T::store(
241        out_ptr.add_offsets(offsets),
242        a * b,
243        Some(in_bounds),
244        &[],
245        None,
246        None,
247    );
248}
249
250/// Backward: da = dy * b,  db = dy * a
251#[kernel]
252pub fn elemwise_mul_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
253    dy_ptr: T::Pointer<D>,
254    a_ptr: T::Pointer<D>,
255    b_ptr: T::Pointer<D>,
256    da_ptr: T::Pointer<D>,
257    db_ptr: T::Pointer<D>,
258    n_elements: i32,
259) where
260    T::I32Tensor: types::Tensor<i32, 1>,
261    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
262    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
263{
264    let pid = T::program_id(Axis::X);
265    let block_start = pid * BLOCK_SIZE;
266    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
267    let in_bounds = offsets.lt(n_elements);
268    let dy = T::load(
269        dy_ptr.add_offsets(offsets),
270        Some(in_bounds),
271        None,
272        &[],
273        None,
274        None,
275        None,
276        false,
277    );
278    let a = T::load(
279        a_ptr.add_offsets(offsets),
280        Some(in_bounds),
281        None,
282        &[],
283        None,
284        None,
285        None,
286        false,
287    );
288    let b = T::load(
289        b_ptr.add_offsets(offsets),
290        Some(in_bounds),
291        None,
292        &[],
293        None,
294        None,
295        None,
296        false,
297    );
298    T::store(
299        da_ptr.add_offsets(offsets),
300        dy * b,
301        Some(in_bounds),
302        &[],
303        None,
304        None,
305    );
306    T::store(
307        db_ptr.add_offsets(offsets),
308        dy * a,
309        Some(in_bounds),
310        &[],
311        None,
312        None,
313    );
314}
315
316impl_binary_num_runtime_op_with_bwd!(ElemwiseMulForward);
317
318// ── Sub ───────────────────────────────────────────────────────────────────────
319
320/// Forward: out = a - b
321#[kernel]
322pub fn elemwise_sub_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
323    a_ptr: T::Pointer<D>,
324    b_ptr: T::Pointer<D>,
325    out_ptr: T::Pointer<D>,
326    n_elements: i32,
327) where
328    T::I32Tensor: types::Tensor<i32, 1>,
329    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
330    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
331{
332    let pid = T::program_id(Axis::X);
333    let block_start = pid * BLOCK_SIZE;
334    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
335    let in_bounds = offsets.lt(n_elements);
336    let a = T::load(
337        a_ptr.add_offsets(offsets),
338        Some(in_bounds),
339        None,
340        &[],
341        None,
342        None,
343        None,
344        false,
345    );
346    let b = T::load(
347        b_ptr.add_offsets(offsets),
348        Some(in_bounds),
349        None,
350        &[],
351        None,
352        None,
353        None,
354        false,
355    );
356    T::store(
357        out_ptr.add_offsets(offsets),
358        a - b,
359        Some(in_bounds),
360        &[],
361        None,
362        None,
363    );
364}
365
366/// Backward: da = dy,  db = -dy
367#[kernel]
368pub fn elemwise_sub_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
369    dy_ptr: T::Pointer<D>,
370    da_ptr: T::Pointer<D>,
371    db_ptr: T::Pointer<D>,
372    n_elements: i32,
373) where
374    T::I32Tensor: types::Tensor<i32, 1>,
375    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
376    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
377{
378    let pid = T::program_id(Axis::X);
379    let block_start = pid * BLOCK_SIZE;
380    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
381    let in_bounds = offsets.lt(n_elements);
382    let dy = T::load(
383        dy_ptr.add_offsets(offsets),
384        Some(in_bounds),
385        None,
386        &[],
387        None,
388        None,
389        None,
390        false,
391    );
392    T::store(
393        da_ptr.add_offsets(offsets),
394        dy,
395        Some(in_bounds),
396        &[],
397        None,
398        None,
399    );
400    T::store(
401        db_ptr.add_offsets(offsets),
402        -dy,
403        Some(in_bounds),
404        &[],
405        None,
406        None,
407    );
408}
409
410impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for ElemwiseSubForward<D> {
411    fn n_activation_inputs(&self) -> usize {
412        2
413    }
414    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
415        vec![]
416    }
417    fn pack_args(
418        &self,
419        inputs: &[(teeny_core::model::RawPtr, &[usize])],
420        _: &[teeny_core::model::RawPtr],
421        output: teeny_core::model::RawPtr,
422        output_shape: &[usize],
423        _: i32,
424        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
425    ) {
426        let n: usize = output_shape.iter().product();
427        visitor.visit_ptr(inputs[0].0);
428        visitor.visit_ptr(inputs[1].0);
429        visitor.visit_ptr(output);
430        visitor.visit_i32(n as i32);
431    }
432    fn block(&self) -> [u32; 3] {
433        [self.block_size as u32, 1, 1]
434    }
435    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
436        let n: usize = output_shape.iter().product();
437        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
438    }
439    #[cfg(feature = "training")]
440    fn has_backward(&self) -> bool {
441        true
442    }
443    #[cfg(feature = "training")]
444    fn pack_backward_args(
445        &self,
446        _: &[(teeny_core::model::RawPtr, &[usize])],
447        _: &[teeny_core::model::RawPtr],
448        _: teeny_core::model::RawPtr,
449        output_shape: &[usize],
450        grad_output: teeny_core::model::RawPtr,
451        _: i32,
452        grad_inputs: &[teeny_core::model::RawPtr],
453        _: &[teeny_core::model::RawPtr],
454        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
455    ) {
456        let n: usize = output_shape.iter().product();
457        visitor.visit_ptr(grad_output);
458        visitor.visit_ptr(grad_inputs[0]);
459        visitor.visit_ptr(grad_inputs[1]);
460        visitor.visit_i32(n as i32);
461    }
462    #[cfg(feature = "training")]
463    fn backward_block(&self) -> [u32; 3] {
464        [self.block_size as u32, 1, 1]
465    }
466    #[cfg(feature = "training")]
467    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
468        let n: usize = output_shape.iter().product();
469        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
470    }
471}
472
473// ── Div ───────────────────────────────────────────────────────────────────────
474
475/// Forward: out = a / b
476#[kernel]
477pub fn elemwise_div_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
478    a_ptr: T::Pointer<D>,
479    b_ptr: T::Pointer<D>,
480    out_ptr: T::Pointer<D>,
481    n_elements: i32,
482) where
483    T::I32Tensor: types::Tensor<i32, 1>,
484    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
485    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
486{
487    let pid = T::program_id(Axis::X);
488    let block_start = pid * BLOCK_SIZE;
489    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
490    let in_bounds = offsets.lt(n_elements);
491    let a = T::load(
492        a_ptr.add_offsets(offsets),
493        Some(in_bounds),
494        None,
495        &[],
496        None,
497        None,
498        None,
499        false,
500    );
501    let b = T::load(
502        b_ptr.add_offsets(offsets),
503        Some(in_bounds),
504        None,
505        &[],
506        None,
507        None,
508        None,
509        false,
510    );
511    T::store(
512        out_ptr.add_offsets(offsets),
513        a / b,
514        Some(in_bounds),
515        &[],
516        None,
517        None,
518    );
519}
520
521/// Backward: da = dy / b,  db = -a * dy / b^2
522#[kernel]
523pub fn elemwise_div_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
524    dy_ptr: T::Pointer<D>,
525    a_ptr: T::Pointer<D>,
526    b_ptr: T::Pointer<D>,
527    da_ptr: T::Pointer<D>,
528    db_ptr: T::Pointer<D>,
529    n_elements: i32,
530) where
531    T::I32Tensor: types::Tensor<i32, 1>,
532    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
533    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
534{
535    let pid = T::program_id(Axis::X);
536    let block_start = pid * BLOCK_SIZE;
537    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
538    let in_bounds = offsets.lt(n_elements);
539    let dy = T::load(
540        dy_ptr.add_offsets(offsets),
541        Some(in_bounds),
542        None,
543        &[],
544        None,
545        None,
546        None,
547        false,
548    );
549    let a = T::load(
550        a_ptr.add_offsets(offsets),
551        Some(in_bounds),
552        None,
553        &[],
554        None,
555        None,
556        None,
557        false,
558    );
559    let b = T::load(
560        b_ptr.add_offsets(offsets),
561        Some(in_bounds),
562        None,
563        &[],
564        None,
565        None,
566        None,
567        false,
568    );
569    T::store(
570        da_ptr.add_offsets(offsets),
571        dy / b,
572        Some(in_bounds),
573        &[],
574        None,
575        None,
576    );
577    T::store(
578        db_ptr.add_offsets(offsets),
579        -(a * dy / (b * b)),
580        Some(in_bounds),
581        &[],
582        None,
583        None,
584    );
585}
586
587impl_binary_float_runtime_op_with_bwd!(ElemwiseDivForward);
588
589// ── Pow (D: Float) ────────────────────────────────────────────────────────────
590
591/// Forward: out = a ^ b = exp(b * log(a))
592#[kernel]
593pub fn elemwise_pow_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
594    a_ptr: T::Pointer<D>,
595    b_ptr: T::Pointer<D>,
596    out_ptr: T::Pointer<D>,
597    n_elements: i32,
598) where
599    T::I32Tensor: types::Tensor<i32, 1>,
600    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
601    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
602{
603    let pid = T::program_id(Axis::X);
604    let block_start = pid * BLOCK_SIZE;
605    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
606    let in_bounds = offsets.lt(n_elements);
607    let a = T::load(
608        a_ptr.add_offsets(offsets),
609        Some(in_bounds),
610        None,
611        &[],
612        None,
613        None,
614        None,
615        false,
616    );
617    let b = T::load(
618        b_ptr.add_offsets(offsets),
619        Some(in_bounds),
620        None,
621        &[],
622        None,
623        None,
624        None,
625        false,
626    );
627    let y = T::exp(b * T::log(a));
628    T::store(
629        out_ptr.add_offsets(offsets),
630        y,
631        Some(in_bounds),
632        &[],
633        None,
634        None,
635    );
636}
637
638/// Backward: da = b * a^(b-1) * dy,  db = log(a) * a^b * dy
639#[kernel]
640pub fn elemwise_pow_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
641    dy_ptr: T::Pointer<D>,
642    a_ptr: T::Pointer<D>,
643    b_ptr: T::Pointer<D>,
644    da_ptr: T::Pointer<D>,
645    db_ptr: T::Pointer<D>,
646    n_elements: i32,
647) where
648    T::I32Tensor: types::Tensor<i32, 1>,
649    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
650    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
651{
652    let pid = T::program_id(Axis::X);
653    let block_start = pid * BLOCK_SIZE;
654    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
655    let in_bounds = offsets.lt(n_elements);
656    let dy = T::load(
657        dy_ptr.add_offsets(offsets),
658        Some(in_bounds),
659        None,
660        &[],
661        None,
662        None,
663        None,
664        false,
665    );
666    let a = T::load(
667        a_ptr.add_offsets(offsets),
668        Some(in_bounds),
669        None,
670        &[],
671        None,
672        None,
673        None,
674        false,
675    );
676    let b = T::load(
677        b_ptr.add_offsets(offsets),
678        Some(in_bounds),
679        None,
680        &[],
681        None,
682        None,
683        None,
684        false,
685    );
686    let a_pow_b = T::exp(b * T::log(a)); // a^b
687    // a^(b-1) = a^b / a  (avoids generic constant D::ONE)
688    let a_pow_bm1 = a_pow_b / a;
689    T::store(
690        da_ptr.add_offsets(offsets),
691        b * a_pow_bm1 * dy,
692        Some(in_bounds),
693        &[],
694        None,
695        None,
696    );
697    T::store(
698        db_ptr.add_offsets(offsets),
699        T::log(a) * a_pow_b * dy,
700        Some(in_bounds),
701        &[],
702        None,
703        None,
704    );
705}
706
707impl_binary_float_runtime_op_with_bwd!(ElemwisePowForward);
708
709// ── Mod ───────────────────────────────────────────────────────────────────────
710
711/// Forward fmod: out = a - trunc(a/b)*b  (C-style float remainder)
712#[kernel]
713pub fn elemwise_fmod_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
714    a_ptr: T::Pointer<D>,
715    b_ptr: T::Pointer<D>,
716    out_ptr: T::Pointer<D>,
717    n_elements: i32,
718) where
719    T::I32Tensor: types::Tensor<i32, 1>,
720    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
721    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
722{
723    let pid = T::program_id(Axis::X);
724    let block_start = pid * BLOCK_SIZE;
725    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
726    let in_bounds = offsets.lt(n_elements);
727    let a = T::load(
728        a_ptr.add_offsets(offsets),
729        Some(in_bounds),
730        None,
731        &[],
732        None,
733        None,
734        None,
735        false,
736    );
737    let b = T::load(
738        b_ptr.add_offsets(offsets),
739        Some(in_bounds),
740        None,
741        &[],
742        None,
743        None,
744        None,
745        false,
746    );
747    // fmod: a - floor(a/b)*b  (use floor here; for true C fmod we'd need trunc)
748    // Using floor makes this the Python-style modulo which is more broadly useful.
749    let y = a - T::floor(a / b) * b;
750    T::store(
751        out_ptr.add_offsets(offsets),
752        y,
753        Some(in_bounds),
754        &[],
755        None,
756        None,
757    );
758}
759
760impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for ElemwiseFmodForward<D> {
761    fn n_activation_inputs(&self) -> usize {
762        2
763    }
764    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
765        vec![]
766    }
767    fn pack_args(
768        &self,
769        inputs: &[(teeny_core::model::RawPtr, &[usize])],
770        _: &[teeny_core::model::RawPtr],
771        output: teeny_core::model::RawPtr,
772        output_shape: &[usize],
773        _: i32,
774        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
775    ) {
776        let n: usize = output_shape.iter().product();
777        visitor.visit_ptr(inputs[0].0);
778        visitor.visit_ptr(inputs[1].0);
779        visitor.visit_ptr(output);
780        visitor.visit_i32(n as i32);
781    }
782    fn block(&self) -> [u32; 3] {
783        [self.block_size as u32, 1, 1]
784    }
785    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
786        let n: usize = output_shape.iter().product();
787        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
788    }
789}
790
791// ── ElemMin / ElemMax (D: Num) ────────────────────────────────────────────────
792
793/// Forward: out = min(a, b)
794#[kernel]
795pub fn elemwise_min_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
796    a_ptr: T::Pointer<D>,
797    b_ptr: T::Pointer<D>,
798    out_ptr: T::Pointer<D>,
799    n_elements: i32,
800) where
801    T::I32Tensor: types::Tensor<i32, 1>,
802    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
803    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
804{
805    let pid = T::program_id(Axis::X);
806    let block_start = pid * BLOCK_SIZE;
807    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
808    let in_bounds = offsets.lt(n_elements);
809    let a = T::load(
810        a_ptr.add_offsets(offsets),
811        Some(in_bounds),
812        None,
813        &[],
814        None,
815        None,
816        None,
817        false,
818    );
819    let b = T::load(
820        b_ptr.add_offsets(offsets),
821        Some(in_bounds),
822        None,
823        &[],
824        None,
825        None,
826        None,
827        false,
828    );
829    T::store(
830        out_ptr.add_offsets(offsets),
831        T::minimum(a, b),
832        Some(in_bounds),
833        &[],
834        None,
835        None,
836    );
837}
838
839/// Backward: pass dy to the input that was smaller, 0 to the other.
840#[kernel]
841pub fn elemwise_min_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
842    dy_ptr: T::Pointer<D>,
843    a_ptr: T::Pointer<D>,
844    b_ptr: T::Pointer<D>,
845    da_ptr: T::Pointer<D>,
846    db_ptr: T::Pointer<D>,
847    n_elements: i32,
848) where
849    T::I32Tensor: types::Tensor<i32, 1>,
850    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
851    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
852{
853    let pid = T::program_id(Axis::X);
854    let block_start = pid * BLOCK_SIZE;
855    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
856    let in_bounds = offsets.lt(n_elements);
857    let dy = T::load(
858        dy_ptr.add_offsets(offsets),
859        Some(in_bounds),
860        None,
861        &[],
862        None,
863        None,
864        None,
865        false,
866    );
867    let a = T::load(
868        a_ptr.add_offsets(offsets),
869        Some(in_bounds),
870        None,
871        &[],
872        None,
873        None,
874        None,
875        false,
876    );
877    let b = T::load(
878        b_ptr.add_offsets(offsets),
879        Some(in_bounds),
880        None,
881        &[],
882        None,
883        None,
884        None,
885        false,
886    );
887    let z = T::zeros_like(dy);
888    let a_is_min = T::le(a, b);
889    T::store(
890        da_ptr.add_offsets(offsets),
891        T::where_(a_is_min, dy, z),
892        Some(in_bounds),
893        &[],
894        None,
895        None,
896    );
897    T::store(
898        db_ptr.add_offsets(offsets),
899        T::where_(a_is_min, z, dy),
900        Some(in_bounds),
901        &[],
902        None,
903        None,
904    );
905}
906
907impl_binary_num_runtime_op_with_bwd!(ElemwiseMinForward);
908
909/// Forward: out = max(a, b)
910#[kernel]
911pub fn elemwise_max_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
912    a_ptr: T::Pointer<D>,
913    b_ptr: T::Pointer<D>,
914    out_ptr: T::Pointer<D>,
915    n_elements: i32,
916) where
917    T::I32Tensor: types::Tensor<i32, 1>,
918    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
919    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
920{
921    let pid = T::program_id(Axis::X);
922    let block_start = pid * BLOCK_SIZE;
923    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
924    let in_bounds = offsets.lt(n_elements);
925    let a = T::load(
926        a_ptr.add_offsets(offsets),
927        Some(in_bounds),
928        None,
929        &[],
930        None,
931        None,
932        None,
933        false,
934    );
935    let b = T::load(
936        b_ptr.add_offsets(offsets),
937        Some(in_bounds),
938        None,
939        &[],
940        None,
941        None,
942        None,
943        false,
944    );
945    T::store(
946        out_ptr.add_offsets(offsets),
947        T::maximum(a, b),
948        Some(in_bounds),
949        &[],
950        None,
951        None,
952    );
953}
954
955/// Backward: pass dy to the input that was larger, 0 to the other.
956#[kernel]
957pub fn elemwise_max_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
958    dy_ptr: T::Pointer<D>,
959    a_ptr: T::Pointer<D>,
960    b_ptr: T::Pointer<D>,
961    da_ptr: T::Pointer<D>,
962    db_ptr: T::Pointer<D>,
963    n_elements: i32,
964) where
965    T::I32Tensor: types::Tensor<i32, 1>,
966    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
967    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
968{
969    let pid = T::program_id(Axis::X);
970    let block_start = pid * BLOCK_SIZE;
971    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
972    let in_bounds = offsets.lt(n_elements);
973    let dy = T::load(
974        dy_ptr.add_offsets(offsets),
975        Some(in_bounds),
976        None,
977        &[],
978        None,
979        None,
980        None,
981        false,
982    );
983    let a = T::load(
984        a_ptr.add_offsets(offsets),
985        Some(in_bounds),
986        None,
987        &[],
988        None,
989        None,
990        None,
991        false,
992    );
993    let b = T::load(
994        b_ptr.add_offsets(offsets),
995        Some(in_bounds),
996        None,
997        &[],
998        None,
999        None,
1000        None,
1001        false,
1002    );
1003    let z = T::zeros_like(dy);
1004    let a_is_max = T::ge(a, b);
1005    T::store(
1006        da_ptr.add_offsets(offsets),
1007        T::where_(a_is_max, dy, z),
1008        Some(in_bounds),
1009        &[],
1010        None,
1011        None,
1012    );
1013    T::store(
1014        db_ptr.add_offsets(offsets),
1015        T::where_(a_is_max, z, dy),
1016        Some(in_bounds),
1017        &[],
1018        None,
1019        None,
1020    );
1021}
1022
1023impl_binary_num_runtime_op_with_bwd!(ElemwiseMaxForward);
1024
1025// ── ElemMean ──────────────────────────────────────────────────────────────────
1026
1027/// Forward: out = (a + b) / 2
1028#[kernel]
1029pub fn elemwise_mean_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1030    a_ptr: T::Pointer<D>,
1031    b_ptr: T::Pointer<D>,
1032    out_ptr: T::Pointer<D>,
1033    n_elements: i32,
1034) where
1035    T::I32Tensor: types::Tensor<i32, 1>,
1036    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1037    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1038{
1039    let pid = T::program_id(Axis::X);
1040    let block_start = pid * BLOCK_SIZE;
1041    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1042    let in_bounds = offsets.lt(n_elements);
1043    let a = T::load(
1044        a_ptr.add_offsets(offsets),
1045        Some(in_bounds),
1046        None,
1047        &[],
1048        None,
1049        None,
1050        None,
1051        false,
1052    );
1053    let b = T::load(
1054        b_ptr.add_offsets(offsets),
1055        Some(in_bounds),
1056        None,
1057        &[],
1058        None,
1059        None,
1060        None,
1061        false,
1062    );
1063    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
1064    T::store(
1065        out_ptr.add_offsets(offsets),
1066        (a + b) / two,
1067        Some(in_bounds),
1068        &[],
1069        None,
1070        None,
1071    );
1072}
1073
1074/// Backward: da = db = dy / 2
1075#[kernel]
1076pub fn elemwise_mean_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1077    dy_ptr: T::Pointer<D>,
1078    da_ptr: T::Pointer<D>,
1079    db_ptr: T::Pointer<D>,
1080    n_elements: i32,
1081) where
1082    T::I32Tensor: types::Tensor<i32, 1>,
1083    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1084    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1085{
1086    let pid = T::program_id(Axis::X);
1087    let block_start = pid * BLOCK_SIZE;
1088    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1089    let in_bounds = offsets.lt(n_elements);
1090    let dy = T::load(
1091        dy_ptr.add_offsets(offsets),
1092        Some(in_bounds),
1093        None,
1094        &[],
1095        None,
1096        None,
1097        None,
1098        false,
1099    );
1100    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
1101    let half_dy = dy / two;
1102    T::store(
1103        da_ptr.add_offsets(offsets),
1104        half_dy,
1105        Some(in_bounds),
1106        &[],
1107        None,
1108        None,
1109    );
1110    T::store(
1111        db_ptr.add_offsets(offsets),
1112        half_dy,
1113        Some(in_bounds),
1114        &[],
1115        None,
1116        None,
1117    );
1118}
1119
1120impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for ElemwiseMeanForward<D> {
1121    fn n_activation_inputs(&self) -> usize {
1122        2
1123    }
1124    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
1125        vec![]
1126    }
1127    fn pack_args(
1128        &self,
1129        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1130        _: &[teeny_core::model::RawPtr],
1131        output: teeny_core::model::RawPtr,
1132        output_shape: &[usize],
1133        _: i32,
1134        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1135    ) {
1136        let n: usize = output_shape.iter().product();
1137        visitor.visit_ptr(inputs[0].0);
1138        visitor.visit_ptr(inputs[1].0);
1139        visitor.visit_ptr(output);
1140        visitor.visit_i32(n as i32);
1141    }
1142    fn block(&self) -> [u32; 3] {
1143        [self.block_size as u32, 1, 1]
1144    }
1145    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
1146        let n: usize = output_shape.iter().product();
1147        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
1148    }
1149    #[cfg(feature = "training")]
1150    fn has_backward(&self) -> bool {
1151        true
1152    }
1153    #[cfg(feature = "training")]
1154    fn pack_backward_args(
1155        &self,
1156        _: &[(teeny_core::model::RawPtr, &[usize])],
1157        _: &[teeny_core::model::RawPtr],
1158        _: teeny_core::model::RawPtr,
1159        output_shape: &[usize],
1160        grad_output: teeny_core::model::RawPtr,
1161        _: i32,
1162        grad_inputs: &[teeny_core::model::RawPtr],
1163        _: &[teeny_core::model::RawPtr],
1164        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1165    ) {
1166        let n: usize = output_shape.iter().product();
1167        visitor.visit_ptr(grad_output);
1168        visitor.visit_ptr(grad_inputs[0]);
1169        visitor.visit_ptr(grad_inputs[1]);
1170        visitor.visit_i32(n as i32);
1171    }
1172    #[cfg(feature = "training")]
1173    fn backward_block(&self) -> [u32; 3] {
1174        [self.block_size as u32, 1, 1]
1175    }
1176    #[cfg(feature = "training")]
1177    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
1178        let n: usize = output_shape.iter().product();
1179        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
1180    }
1181}
1182
1183// ── ElemSum (binary add — identical to ElemwiseAdd semantics) ─────────────────
1184
1185/// Forward: out = a + b  (binary ElemSum)
1186#[kernel]
1187pub fn elemwise_sum_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1188    a_ptr: T::Pointer<D>,
1189    b_ptr: T::Pointer<D>,
1190    out_ptr: T::Pointer<D>,
1191    n_elements: i32,
1192) where
1193    T::I32Tensor: types::Tensor<i32, 1>,
1194    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1195    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1196{
1197    let pid = T::program_id(Axis::X);
1198    let block_start = pid * BLOCK_SIZE;
1199    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1200    let in_bounds = offsets.lt(n_elements);
1201    let a = T::load(
1202        a_ptr.add_offsets(offsets),
1203        Some(in_bounds),
1204        None,
1205        &[],
1206        None,
1207        None,
1208        None,
1209        false,
1210    );
1211    let b = T::load(
1212        b_ptr.add_offsets(offsets),
1213        Some(in_bounds),
1214        None,
1215        &[],
1216        None,
1217        None,
1218        None,
1219        false,
1220    );
1221    T::store(
1222        out_ptr.add_offsets(offsets),
1223        a + b,
1224        Some(in_bounds),
1225        &[],
1226        None,
1227        None,
1228    );
1229}
1230
1231/// Backward: da = db = dy
1232#[kernel]
1233pub fn elemwise_sum_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1234    dy_ptr: T::Pointer<D>,
1235    da_ptr: T::Pointer<D>,
1236    db_ptr: T::Pointer<D>,
1237    n_elements: i32,
1238) where
1239    T::I32Tensor: types::Tensor<i32, 1>,
1240    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1241    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1242{
1243    let pid = T::program_id(Axis::X);
1244    let block_start = pid * BLOCK_SIZE;
1245    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1246    let in_bounds = offsets.lt(n_elements);
1247    let dy = T::load(
1248        dy_ptr.add_offsets(offsets),
1249        Some(in_bounds),
1250        None,
1251        &[],
1252        None,
1253        None,
1254        None,
1255        false,
1256    );
1257    T::store(
1258        da_ptr.add_offsets(offsets),
1259        dy,
1260        Some(in_bounds),
1261        &[],
1262        None,
1263        None,
1264    );
1265    T::store(
1266        db_ptr.add_offsets(offsets),
1267        dy,
1268        Some(in_bounds),
1269        &[],
1270        None,
1271        None,
1272    );
1273}
1274
1275impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for ElemwiseSumForward<D> {
1276    fn n_activation_inputs(&self) -> usize {
1277        2
1278    }
1279    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
1280        vec![]
1281    }
1282    fn pack_args(
1283        &self,
1284        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1285        _: &[teeny_core::model::RawPtr],
1286        output: teeny_core::model::RawPtr,
1287        output_shape: &[usize],
1288        _: i32,
1289        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1290    ) {
1291        let n: usize = output_shape.iter().product();
1292        visitor.visit_ptr(inputs[0].0);
1293        visitor.visit_ptr(inputs[1].0);
1294        visitor.visit_ptr(output);
1295        visitor.visit_i32(n as i32);
1296    }
1297    fn block(&self) -> [u32; 3] {
1298        [self.block_size as u32, 1, 1]
1299    }
1300    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
1301        let n: usize = output_shape.iter().product();
1302        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
1303    }
1304    #[cfg(feature = "training")]
1305    fn has_backward(&self) -> bool {
1306        true
1307    }
1308    #[cfg(feature = "training")]
1309    fn pack_backward_args(
1310        &self,
1311        _: &[(teeny_core::model::RawPtr, &[usize])],
1312        _: &[teeny_core::model::RawPtr],
1313        _: teeny_core::model::RawPtr,
1314        output_shape: &[usize],
1315        grad_output: teeny_core::model::RawPtr,
1316        _: i32,
1317        grad_inputs: &[teeny_core::model::RawPtr],
1318        _: &[teeny_core::model::RawPtr],
1319        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1320    ) {
1321        let n: usize = output_shape.iter().product();
1322        visitor.visit_ptr(grad_output);
1323        visitor.visit_ptr(grad_inputs[0]);
1324        visitor.visit_ptr(grad_inputs[1]);
1325        visitor.visit_i32(n as i32);
1326    }
1327    #[cfg(feature = "training")]
1328    fn backward_block(&self) -> [u32; 3] {
1329        [self.block_size as u32, 1, 1]
1330    }
1331    #[cfg(feature = "training")]
1332    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
1333        let n: usize = output_shape.iter().product();
1334        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
1335    }
1336}
1337
1338// ── Comparison ops (output 0.0/1.0 as float) ──────────────────────────────────
1339
1340/// Forward: out = 1.0 if a == b else 0.0
1341#[kernel]
1342pub fn elemwise_equal_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1343    a_ptr: T::Pointer<D>,
1344    b_ptr: T::Pointer<D>,
1345    out_ptr: T::Pointer<D>,
1346    n_elements: i32,
1347) where
1348    T::I32Tensor: types::Tensor<i32, 1>,
1349    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1350    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1351{
1352    let pid = T::program_id(Axis::X);
1353    let block_start = pid * BLOCK_SIZE;
1354    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1355    let in_bounds = offsets.lt(n_elements);
1356    let a = T::load(
1357        a_ptr.add_offsets(offsets),
1358        Some(in_bounds),
1359        None,
1360        &[],
1361        None,
1362        None,
1363        None,
1364        false,
1365    );
1366    let b = T::load(
1367        b_ptr.add_offsets(offsets),
1368        Some(in_bounds),
1369        None,
1370        &[],
1371        None,
1372        None,
1373        None,
1374        false,
1375    );
1376    let cond = T::eq(a, b);
1377    let one = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
1378    let zero = T::zeros_like(a);
1379    T::store(
1380        out_ptr.add_offsets(offsets),
1381        T::where_(cond, one, zero),
1382        Some(in_bounds),
1383        &[],
1384        None,
1385        None,
1386    );
1387}
1388
1389impl_binary_num_runtime_op_no_bwd!(ElemwiseEqualForward);
1390
1391/// Forward: out = 1.0 if a > b else 0.0
1392#[kernel]
1393pub fn elemwise_greater_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1394    a_ptr: T::Pointer<D>,
1395    b_ptr: T::Pointer<D>,
1396    out_ptr: T::Pointer<D>,
1397    n_elements: i32,
1398) where
1399    T::I32Tensor: types::Tensor<i32, 1>,
1400    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1401    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1402{
1403    let pid = T::program_id(Axis::X);
1404    let block_start = pid * BLOCK_SIZE;
1405    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1406    let in_bounds = offsets.lt(n_elements);
1407    let a = T::load(
1408        a_ptr.add_offsets(offsets),
1409        Some(in_bounds),
1410        None,
1411        &[],
1412        None,
1413        None,
1414        None,
1415        false,
1416    );
1417    let b = T::load(
1418        b_ptr.add_offsets(offsets),
1419        Some(in_bounds),
1420        None,
1421        &[],
1422        None,
1423        None,
1424        None,
1425        false,
1426    );
1427    let cond = T::gt(a, b);
1428    let one = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
1429    let zero = T::zeros_like(a);
1430    T::store(
1431        out_ptr.add_offsets(offsets),
1432        T::where_(cond, one, zero),
1433        Some(in_bounds),
1434        &[],
1435        None,
1436        None,
1437    );
1438}
1439
1440impl_binary_num_runtime_op_no_bwd!(ElemwiseGreaterForward);
1441
1442/// Forward: out = 1.0 if a >= b else 0.0
1443#[kernel]
1444pub fn elemwise_greater_equal_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1445    a_ptr: T::Pointer<D>,
1446    b_ptr: T::Pointer<D>,
1447    out_ptr: T::Pointer<D>,
1448    n_elements: i32,
1449) where
1450    T::I32Tensor: types::Tensor<i32, 1>,
1451    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1452    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1453{
1454    let pid = T::program_id(Axis::X);
1455    let block_start = pid * BLOCK_SIZE;
1456    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1457    let in_bounds = offsets.lt(n_elements);
1458    let a = T::load(
1459        a_ptr.add_offsets(offsets),
1460        Some(in_bounds),
1461        None,
1462        &[],
1463        None,
1464        None,
1465        None,
1466        false,
1467    );
1468    let b = T::load(
1469        b_ptr.add_offsets(offsets),
1470        Some(in_bounds),
1471        None,
1472        &[],
1473        None,
1474        None,
1475        None,
1476        false,
1477    );
1478    let cond = T::ge(a, b);
1479    let one = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
1480    let zero = T::zeros_like(a);
1481    T::store(
1482        out_ptr.add_offsets(offsets),
1483        T::where_(cond, one, zero),
1484        Some(in_bounds),
1485        &[],
1486        None,
1487        None,
1488    );
1489}
1490
1491impl_binary_num_runtime_op_no_bwd!(ElemwiseGreaterEqualForward);
1492
1493/// Forward: out = 1.0 if a < b else 0.0
1494#[kernel]
1495pub fn elemwise_less_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1496    a_ptr: T::Pointer<D>,
1497    b_ptr: T::Pointer<D>,
1498    out_ptr: T::Pointer<D>,
1499    n_elements: i32,
1500) where
1501    T::I32Tensor: types::Tensor<i32, 1>,
1502    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1503    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1504{
1505    let pid = T::program_id(Axis::X);
1506    let block_start = pid * BLOCK_SIZE;
1507    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1508    let in_bounds = offsets.lt(n_elements);
1509    let a = T::load(
1510        a_ptr.add_offsets(offsets),
1511        Some(in_bounds),
1512        None,
1513        &[],
1514        None,
1515        None,
1516        None,
1517        false,
1518    );
1519    let b = T::load(
1520        b_ptr.add_offsets(offsets),
1521        Some(in_bounds),
1522        None,
1523        &[],
1524        None,
1525        None,
1526        None,
1527        false,
1528    );
1529    let cond = T::lt(a, b);
1530    let one = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
1531    let zero = T::zeros_like(a);
1532    T::store(
1533        out_ptr.add_offsets(offsets),
1534        T::where_(cond, one, zero),
1535        Some(in_bounds),
1536        &[],
1537        None,
1538        None,
1539    );
1540}
1541
1542impl_binary_num_runtime_op_no_bwd!(ElemwiseLessForward);
1543
1544/// Forward: out = 1.0 if a <= b else 0.0
1545#[kernel]
1546pub fn elemwise_less_equal_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
1547    a_ptr: T::Pointer<D>,
1548    b_ptr: T::Pointer<D>,
1549    out_ptr: T::Pointer<D>,
1550    n_elements: i32,
1551) where
1552    T::I32Tensor: types::Tensor<i32, 1>,
1553    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1554    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1555{
1556    let pid = T::program_id(Axis::X);
1557    let block_start = pid * BLOCK_SIZE;
1558    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1559    let in_bounds = offsets.lt(n_elements);
1560    let a = T::load(
1561        a_ptr.add_offsets(offsets),
1562        Some(in_bounds),
1563        None,
1564        &[],
1565        None,
1566        None,
1567        None,
1568        false,
1569    );
1570    let b = T::load(
1571        b_ptr.add_offsets(offsets),
1572        Some(in_bounds),
1573        None,
1574        &[],
1575        None,
1576        None,
1577        None,
1578        false,
1579    );
1580    let cond = T::le(a, b);
1581    let one = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
1582    let zero = T::zeros_like(a);
1583    T::store(
1584        out_ptr.add_offsets(offsets),
1585        T::where_(cond, one, zero),
1586        Some(in_bounds),
1587        &[],
1588        None,
1589        None,
1590    );
1591}
1592
1593impl_binary_num_runtime_op_no_bwd!(ElemwiseLessEqualForward);
1594
1595// ── Where (3-input: condition, x, y) ─────────────────────────────────────────
1596//
1597// Condition is stored as same D type: 0 = false, non-zero = true.
1598
1599/// Forward: out = x where cond != 0 else y
1600#[kernel]
1601pub fn elemwise_where_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1602    cond_ptr: T::Pointer<D>,
1603    x_ptr: T::Pointer<D>,
1604    y_ptr: T::Pointer<D>,
1605    out_ptr: T::Pointer<D>,
1606    n_elements: i32,
1607) where
1608    T::I32Tensor: types::Tensor<i32, 1>,
1609    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1610    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1611{
1612    let pid = T::program_id(Axis::X);
1613    let block_start = pid * BLOCK_SIZE;
1614    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1615    let in_bounds = offsets.lt(n_elements);
1616    let cond = T::load(
1617        cond_ptr.add_offsets(offsets),
1618        Some(in_bounds),
1619        None,
1620        &[],
1621        None,
1622        None,
1623        None,
1624        false,
1625    );
1626    let x = T::load(
1627        x_ptr.add_offsets(offsets),
1628        Some(in_bounds),
1629        None,
1630        &[],
1631        None,
1632        None,
1633        None,
1634        false,
1635    );
1636    let y = T::load(
1637        y_ptr.add_offsets(offsets),
1638        Some(in_bounds),
1639        None,
1640        &[],
1641        None,
1642        None,
1643        None,
1644        false,
1645    );
1646    let zero = T::zeros_like(cond);
1647    let bool_cond = T::ne(cond, zero);
1648    T::store(
1649        out_ptr.add_offsets(offsets),
1650        T::where_(bool_cond, x, y),
1651        Some(in_bounds),
1652        &[],
1653        None,
1654        None,
1655    );
1656}
1657
1658/// Backward: dx = where(cond, dy, 0),  dy_in = where(cond, 0, dy)
1659#[kernel]
1660pub fn elemwise_where_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1661    dy_ptr: T::Pointer<D>,
1662    cond_ptr: T::Pointer<D>,
1663    dx_ptr: T::Pointer<D>,
1664    dy_in_ptr: T::Pointer<D>,
1665    n_elements: i32,
1666) where
1667    T::I32Tensor: types::Tensor<i32, 1>,
1668    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1669    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1670{
1671    let pid = T::program_id(Axis::X);
1672    let block_start = pid * BLOCK_SIZE;
1673    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1674    let in_bounds = offsets.lt(n_elements);
1675    let dy = T::load(
1676        dy_ptr.add_offsets(offsets),
1677        Some(in_bounds),
1678        None,
1679        &[],
1680        None,
1681        None,
1682        None,
1683        false,
1684    );
1685    let cond = T::load(
1686        cond_ptr.add_offsets(offsets),
1687        Some(in_bounds),
1688        None,
1689        &[],
1690        None,
1691        None,
1692        None,
1693        false,
1694    );
1695    let zero = T::zeros_like(dy);
1696    let bool_cond = T::ne(cond, T::zeros_like(cond));
1697    T::store(
1698        dx_ptr.add_offsets(offsets),
1699        T::where_(bool_cond, dy, zero),
1700        Some(in_bounds),
1701        &[],
1702        None,
1703        None,
1704    );
1705    T::store(
1706        dy_in_ptr.add_offsets(offsets),
1707        T::where_(bool_cond, zero, dy),
1708        Some(in_bounds),
1709        &[],
1710        None,
1711        None,
1712    );
1713}
1714
1715impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for ElemwiseWhereForward<D> {
1716    fn n_activation_inputs(&self) -> usize {
1717        3
1718    }
1719    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
1720        vec![]
1721    }
1722    fn pack_args(
1723        &self,
1724        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1725        _: &[teeny_core::model::RawPtr],
1726        output: teeny_core::model::RawPtr,
1727        output_shape: &[usize],
1728        _: i32,
1729        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1730    ) {
1731        let n: usize = output_shape.iter().product();
1732        visitor.visit_ptr(inputs[0].0); // cond
1733        visitor.visit_ptr(inputs[1].0); // x
1734        visitor.visit_ptr(inputs[2].0); // y
1735        visitor.visit_ptr(output);
1736        visitor.visit_i32(n as i32);
1737    }
1738    fn block(&self) -> [u32; 3] {
1739        [self.block_size as u32, 1, 1]
1740    }
1741    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
1742        let n: usize = output_shape.iter().product();
1743        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
1744    }
1745    #[cfg(feature = "training")]
1746    fn has_backward(&self) -> bool {
1747        true
1748    }
1749    #[cfg(feature = "training")]
1750    fn pack_backward_args(
1751        &self,
1752        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1753        _: &[teeny_core::model::RawPtr],
1754        _: teeny_core::model::RawPtr,
1755        output_shape: &[usize],
1756        grad_output: teeny_core::model::RawPtr,
1757        _: i32,
1758        grad_inputs: &[teeny_core::model::RawPtr],
1759        _: &[teeny_core::model::RawPtr],
1760        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1761    ) {
1762        let n: usize = output_shape.iter().product();
1763        visitor.visit_ptr(grad_output);
1764        visitor.visit_ptr(inputs[0].0); // cond
1765        visitor.visit_ptr(grad_inputs[1]); // dx
1766        visitor.visit_ptr(grad_inputs[2]); // dy_in
1767        visitor.visit_i32(n as i32);
1768    }
1769    #[cfg(feature = "training")]
1770    fn backward_block(&self) -> [u32; 3] {
1771        [self.block_size as u32, 1, 1]
1772    }
1773    #[cfg(feature = "training")]
1774    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
1775        let n: usize = output_shape.iter().product();
1776        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
1777    }
1778}
1779
1780// ── Clip (3-input: x, min_val, max_val) ───────────────────────────────────────
1781//
1782// For simplicity this kernel takes min and max as f32 scalar kernel params
1783// rather than tensors.  The lowering packs them as f32.
1784
1785/// Forward: out = clamp(x, min_val, max_val)
1786#[kernel]
1787pub fn elemwise_clip_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1788    x_ptr: T::Pointer<D>,
1789    out_ptr: T::Pointer<D>,
1790    n_elements: i32,
1791    min_val: f32,
1792    max_val: f32,
1793) where
1794    T::I32Tensor: types::Tensor<i32, 1>,
1795    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1796    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1797{
1798    let pid = T::program_id(Axis::X);
1799    let block_start = pid * BLOCK_SIZE;
1800    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1801    let in_bounds = offsets.lt(n_elements);
1802    let x = T::load(
1803        x_ptr.add_offsets(offsets),
1804        Some(in_bounds),
1805        None,
1806        &[],
1807        None,
1808        None,
1809        None,
1810        false,
1811    );
1812    let lo = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], min_val), None, false);
1813    let hi = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], max_val), None, false);
1814    let y = T::clamp(x, lo, hi);
1815    T::store(
1816        out_ptr.add_offsets(offsets),
1817        y,
1818        Some(in_bounds),
1819        &[],
1820        None,
1821        None,
1822    );
1823}
1824
1825/// Backward: pass dy through only where x was in [min_val, max_val]
1826#[kernel]
1827pub fn elemwise_clip_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1828    dy_ptr: T::Pointer<D>,
1829    x_ptr: T::Pointer<D>,
1830    dx_ptr: T::Pointer<D>,
1831    n_elements: i32,
1832    min_val: f32,
1833    max_val: f32,
1834) where
1835    T::I32Tensor: types::Tensor<i32, 1>,
1836    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1837    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1838{
1839    let pid = T::program_id(Axis::X);
1840    let block_start = pid * BLOCK_SIZE;
1841    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1842    let in_bounds = offsets.lt(n_elements);
1843    let dy = T::load(
1844        dy_ptr.add_offsets(offsets),
1845        Some(in_bounds),
1846        None,
1847        &[],
1848        None,
1849        None,
1850        None,
1851        false,
1852    );
1853    let x = T::load(
1854        x_ptr.add_offsets(offsets),
1855        Some(in_bounds),
1856        None,
1857        &[],
1858        None,
1859        None,
1860        None,
1861        false,
1862    );
1863    let lo = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], min_val), None, false);
1864    let hi = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], max_val), None, false);
1865    let in_range = T::ge(x, lo) & T::le(x, hi);
1866    let dx = T::where_(in_range, dy, T::zeros_like(dy));
1867    T::store(
1868        dx_ptr.add_offsets(offsets),
1869        dx,
1870        Some(in_bounds),
1871        &[],
1872        None,
1873        None,
1874    );
1875}
1876
1877/// A RuntimeOp wrapper for Clip that stores the min/max scalar params alongside
1878/// the kernel struct (which only stores block_size).
1879pub struct ClipRuntimeOp<D: Float + Send + Sync + 'static> {
1880    pub kernel: ElemwiseClipForward<D>,
1881    pub backward_kernel: ElemwiseClipBackward<D>,
1882    pub min_val: f32,
1883    pub max_val: f32,
1884}
1885
1886impl<D: Float + Send + Sync + 'static> ClipRuntimeOp<D> {
1887    pub fn new(block_size: i32, min_val: f32, max_val: f32) -> Self {
1888        Self {
1889            kernel: ElemwiseClipForward::<D>::new(block_size),
1890            backward_kernel: ElemwiseClipBackward::<D>::new(block_size),
1891            min_val,
1892            max_val,
1893        }
1894    }
1895
1896    pub fn forward_source(&self) -> &str {
1897        &self.kernel.source
1898    }
1899
1900    pub fn backward_source(&self) -> &str {
1901        &self.backward_kernel.source
1902    }
1903
1904    pub fn kernel_name(&self) -> &str {
1905        self.kernel.name
1906    }
1907}
1908
1909impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for ClipRuntimeOp<D> {
1910    fn n_activation_inputs(&self) -> usize {
1911        1
1912    }
1913    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
1914        vec![]
1915    }
1916    fn pack_args(
1917        &self,
1918        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1919        _: &[teeny_core::model::RawPtr],
1920        output: teeny_core::model::RawPtr,
1921        output_shape: &[usize],
1922        _: i32,
1923        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1924    ) {
1925        let n: usize = output_shape.iter().product();
1926        visitor.visit_ptr(inputs[0].0);
1927        visitor.visit_ptr(output);
1928        visitor.visit_i32(n as i32);
1929        visitor.visit_f32(self.min_val);
1930        visitor.visit_f32(self.max_val);
1931    }
1932    fn block(&self) -> [u32; 3] {
1933        [self.kernel.block_size as u32, 1, 1]
1934    }
1935    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
1936        let n: usize = output_shape.iter().product();
1937        [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
1938    }
1939    #[cfg(feature = "training")]
1940    fn has_backward(&self) -> bool {
1941        true
1942    }
1943    #[cfg(feature = "training")]
1944    fn pack_backward_args(
1945        &self,
1946        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1947        _: &[teeny_core::model::RawPtr],
1948        _: teeny_core::model::RawPtr,
1949        output_shape: &[usize],
1950        grad_output: teeny_core::model::RawPtr,
1951        _: i32,
1952        grad_inputs: &[teeny_core::model::RawPtr],
1953        _: &[teeny_core::model::RawPtr],
1954        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1955    ) {
1956        let n: usize = output_shape.iter().product();
1957        visitor.visit_ptr(grad_output);
1958        visitor.visit_ptr(inputs[0].0);
1959        visitor.visit_ptr(grad_inputs[0]);
1960        visitor.visit_i32(n as i32);
1961        visitor.visit_f32(self.min_val);
1962        visitor.visit_f32(self.max_val);
1963    }
1964    #[cfg(feature = "training")]
1965    fn backward_block(&self) -> [u32; 3] {
1966        [self.kernel.block_size as u32, 1, 1]
1967    }
1968    #[cfg(feature = "training")]
1969    fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
1970        let n: usize = output_shape.iter().product();
1971        [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
1972    }
1973}