Skip to main content

vision_rs/models/yolo/kernels/loss/
ciou.rs

1//! Fused CIoU loss Triton kernels for YOLO detection training.
2//!
3//! Computes Complete-IoU (CIoU) loss between predicted and target boxes,
4//! both in XYWH world-coordinate format, in a single fused kernel pass.
5//!
6//! CIoU = 1 − IoU + d²/c² + α·v
7//!
8//! where:
9//!   - IoU  = intersection / union
10//!   - d²   = squared Euclidean distance between box centers
11//!   - c²   = squared diagonal of the smallest enclosing box
12//!   - v    = (4/π²) · (atan(tw/th) − atan(pw/ph))²  (aspect-ratio consistency)
13//!   - α    = v / (1 − IoU + v + ε)
14//!
15//! Layout:
16//!   - `pred`:   `[4, N]` — predicted (cx, cy, w, h) per anchor
17//!   - `target`: `[4, N]` — target (cx, cy, w, h) per anchor
18//!   - `loss`:   `[N]`    — per-anchor CIoU loss
19//!   - `iou`:    `[N]`    — saved IoU (forward → backward)
20//!   - `v`:      `[N]`    — saved aspect-ratio term (forward → backward)
21//!   - `alpha`:  `[N]`    — saved α coefficient (forward → backward)
22//!
23//! Parallelism: one CTA per BLOCK_N-wide anchor tile.
24//! Grid: `cdiv(N, BLOCK_N)` flat CTAs.
25
26#![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/// Fused CIoU loss forward: predicted and target XYWH boxes → per-anchor loss
36/// plus saved activations (iou, v, alpha) needed by the backward pass.
37///
38/// Grid: `cdiv(N, BLOCK_N)` — one CTA per anchor tile.
39#[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    // Pred XYWH — layout [4, N]: channel c is at base offset c*N.
59    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    // Target XYWH.
65    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    // Pred corners.
75    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    // Target corners.
81    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    // Intersection.
87    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    // Union.
96    let pred_area   = pw * ph;
97    let target_area = tw * th;
98    let union = pred_area + target_area - inter;
99
100    // IoU.
101    let iou = inter / (union + eps);
102
103    // Center distance squared.
104    let dx = px - tx;
105    let dy = py - ty;
106    let d2 = dx * dx + dy * dy;
107
108    // Smallest enclosing box diagonal squared.
109    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    // Aspect-ratio consistency term (requires atan).
118    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    // 4 / π²  ≈ 0.405284734
122    let pi2_inv4 = T::full(&[BLOCK_N], D::from_f64(0.405_284_73));
123    let v = pi2_inv4 * diff * diff;
124
125    // α = v / (1 − IoU + v + ε)
126    let alpha = v / (ones - iou + v + eps);
127
128    // Save activations for the backward pass.
129    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    // CIoU loss = 1 − IoU + d²/(c² + ε) + α·v
134    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/// Fused CIoU loss backward: computes ∂L/∂pred given the upstream gradient
140/// and the saved activations from the forward pass.
141///
142/// Only produces gradients w.r.t. `pred`; target boxes are treated as constants.
143///
144/// Gradient decomposes into three independent parts:
145///
146///   (a) IoU term:            ∂(−IoU)/∂pred
147///   (b) Center-distance term: ∂(d²/(c²+ε))/∂pred
148///   (c) α·v term:            ∂(α·v)/∂(pw,ph)  — zero for (px,py)
149///
150/// The min/max branching for intersection and enclosing-box corners is
151/// re-derived from (pred, target) rather than saved, since pred+target
152/// fully determine which branch was active.
153///
154/// Grid: `cdiv(N, BLOCK_N)` — one CTA per anchor tile.
155#[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    // Reload pred XYWH.
178    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    // Reload target XYWH.
184    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    // Saved activations.
190    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    // Re-derive corners.
200    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    // ── Intersection geometry ────────────────────────────────────────────────
210
211    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    // Boolean masks for min/max branch routing (1.0 if that branch was active).
223    let iw_pos    = T::gt(inter_w, zeros);
224    let px2_wins  = T::lt(px2, tx2);   // ix2 = px2 (pred right edge tighter)
225    let px1_loses = T::gt(px1, tx1);   // ix1 = px1 (pred left  edge tighter)
226    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    // ∂inter_w/∂px = ±1 gated by iw_pos; ∂inter_w/∂pw = ½·(wins+loses) gated by iw_pos
231    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    // ∂union/∂pred: union = pw·ph + tw·th − inter
246    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    // ∂IoU/∂pred_i = (∂inter/∂pred_i − iou·∂union/∂pred_i) / (union+ε)
252    //
253    // The total IoU contribution to ∂loss/∂pred_i also includes a cross-term
254    // from ∂α/∂IoU (since α = v/D, D = 1−IoU+v+ε):
255    //
256    //   ∂loss/∂pred_i|_{IoU} = (v²/D² − 1) · ∂IoU/∂pred_i
257    //
258    // The (−1) is the direct ∂(−IoU)/∂IoU = −1 term.
259    // The (+v²/D²) is the cross-term from ∂(α·v)/∂IoU = v²/D².
260    // D is computed below alongside the α·v section; we use d_denom for it.
261    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    // ── Enclosing-box diagonal squared ───────────────────────────────────────
267
268    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    // Center distance squared.
277    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    // Boolean masks for enclosing-box corner routing.
284    let ex2_from_pred = T::gt(px2, tx2);    // ex2 = px2
285    let ex1_from_pred = T::lt(px1, tx1);    // ex1 = px1
286    let ey2_from_pred = T::gt(py2, ty2);
287    let ey1_from_pred = T::lt(py1, ty1);
288
289    // ∂c²/∂px  = 2·ecw·((ex2_from_pred?1:0) − (ex1_from_pred?1:0))
290    // ∂c²/∂pw  = 2·ecw·(½·ex2_from_pred + ½·ex1_from_pred)  = ecw·(ex2+ex1 flags)
291    //   [px2 = px+pw/2 → ∂px2/∂pw = ½;  px1 = px−pw/2 → ∂ex1/∂pw = −∂px1/∂pw = +½]
292    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    // ∂(d²/(c²+ε))/∂px = 2·(px−tx)/(c²+ε) − d²/(c²+ε)²·∂c²/∂px
298    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    // ── α·v term ─────────────────────────────────────────────────────────────
306    // Only pw and ph carry gradients (atan terms depend on w/h only).
307    //
308    // v = (4/π²)·(atan(tw/(th+ε)) − atan(pw/(ph+ε)))²
309    // Let Δ = atan_t − atan_p,  u_p = pw/(ph+ε)
310    // ∂v/∂pw = (4/π²)·2Δ·(−1/(1+u_p²))·(1/(ph+ε))
311    //        = −(8/π²)·Δ·(ph+ε)/((ph+ε)²+pw²)
312    // ∂v/∂ph = (4/π²)·2Δ·(−1/(1+u_p²))·(−pw/(ph+ε)²)
313    //        = +(8/π²)·Δ·pw/((ph+ε)²+pw²)
314    //
315    // ∂(α·v)/∂pw = ∂v/∂pw · [alpha + v·(1−iou+ε)/(1−iou+v+ε)²]
316    //            = ∂v/∂pw · [alpha + v·(1−iou+ε)/((1−iou+v+ε)·(1−iou+v+ε))]
317    // but  alpha = v/(1−iou+v+ε), so (1−iou+v+ε) = v/alpha  (when alpha≠0)
318    // Simpler: let D = 1−iou+v+ε
319    //   ∂(α·v)/∂pw = ∂v/∂pw · (alpha + v·(D−v)/D²)
320    //              = ∂v/∂pw · (alpha + v·(1−iou+ε)/D²)
321
322    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;        // (ph+ε)² + pw²
328
329    // 8/π²  ≈ 0.810569466
330    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    // D = (1 − iou + v + ε)
336    let d_denom = ones - iou + v + eps;
337    // ∂(α·v)/∂pw = ∂v/∂pw · (alpha + v·(1−iou+ε)/D²)
338    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    // ── IoU scale: combines direct ∂(−IoU) and the cross-term ∂(α·v)/∂IoU ────
345    // Full IoU contribution: (v²/D² − 1) · ∂IoU/∂pred_i
346    let iou_scale = v * v / (d_denom * d_denom) - ones;
347
348    // ── Combine and write ─────────────────────────────────────────────────────
349
350    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}