vision_rs/models/yolo/kernels/loss/
cls.rs1#![allow(non_snake_case)]
20
21use teeny_core::dtype::Float;
22use teeny_macros::kernel;
23use teeny_triton::triton::{
24 types::{AddOffsets, Comparison},
25 *,
26};
27
28#[allow(clippy::erasing_op, clippy::identity_op)]
32#[kernel]
33pub fn yolo_bce_cls_loss_forward<T: Triton, D: Float, const BLOCK_N: i32>(
34 pred_ptr: T::Pointer<D>,
35 target_ptr: T::Pointer<D>,
36 loss_ptr: T::Pointer<D>,
37 N: i32,
38 C: i32,
39) where
40 T::I32Tensor: types::Tensor<i32, 1>,
41 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
42 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
43{
44 let n_start = T::program_id(Axis::X) * BLOCK_N;
45 let n_offs = T::arange(0, BLOCK_N) + n_start;
46 let mask = n_offs.lt(N);
47 let zeros = T::zeros::<D>(&[BLOCK_N]);
48 let ones = T::full(&[BLOCK_N], D::from_f64(1.0));
49
50 let mut acc = zeros;
52 let mut c: i32 = 0;
53 while c < C {
54 let base = c * N;
56 let x = T::load(pred_ptr.add_offsets(n_offs + base), Some(mask), Some(zeros), &[], None, None, None, false);
57 let t = T::load(target_ptr.add_offsets(n_offs + base), Some(mask), Some(zeros), &[], None, None, None, false);
58
59 let relu_x = T::maximum(x, zeros);
61 let log1p_exp = T::log(ones + T::exp(zeros - T::abs(x)));
62 let bce = relu_x - x * t + log1p_exp;
63
64 acc = acc + bce;
65 c += 1;
66 }
67
68 T::store(loss_ptr.add_offsets(n_offs), acc, Some(mask), &[], None, None);
69}
70
71#[allow(clippy::erasing_op, clippy::identity_op)]
81#[kernel]
82pub fn yolo_bce_cls_loss_backward<T: Triton, D: Float, const BLOCK_N: i32>(
83 dy_ptr: T::Pointer<D>,
84 pred_ptr: T::Pointer<D>,
85 target_ptr: T::Pointer<D>,
86 d_pred_ptr: T::Pointer<D>,
87 N: i32,
88 C: i32,
89) where
90 T::I32Tensor: types::Tensor<i32, 1>,
91 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
92 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
93{
94 let n_start = T::program_id(Axis::X) * BLOCK_N;
95 let n_offs = T::arange(0, BLOCK_N) + n_start;
96 let mask = n_offs.lt(N);
97 let zeros = T::zeros::<D>(&[BLOCK_N]);
98 let ones = T::full(&[BLOCK_N], D::from_f64(1.0));
99 let neg_one = T::full(&[BLOCK_N], D::from_f64(-1.0));
100
101 let dy = T::load(dy_ptr.add_offsets(n_offs), Some(mask), Some(zeros), &[], None, None, None, false);
102
103 let mut c: i32 = 0;
104 while c < C {
105 let base = c * N;
106 let x = T::load(pred_ptr.add_offsets(n_offs + base), Some(mask), Some(zeros), &[], None, None, None, false);
107 let t = T::load(target_ptr.add_offsets(n_offs + base), Some(mask), Some(zeros), &[], None, None, None, false);
108
109 let sig = ones / (ones + T::exp(neg_one * x));
111 let grad = dy * (sig - t);
112
113 T::store(d_pred_ptr.add_offsets(n_offs + base), grad, Some(mask), &[], None, None);
114 c += 1;
115 }
116}