1#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21 types::{AddOffsets, Comparison},
22 *,
23};
24
25#[kernel]
37pub fn bce_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
38 input_ptr: T::Pointer<f32>,
39 target_ptr: T::Pointer<f32>,
40 out_ptr: T::Pointer<f32>,
41 n_elements: i32,
42) where
43 T::I32Tensor: types::Tensor<i32, 1>,
44 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
45 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
46{
47 let pid = T::program_id(Axis::X);
48 let block_start = pid * BLOCK_SIZE;
49 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
50 let in_bounds = offsets.lt(n_elements);
51 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
52
53 let inp = T::load(
54 input_ptr.add_offsets(offsets),
55 Some(in_bounds),
56 Some(zeros),
57 &[],
58 None,
59 None,
60 None,
61 false,
62 );
63 let tgt = T::load(
64 target_ptr.add_offsets(offsets),
65 Some(in_bounds),
66 Some(zeros),
67 &[],
68 None,
69 None,
70 None,
71 false,
72 );
73
74 let one = T::full(&[BLOCK_SIZE], 1.0_f32);
75 let eps = T::full(&[BLOCK_SIZE], 1e-7_f32);
77 let one_minus_eps = T::full(&[BLOCK_SIZE], 1.0_f32 - 1e-7_f32);
78 let inp_c = T::clamp(inp, eps, one_minus_eps);
79
80 let loss = T::full(&[BLOCK_SIZE], -1.0_f32)
81 * (tgt * T::log(inp_c) + (one - tgt) * T::log(one - inp_c));
82 T::store(
83 out_ptr.add_offsets(offsets),
84 loss,
85 Some(in_bounds),
86 &[],
87 None,
88 None,
89 );
90}
91
92#[kernel]
98pub fn bce_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
99 dy_ptr: T::Pointer<f32>,
100 input_ptr: T::Pointer<f32>,
101 target_ptr: T::Pointer<f32>,
102 dx_ptr: T::Pointer<f32>,
103 n_elements: i32,
104) where
105 T::I32Tensor: types::Tensor<i32, 1>,
106 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
107 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
108{
109 let pid = T::program_id(Axis::X);
110 let block_start = pid * BLOCK_SIZE;
111 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
112 let in_bounds = offsets.lt(n_elements);
113 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
114
115 let dy = T::load(
116 dy_ptr.add_offsets(offsets),
117 Some(in_bounds),
118 Some(zeros),
119 &[],
120 None,
121 None,
122 None,
123 false,
124 );
125 let inp = T::load(
126 input_ptr.add_offsets(offsets),
127 Some(in_bounds),
128 Some(zeros),
129 &[],
130 None,
131 None,
132 None,
133 false,
134 );
135 let tgt = T::load(
136 target_ptr.add_offsets(offsets),
137 Some(in_bounds),
138 Some(zeros),
139 &[],
140 None,
141 None,
142 None,
143 false,
144 );
145
146 let one = T::full(&[BLOCK_SIZE], 1.0_f32);
147 let eps = T::full(&[BLOCK_SIZE], 1e-7_f32);
148 let one_minus_eps = T::full(&[BLOCK_SIZE], 1.0_f32 - 1e-7_f32);
149 let inp_c = T::clamp(inp, eps, one_minus_eps);
150
151 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
153 let dx_raw = neg_one * (tgt / inp_c - (one - tgt) / (one - inp_c));
154 let dx = dx_raw * dy;
155 T::store(
156 dx_ptr.add_offsets(offsets),
157 dx,
158 Some(in_bounds),
159 &[],
160 None,
161 None,
162 );
163}
164
165#[kernel]
174pub fn bce_with_logits_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
175 input_ptr: T::Pointer<f32>,
176 target_ptr: T::Pointer<f32>,
177 out_ptr: T::Pointer<f32>,
178 n_elements: i32,
179) where
180 T::I32Tensor: types::Tensor<i32, 1>,
181 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
182 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
183{
184 let pid = T::program_id(Axis::X);
185 let block_start = pid * BLOCK_SIZE;
186 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
187 let in_bounds = offsets.lt(n_elements);
188 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
189
190 let inp = T::load(
191 input_ptr.add_offsets(offsets),
192 Some(in_bounds),
193 Some(zeros),
194 &[],
195 None,
196 None,
197 None,
198 false,
199 );
200 let tgt = T::load(
201 target_ptr.add_offsets(offsets),
202 Some(in_bounds),
203 Some(zeros),
204 &[],
205 None,
206 None,
207 None,
208 false,
209 );
210
211 let one = T::full(&[BLOCK_SIZE], 1.0_f32);
212 let relu_x = T::maximum(inp, zeros);
214 let neg_abs_x = T::full(&[BLOCK_SIZE], -1.0_f32) * T::abs(inp);
215 let loss = relu_x - inp * tgt + T::log(one + T::exp(neg_abs_x));
216 T::store(
217 out_ptr.add_offsets(offsets),
218 loss,
219 Some(in_bounds),
220 &[],
221 None,
222 None,
223 );
224}
225
226#[kernel]
230pub fn bce_with_logits_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
231 dy_ptr: T::Pointer<f32>,
232 input_ptr: T::Pointer<f32>,
233 target_ptr: T::Pointer<f32>,
234 dx_ptr: T::Pointer<f32>,
235 n_elements: i32,
236) where
237 T::I32Tensor: types::Tensor<i32, 1>,
238 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
239 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
240{
241 let pid = T::program_id(Axis::X);
242 let block_start = pid * BLOCK_SIZE;
243 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
244 let in_bounds = offsets.lt(n_elements);
245 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
246
247 let dy = T::load(
248 dy_ptr.add_offsets(offsets),
249 Some(in_bounds),
250 Some(zeros),
251 &[],
252 None,
253 None,
254 None,
255 false,
256 );
257 let inp = T::load(
258 input_ptr.add_offsets(offsets),
259 Some(in_bounds),
260 Some(zeros),
261 &[],
262 None,
263 None,
264 None,
265 false,
266 );
267 let tgt = T::load(
268 target_ptr.add_offsets(offsets),
269 Some(in_bounds),
270 Some(zeros),
271 &[],
272 None,
273 None,
274 None,
275 false,
276 );
277
278 let one = T::full(&[BLOCK_SIZE], 1.0_f32);
279 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
280 let sig = one / (one + T::exp(neg_one * inp));
282 let dx = (sig - tgt) * dy;
283 T::store(
284 dx_ptr.add_offsets(offsets),
285 dx,
286 Some(in_bounds),
287 &[],
288 None,
289 None,
290 );
291}
292
293#[kernel]
303pub fn soft_margin_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
304 input_ptr: T::Pointer<f32>,
305 target_ptr: T::Pointer<f32>,
306 out_ptr: T::Pointer<f32>,
307 n_elements: i32,
308) where
309 T::I32Tensor: types::Tensor<i32, 1>,
310 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
311 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
312{
313 let pid = T::program_id(Axis::X);
314 let block_start = pid * BLOCK_SIZE;
315 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
316 let in_bounds = offsets.lt(n_elements);
317 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
318
319 let inp = T::load(
320 input_ptr.add_offsets(offsets),
321 Some(in_bounds),
322 Some(zeros),
323 &[],
324 None,
325 None,
326 None,
327 false,
328 );
329 let tgt = T::load(
330 target_ptr.add_offsets(offsets),
331 Some(in_bounds),
332 Some(zeros),
333 &[],
334 None,
335 None,
336 None,
337 false,
338 );
339
340 let one = T::full(&[BLOCK_SIZE], 1.0_f32);
341 let neg_tx = T::full(&[BLOCK_SIZE], -1.0_f32) * tgt * inp;
343 let loss = T::maximum(neg_tx, zeros)
344 + T::log(one + T::exp(T::full(&[BLOCK_SIZE], -1.0_f32) * T::abs(tgt * inp)));
345 T::store(
346 out_ptr.add_offsets(offsets),
347 loss,
348 Some(in_bounds),
349 &[],
350 None,
351 None,
352 );
353}
354
355#[kernel]
361pub fn soft_margin_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
362 dy_ptr: T::Pointer<f32>,
363 input_ptr: T::Pointer<f32>,
364 target_ptr: T::Pointer<f32>,
365 dx_ptr: T::Pointer<f32>,
366 n_elements: i32,
367) where
368 T::I32Tensor: types::Tensor<i32, 1>,
369 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
370 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
371{
372 let pid = T::program_id(Axis::X);
373 let block_start = pid * BLOCK_SIZE;
374 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
375 let in_bounds = offsets.lt(n_elements);
376 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
377
378 let dy = T::load(
379 dy_ptr.add_offsets(offsets),
380 Some(in_bounds),
381 Some(zeros),
382 &[],
383 None,
384 None,
385 None,
386 false,
387 );
388 let inp = T::load(
389 input_ptr.add_offsets(offsets),
390 Some(in_bounds),
391 Some(zeros),
392 &[],
393 None,
394 None,
395 None,
396 false,
397 );
398 let tgt = T::load(
399 target_ptr.add_offsets(offsets),
400 Some(in_bounds),
401 Some(zeros),
402 &[],
403 None,
404 None,
405 None,
406 false,
407 );
408
409 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
410 let one = T::full(&[BLOCK_SIZE], 1.0_f32);
411 let neg_tx = neg_one * tgt * inp;
413 let sig_neg_tx = one / (one + T::exp(neg_one * neg_tx));
414 let dx = neg_one * tgt * sig_neg_tx * dy;
416 T::store(
417 dx_ptr.add_offsets(offsets),
418 dx,
419 Some(in_bounds),
420 &[],
421 None,
422 None,
423 );
424}
425
426#[kernel]
438pub fn kl_div_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
439 input_ptr: T::Pointer<f32>,
440 target_ptr: T::Pointer<f32>,
441 out_ptr: T::Pointer<f32>,
442 n_elements: i32,
443) where
444 T::I32Tensor: types::Tensor<i32, 1>,
445 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
446 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
447{
448 let pid = T::program_id(Axis::X);
449 let block_start = pid * BLOCK_SIZE;
450 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
451 let in_bounds = offsets.lt(n_elements);
452 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
453
454 let inp = T::load(
455 input_ptr.add_offsets(offsets),
456 Some(in_bounds),
457 Some(zeros),
458 &[],
459 None,
460 None,
461 None,
462 false,
463 );
464 let tgt = T::load(
465 target_ptr.add_offsets(offsets),
466 Some(in_bounds),
467 Some(zeros),
468 &[],
469 None,
470 None,
471 None,
472 false,
473 );
474
475 let eps = T::full(&[BLOCK_SIZE], 1e-10_f32);
477 let tgt_safe = T::maximum(tgt, eps);
478 let loss_raw = tgt * (T::log(tgt_safe) - inp);
479 let positive = T::gt(tgt, zeros);
480 let loss = T::where_(positive, loss_raw, zeros);
481 T::store(
482 out_ptr.add_offsets(offsets),
483 loss,
484 Some(in_bounds),
485 &[],
486 None,
487 None,
488 );
489}
490
491#[kernel]
495pub fn kl_div_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
496 dy_ptr: T::Pointer<f32>,
497 target_ptr: T::Pointer<f32>,
498 dx_ptr: T::Pointer<f32>,
499 n_elements: i32,
500) where
501 T::I32Tensor: types::Tensor<i32, 1>,
502 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
503 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
504{
505 let pid = T::program_id(Axis::X);
506 let block_start = pid * BLOCK_SIZE;
507 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
508 let in_bounds = offsets.lt(n_elements);
509 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
510
511 let dy = T::load(
512 dy_ptr.add_offsets(offsets),
513 Some(in_bounds),
514 Some(zeros),
515 &[],
516 None,
517 None,
518 None,
519 false,
520 );
521 let tgt = T::load(
522 target_ptr.add_offsets(offsets),
523 Some(in_bounds),
524 Some(zeros),
525 &[],
526 None,
527 None,
528 None,
529 false,
530 );
531
532 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
533 let dx = neg_one * tgt * dy;
534 T::store(
535 dx_ptr.add_offsets(offsets),
536 dx,
537 Some(in_bounds),
538 &[],
539 None,
540 None,
541 );
542}
543
544#[kernel]
552pub fn poisson_nll_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
553 input_ptr: T::Pointer<f32>,
554 target_ptr: T::Pointer<f32>,
555 out_ptr: T::Pointer<f32>,
556 n_elements: i32,
557) where
558 T::I32Tensor: types::Tensor<i32, 1>,
559 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
560 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
561{
562 let pid = T::program_id(Axis::X);
563 let block_start = pid * BLOCK_SIZE;
564 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
565 let in_bounds = offsets.lt(n_elements);
566 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
567
568 let inp = T::load(
569 input_ptr.add_offsets(offsets),
570 Some(in_bounds),
571 Some(zeros),
572 &[],
573 None,
574 None,
575 None,
576 false,
577 );
578 let tgt = T::load(
579 target_ptr.add_offsets(offsets),
580 Some(in_bounds),
581 Some(zeros),
582 &[],
583 None,
584 None,
585 None,
586 false,
587 );
588
589 let loss = T::exp(inp) - tgt * inp;
591 T::store(
592 out_ptr.add_offsets(offsets),
593 loss,
594 Some(in_bounds),
595 &[],
596 None,
597 None,
598 );
599}
600
601#[kernel]
605pub fn poisson_nll_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
606 dy_ptr: T::Pointer<f32>,
607 input_ptr: T::Pointer<f32>,
608 target_ptr: T::Pointer<f32>,
609 dx_ptr: T::Pointer<f32>,
610 n_elements: i32,
611) where
612 T::I32Tensor: types::Tensor<i32, 1>,
613 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
614 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
615{
616 let pid = T::program_id(Axis::X);
617 let block_start = pid * BLOCK_SIZE;
618 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
619 let in_bounds = offsets.lt(n_elements);
620 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
621
622 let dy = T::load(
623 dy_ptr.add_offsets(offsets),
624 Some(in_bounds),
625 Some(zeros),
626 &[],
627 None,
628 None,
629 None,
630 false,
631 );
632 let inp = T::load(
633 input_ptr.add_offsets(offsets),
634 Some(in_bounds),
635 Some(zeros),
636 &[],
637 None,
638 None,
639 None,
640 false,
641 );
642 let tgt = T::load(
643 target_ptr.add_offsets(offsets),
644 Some(in_bounds),
645 Some(zeros),
646 &[],
647 None,
648 None,
649 None,
650 false,
651 );
652
653 let dx = (T::exp(inp) - tgt) * dy;
654 T::store(
655 dx_ptr.add_offsets(offsets),
656 dx,
657 Some(in_bounds),
658 &[],
659 None,
660 None,
661 );
662}
663
664#[kernel]
674pub fn gaussian_nll_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
675 input_ptr: T::Pointer<f32>,
676 target_ptr: T::Pointer<f32>,
677 var_ptr: T::Pointer<f32>,
678 out_ptr: T::Pointer<f32>,
679 n_elements: i32,
680 eps_var: f32,
681) where
682 T::I32Tensor: types::Tensor<i32, 1>,
683 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
684 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
685{
686 let pid = T::program_id(Axis::X);
687 let block_start = pid * BLOCK_SIZE;
688 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
689 let in_bounds = offsets.lt(n_elements);
690 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
691
692 let inp = T::load(
693 input_ptr.add_offsets(offsets),
694 Some(in_bounds),
695 Some(zeros),
696 &[],
697 None,
698 None,
699 None,
700 false,
701 );
702 let tgt = T::load(
703 target_ptr.add_offsets(offsets),
704 Some(in_bounds),
705 Some(zeros),
706 &[],
707 None,
708 None,
709 None,
710 false,
711 );
712 let var = T::load(
713 var_ptr.add_offsets(offsets),
714 Some(in_bounds),
715 Some(zeros),
716 &[],
717 None,
718 None,
719 None,
720 false,
721 );
722
723 let eps_t = T::full(&[BLOCK_SIZE], eps_var);
724 let half = T::full(&[BLOCK_SIZE], 0.5_f32);
725 let var_c = T::maximum(var, eps_t);
726 let diff = inp - tgt;
727 let loss = half * (T::log(var_c) + diff * diff / var_c);
728 T::store(
729 out_ptr.add_offsets(offsets),
730 loss,
731 Some(in_bounds),
732 &[],
733 None,
734 None,
735 );
736}
737
738#[kernel]
742pub fn gaussian_nll_loss_backward_input<T: Triton, const BLOCK_SIZE: i32>(
743 dy_ptr: T::Pointer<f32>,
744 input_ptr: T::Pointer<f32>,
745 target_ptr: T::Pointer<f32>,
746 var_ptr: T::Pointer<f32>,
747 dx_ptr: T::Pointer<f32>,
748 n_elements: i32,
749 eps_var: f32,
750) where
751 T::I32Tensor: types::Tensor<i32, 1>,
752 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
753 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
754{
755 let pid = T::program_id(Axis::X);
756 let block_start = pid * BLOCK_SIZE;
757 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
758 let in_bounds = offsets.lt(n_elements);
759 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
760
761 let dy = T::load(
762 dy_ptr.add_offsets(offsets),
763 Some(in_bounds),
764 Some(zeros),
765 &[],
766 None,
767 None,
768 None,
769 false,
770 );
771 let inp = T::load(
772 input_ptr.add_offsets(offsets),
773 Some(in_bounds),
774 Some(zeros),
775 &[],
776 None,
777 None,
778 None,
779 false,
780 );
781 let tgt = T::load(
782 target_ptr.add_offsets(offsets),
783 Some(in_bounds),
784 Some(zeros),
785 &[],
786 None,
787 None,
788 None,
789 false,
790 );
791 let var = T::load(
792 var_ptr.add_offsets(offsets),
793 Some(in_bounds),
794 Some(zeros),
795 &[],
796 None,
797 None,
798 None,
799 false,
800 );
801
802 let eps_t = T::full(&[BLOCK_SIZE], eps_var);
803 let var_c = T::maximum(var, eps_t);
804 let dx = (inp - tgt) / var_c * dy;
805 T::store(
806 dx_ptr.add_offsets(offsets),
807 dx,
808 Some(in_bounds),
809 &[],
810 None,
811 None,
812 );
813}