1#![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
26macro_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); visitor.visit_ptr(inputs[1].0); 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
166macro_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#[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#[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#[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#[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#[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#[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#[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#[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)); 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#[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 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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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); visitor.visit_ptr(inputs[1].0); visitor.visit_ptr(inputs[2].0); 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); visitor.visit_ptr(grad_inputs[1]); visitor.visit_ptr(grad_inputs[2]); 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#[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#[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
1877pub 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}