1#![allow(non_snake_case, clippy::erasing_op, clippy::identity_op)]
27
28use teeny_core::dtype::Float;
29use teeny_macros::kernel;
30use teeny_triton::triton::{
31 types::{AddOffsets, Comparison},
32 *,
33};
34
35#[kernel]
40pub fn yolo_ciou_loss_forward<T: Triton, D: Float, const BLOCK_N: i32>(
41 pred_ptr: T::Pointer<D>,
42 target_ptr: T::Pointer<D>,
43 loss_ptr: T::Pointer<D>,
44 iou_ptr: T::Pointer<D>,
45 v_ptr: T::Pointer<D>,
46 alpha_ptr: T::Pointer<D>,
47 N: i32,
48) where
49 T::I32Tensor: types::Tensor<i32, 1>,
50 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
51 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
52{
53 let n_start = T::program_id(Axis::X) * BLOCK_N;
54 let n_offs = T::arange(0, BLOCK_N) + n_start;
55 let mask = n_offs.lt(N);
56 let zeros = T::zeros::<D>(&[BLOCK_N]);
57
58 let px = T::load(pred_ptr.add_offsets(n_offs + 0 * N), Some(mask), Some(zeros), &[], None, None, None, false);
60 let py = T::load(pred_ptr.add_offsets(n_offs + 1 * N), Some(mask), Some(zeros), &[], None, None, None, false);
61 let pw = T::load(pred_ptr.add_offsets(n_offs + 2 * N), Some(mask), Some(zeros), &[], None, None, None, false);
62 let ph = T::load(pred_ptr.add_offsets(n_offs + 3 * N), Some(mask), Some(zeros), &[], None, None, None, false);
63
64 let tx = T::load(target_ptr.add_offsets(n_offs + 0 * N), Some(mask), Some(zeros), &[], None, None, None, false);
66 let ty = T::load(target_ptr.add_offsets(n_offs + 1 * N), Some(mask), Some(zeros), &[], None, None, None, false);
67 let tw = T::load(target_ptr.add_offsets(n_offs + 2 * N), Some(mask), Some(zeros), &[], None, None, None, false);
68 let th = T::load(target_ptr.add_offsets(n_offs + 3 * N), Some(mask), Some(zeros), &[], None, None, None, false);
69
70 let half = T::full(&[BLOCK_N], D::from_f64(0.5));
71 let eps = T::full(&[BLOCK_N], D::from_f64(1e-7));
72 let ones = T::full(&[BLOCK_N], D::from_f64(1.0));
73
74 let px1 = px - pw * half;
76 let px2 = px + pw * half;
77 let py1 = py - ph * half;
78 let py2 = py + ph * half;
79
80 let tx1 = tx - tw * half;
82 let tx2 = tx + tw * half;
83 let ty1 = ty - th * half;
84 let ty2 = ty + th * half;
85
86 let ix1 = T::maximum(px1, tx1);
88 let ix2 = T::minimum(px2, tx2);
89 let iy1 = T::maximum(py1, ty1);
90 let iy2 = T::minimum(py2, ty2);
91 let inter_w = T::maximum(ix2 - ix1, zeros);
92 let inter_h = T::maximum(iy2 - iy1, zeros);
93 let inter = inter_w * inter_h;
94
95 let pred_area = pw * ph;
97 let target_area = tw * th;
98 let union = pred_area + target_area - inter;
99
100 let iou = inter / (union + eps);
102
103 let dx = px - tx;
105 let dy = py - ty;
106 let d2 = dx * dx + dy * dy;
107
108 let ex1 = T::minimum(px1, tx1);
110 let ex2 = T::maximum(px2, tx2);
111 let ey1 = T::minimum(py1, ty1);
112 let ey2 = T::maximum(py2, ty2);
113 let ecw = ex2 - ex1;
114 let ech = ey2 - ey1;
115 let c2 = ecw * ecw + ech * ech;
116
117 let atan_t = T::atan(tw / (th + eps));
119 let atan_p = T::atan(pw / (ph + eps));
120 let diff = atan_t - atan_p;
121 let pi2_inv4 = T::full(&[BLOCK_N], D::from_f64(0.405_284_73));
123 let v = pi2_inv4 * diff * diff;
124
125 let alpha = v / (ones - iou + v + eps);
127
128 T::store(iou_ptr.add_offsets(n_offs), iou, Some(mask), &[], None, None);
130 T::store(v_ptr.add_offsets(n_offs), v, Some(mask), &[], None, None);
131 T::store(alpha_ptr.add_offsets(n_offs), alpha, Some(mask), &[], None, None);
132
133 let loss = ones - iou + d2 / (c2 + eps) + alpha * v;
135
136 T::store(loss_ptr.add_offsets(n_offs), loss, Some(mask), &[], None, None);
137}
138
139#[kernel]
156pub fn yolo_ciou_loss_backward<T: Triton, D: Float, const BLOCK_N: i32>(
157 dy_ptr: T::Pointer<D>,
158 pred_ptr: T::Pointer<D>,
159 target_ptr: T::Pointer<D>,
160 iou_ptr: T::Pointer<D>,
161 v_ptr: T::Pointer<D>,
162 alpha_ptr: T::Pointer<D>,
163 d_pred_ptr: T::Pointer<D>,
164 N: i32,
165) where
166 T::I32Tensor: types::Tensor<i32, 1>,
167 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
168 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
169{
170 let n_start = T::program_id(Axis::X) * BLOCK_N;
171 let n_offs = T::arange(0, BLOCK_N) + n_start;
172 let mask = n_offs.lt(N);
173 let zeros = T::zeros::<D>(&[BLOCK_N]);
174
175 let dy = T::load(dy_ptr.add_offsets(n_offs), Some(mask), Some(zeros), &[], None, None, None, false);
176
177 let px = T::load(pred_ptr.add_offsets(n_offs + 0 * N), Some(mask), Some(zeros), &[], None, None, None, false);
179 let py = T::load(pred_ptr.add_offsets(n_offs + 1 * N), Some(mask), Some(zeros), &[], None, None, None, false);
180 let pw = T::load(pred_ptr.add_offsets(n_offs + 2 * N), Some(mask), Some(zeros), &[], None, None, None, false);
181 let ph = T::load(pred_ptr.add_offsets(n_offs + 3 * N), Some(mask), Some(zeros), &[], None, None, None, false);
182
183 let tx = T::load(target_ptr.add_offsets(n_offs + 0 * N), Some(mask), Some(zeros), &[], None, None, None, false);
185 let ty = T::load(target_ptr.add_offsets(n_offs + 1 * N), Some(mask), Some(zeros), &[], None, None, None, false);
186 let tw = T::load(target_ptr.add_offsets(n_offs + 2 * N), Some(mask), Some(zeros), &[], None, None, None, false);
187 let th = T::load(target_ptr.add_offsets(n_offs + 3 * N), Some(mask), Some(zeros), &[], None, None, None, false);
188
189 let iou = T::load(iou_ptr.add_offsets(n_offs), Some(mask), Some(zeros), &[], None, None, None, false);
191 let v = T::load(v_ptr.add_offsets(n_offs), Some(mask), Some(zeros), &[], None, None, None, false);
192 let alpha = T::load(alpha_ptr.add_offsets(n_offs), Some(mask), Some(zeros), &[], None, None, None, false);
193
194 let half = T::full(&[BLOCK_N], D::from_f64(0.5));
195 let eps = T::full(&[BLOCK_N], D::from_f64(1e-7));
196 let ones = T::full(&[BLOCK_N], D::from_f64(1.0));
197 let two = T::full(&[BLOCK_N], D::from_f64(2.0));
198
199 let px1 = px - pw * half;
201 let px2 = px + pw * half;
202 let py1 = py - ph * half;
203 let py2 = py + ph * half;
204 let tx1 = tx - tw * half;
205 let tx2 = tx + tw * half;
206 let ty1 = ty - th * half;
207 let ty2 = ty + th * half;
208
209 let ix1 = T::maximum(px1, tx1);
212 let ix2 = T::minimum(px2, tx2);
213 let iy1 = T::maximum(py1, ty1);
214 let iy2 = T::minimum(py2, ty2);
215 let inter_w = T::maximum(ix2 - ix1, zeros);
216 let inter_h = T::maximum(iy2 - iy1, zeros);
217 let inter = inter_w * inter_h;
218 let union = pw * ph + tw * th - inter;
219
220 let union_eps = union + eps;
221
222 let iw_pos = T::gt(inter_w, zeros);
224 let px2_wins = T::lt(px2, tx2); let px1_loses = T::gt(px1, tx1); let py2_wins = T::lt(py2, ty2);
227 let py1_loses = T::gt(py1, ty1);
228 let ih_pos = T::gt(inter_h, zeros);
229
230 let diw_dpx = T::where_(iw_pos,
232 T::where_(px2_wins, ones, zeros) - T::where_(px1_loses, ones, zeros), zeros);
233 let dih_dpy = T::where_(ih_pos,
234 T::where_(py2_wins, ones, zeros) - T::where_(py1_loses, ones, zeros), zeros);
235 let diw_dpw = T::where_(iw_pos,
236 half * (T::where_(px2_wins, ones, zeros) + T::where_(px1_loses, ones, zeros)), zeros);
237 let dih_dph = T::where_(ih_pos,
238 half * (T::where_(py2_wins, ones, zeros) + T::where_(py1_loses, ones, zeros)), zeros);
239
240 let di_dpx = inter_h * diw_dpx;
241 let di_dpy = inter_w * dih_dpy;
242 let di_dpw = inter_h * diw_dpw;
243 let di_dph = inter_w * dih_dph;
244
245 let du_dpx = zeros - di_dpx;
247 let du_dpy = zeros - di_dpy;
248 let du_dpw = ph - di_dpw;
249 let du_dph = pw - di_dph;
250
251 let diou_dpx = (di_dpx - iou * du_dpx) / union_eps;
262 let diou_dpy = (di_dpy - iou * du_dpy) / union_eps;
263 let diou_dpw = (di_dpw - iou * du_dpw) / union_eps;
264 let diou_dph = (di_dph - iou * du_dph) / union_eps;
265
266 let ex1 = T::minimum(px1, tx1);
269 let ex2 = T::maximum(px2, tx2);
270 let ey1 = T::minimum(py1, ty1);
271 let ey2 = T::maximum(py2, ty2);
272 let ecw = ex2 - ex1;
273 let ech = ey2 - ey1;
274 let c2 = ecw * ecw + ech * ech;
275
276 let dcx = px - tx;
278 let dcy = py - ty;
279 let d2 = dcx * dcx + dcy * dcy;
280
281 let c2_eps = c2 + eps;
282
283 let ex2_from_pred = T::gt(px2, tx2); let ex1_from_pred = T::lt(px1, tx1); let ey2_from_pred = T::gt(py2, ty2);
287 let ey1_from_pred = T::lt(py1, ty1);
288
289 let dc2_dpx = two * ecw * (T::where_(ex2_from_pred, ones, zeros) - T::where_(ex1_from_pred, ones, zeros));
293 let dc2_dpy = two * ech * (T::where_(ey2_from_pred, ones, zeros) - T::where_(ey1_from_pred, ones, zeros));
294 let dc2_dpw = ecw * (T::where_(ex2_from_pred, ones, zeros) + T::where_(ex1_from_pred, ones, zeros));
295 let dc2_dph = ech * (T::where_(ey2_from_pred, ones, zeros) + T::where_(ey1_from_pred, ones, zeros));
296
297 let inv_c2 = ones / c2_eps;
299 let d2_over_c2sq = d2 / (c2_eps * c2_eps);
300 let d_dist_dpx = two * dcx * inv_c2 - d2_over_c2sq * dc2_dpx;
301 let d_dist_dpy = two * dcy * inv_c2 - d2_over_c2sq * dc2_dpy;
302 let d_dist_dpw = zeros - d2_over_c2sq * dc2_dpw;
303 let d_dist_dph = zeros - d2_over_c2sq * dc2_dph;
304
305 let atan_p = T::atan(pw / (ph + eps));
323 let atan_t = T::atan(tw / (th + eps));
324 let diff = atan_t - atan_p;
325
326 let ph_eps = ph + eps;
327 let denom_uv = ph_eps * ph_eps + pw * pw; let eight_pi2_inv = T::full(&[BLOCK_N], D::from_f64(0.810_569_46));
331
332 let dv_dpw = zeros - eight_pi2_inv * diff * ph_eps / denom_uv;
333 let dv_dph = eight_pi2_inv * diff * pw / denom_uv;
334
335 let d_denom = ones - iou + v + eps;
337 let one_minus_iou_eps = ones - iou + eps;
339 let factor = alpha + v * one_minus_iou_eps / (d_denom * d_denom);
340
341 let d_av_dpw = dv_dpw * factor;
342 let d_av_dph = dv_dph * factor;
343
344 let iou_scale = v * v / (d_denom * d_denom) - ones;
347
348 let g_px = dy * (iou_scale * diou_dpx + d_dist_dpx);
351 let g_py = dy * (iou_scale * diou_dpy + d_dist_dpy);
352 let g_pw = dy * (iou_scale * diou_dpw + d_dist_dpw + d_av_dpw);
353 let g_ph = dy * (iou_scale * diou_dph + d_dist_dph + d_av_dph);
354
355 T::store(d_pred_ptr.add_offsets(n_offs + 0 * N), g_px, Some(mask), &[], None, None);
356 T::store(d_pred_ptr.add_offsets(n_offs + 1 * N), g_py, Some(mask), &[], None, None);
357 T::store(d_pred_ptr.add_offsets(n_offs + 2 * N), g_pw, Some(mask), &[], None, None);
358 T::store(d_pred_ptr.add_offsets(n_offs + 3 * N), g_ph, Some(mask), &[], None, None);
359}