Skip to main content

teeny_kernels/nn/loss/
embedding.rs

1/*
2 * Copyright (c) 2026 Teenygrad.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *   http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21    types::{AddOffsets, Comparison},
22    *,
23};
24
25// ── CosineEmbeddingLoss ───────────────────────────────────────────────────────
26
27/// Cosine embedding loss forward (per-row).
28///
29/// ```text
30/// cos_sim = dot(x1[n], x2[n]) / (||x1[n]|| * ||x2[n]||)
31/// out[n]  = 1 - cos_sim              if y[n] ==  1
32///         = max(0, cos_sim - margin) if y[n] == -1
33/// ```
34///
35/// Grid: `[n_rows, 1, 1]`.  `BLOCK_SIZE` must equal `next_power_of_two(n_dim)`.
36#[kernel]
37pub fn cosine_embedding_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
38    x1_ptr: T::Pointer<f32>,
39    x2_ptr: T::Pointer<f32>,
40    y_ptr: T::Pointer<f32>,
41    out_ptr: T::Pointer<f32>,
42    _n_rows: i32,
43    n_dim: i32,
44    margin: f32,
45) where
46    T::I32Tensor: types::Tensor<i32, 1>,
47    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
48    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
49{
50    let pid = T::program_id(Axis::X);
51    let row_base = pid * n_dim;
52    let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
53    let row_offs: T::I32Tensor = col_offs + row_base;
54    let in_row = col_offs.lt(n_dim);
55    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
56
57    let x1 = T::load(
58        x1_ptr.add_offsets(row_offs),
59        Some(in_row),
60        Some(zeros),
61        &[],
62        None,
63        None,
64        None,
65        false,
66    );
67    let x2 = T::load(
68        x2_ptr.add_offsets(row_offs),
69        Some(in_row),
70        Some(zeros),
71        &[],
72        None,
73        None,
74        None,
75        false,
76    );
77
78    // Reduction scalars (tt.reduce returns f32 scalar in MLIR)
79    let dot_raw = T::sum(x1 * x2, Some(0), true);
80    let sq1_raw = T::sum(x1 * x1, Some(0), true);
81    let sq2_raw = T::sum(x2 * x2, Some(0), true);
82
83    let dot_t = T::zeros::<f32>(&[1]) + dot_raw;
84    let sq1_t = T::zeros::<f32>(&[1]) + sq1_raw;
85    let sq2_t = T::zeros::<f32>(&[1]) + sq2_raw;
86
87    let norm1 = T::sqrt_rn(sq1_t);
88    let norm2 = T::sqrt_rn(sq2_t);
89    let cos_sim = dot_t / (norm1 * norm2);
90
91    // Load y[pid] → tensor<1xf32>
92    let y_off: T::I32Tensor = T::arange(0, 1) + pid;
93    let y: T::Tensor<f32> = T::load(
94        y_ptr.add_offsets(y_off),
95        None,
96        None,
97        &[],
98        None,
99        None,
100        None,
101        false,
102    );
103
104    let zeros1 = T::zeros::<f32>(&[1]);
105    let margin_t = T::full::<f32>(&[1], margin);
106
107    let y_is_pos = T::gt(y, zeros1);
108    let hinge = T::maximum(cos_sim - margin_t, zeros1);
109    let loss = T::where_(y_is_pos, T::full::<f32>(&[1], 1.0_f32) - cos_sim, hinge);
110
111    let out_off: T::I32Tensor = T::arange(0, 1) + pid;
112    T::store(out_ptr.add_offsets(out_off), loss, None, &[], None, None);
113}
114
115/// Cosine embedding loss backward (per-row).
116///
117/// Let `c = cos_sim`, `n1 = ||x1||`, `n2 = ||x2||`, `r1 = 1/n1`, `r2 = 1/n2`.
118/// ```text
119/// dc/dx1[k] = (x2[k]*r2 - c*x1[k]*r1) * r1
120/// dc/dx2[k] = (x1[k]*r1 - c*x2[k]*r2) * r2
121///
122/// coeff = -dy  if y ==  1
123///       =  dy  if y == -1 and cos_sim > margin
124///       =   0  otherwise
125///
126/// dx1 = coeff * dc/dx1,   dx2 = coeff * dc/dx2
127/// ```
128///
129/// Grid: `[n_rows, 1, 1]`.  `BLOCK_SIZE` must equal `next_power_of_two(n_dim)`.
130#[kernel]
131pub fn cosine_embedding_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
132    dy_ptr: T::Pointer<f32>,
133    x1_ptr: T::Pointer<f32>,
134    x2_ptr: T::Pointer<f32>,
135    y_ptr: T::Pointer<f32>,
136    dx1_ptr: T::Pointer<f32>,
137    dx2_ptr: T::Pointer<f32>,
138    _n_rows: i32,
139    n_dim: i32,
140    margin: f32,
141) where
142    T::I32Tensor: types::Tensor<i32, 1>,
143    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
144    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
145{
146    let pid = T::program_id(Axis::X);
147    let row_base = pid * n_dim;
148    let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
149    let row_offs: T::I32Tensor = col_offs + row_base;
150    let in_row = col_offs.lt(n_dim);
151    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
152
153    let x1 = T::load(
154        x1_ptr.add_offsets(row_offs),
155        Some(in_row),
156        Some(zeros),
157        &[],
158        None,
159        None,
160        None,
161        false,
162    );
163    let x2 = T::load(
164        x2_ptr.add_offsets(row_offs),
165        Some(in_row),
166        Some(zeros),
167        &[],
168        None,
169        None,
170        None,
171        false,
172    );
173
174    // Reductions → scalar f32 in MLIR; force to tensor<1xf32>
175    let dot_t = T::zeros::<f32>(&[1]) + T::sum(x1 * x2, Some(0), true);
176    let sq1_t = T::zeros::<f32>(&[1]) + T::sum(x1 * x1, Some(0), true);
177    let sq2_t = T::zeros::<f32>(&[1]) + T::sum(x2 * x2, Some(0), true);
178
179    let one = T::full::<f32>(&[1], 1.0_f32);
180    let inv_norm1 = one / T::sqrt_rn(sq1_t);
181    let inv_norm2 = one / T::sqrt_rn(sq2_t);
182    let cos_sim = dot_t * inv_norm1 * inv_norm2;
183
184    // Load dy, y → tensor<1xf32>
185    let scalar_off: T::I32Tensor = T::arange(0, 1) + pid;
186    let dy: T::Tensor<f32> = T::load(
187        dy_ptr.add_offsets(scalar_off),
188        None,
189        None,
190        &[],
191        None,
192        None,
193        None,
194        false,
195    );
196    let y: T::Tensor<f32> = T::load(
197        y_ptr.add_offsets(scalar_off),
198        None,
199        None,
200        &[],
201        None,
202        None,
203        None,
204        false,
205    );
206
207    let zeros1 = T::zeros::<f32>(&[1]);
208    let margin_t = T::full::<f32>(&[1], margin);
209
210    let y_is_pos = T::gt(y, zeros1);
211    let cos_gt_margin = T::gt(cos_sim, margin_t);
212    let neg_dy = T::full::<f32>(&[1], -1.0_f32) * dy;
213
214    // coeff: tensor<1xf32>
215    let coeff = T::where_(y_is_pos, neg_dy, T::where_(cos_gt_margin, dy, zeros1));
216
217    // Broadcast tensor<1xf32> values to tensor<BLOCK_SIZExf32> before mixing with x1/x2
218    let inv_norm1_b = T::broadcast_to(inv_norm1, &[BLOCK_SIZE]);
219    let inv_norm2_b = T::broadcast_to(inv_norm2, &[BLOCK_SIZE]);
220    let cos_sim_b = T::broadcast_to(cos_sim, &[BLOCK_SIZE]);
221    let coeff_b = T::broadcast_to(coeff, &[BLOCK_SIZE]);
222
223    // Gradient of cos wrt x1 and x2
224    let d_cos_dx1 = (x2 * inv_norm2_b - cos_sim_b * x1 * inv_norm1_b) * inv_norm1_b;
225    let d_cos_dx2 = (x1 * inv_norm1_b - cos_sim_b * x2 * inv_norm2_b) * inv_norm2_b;
226
227    let dx1 = coeff_b * d_cos_dx1;
228    let dx2 = coeff_b * d_cos_dx2;
229
230    T::store(
231        dx1_ptr.add_offsets(row_offs),
232        dx1,
233        Some(in_row),
234        &[],
235        None,
236        None,
237    );
238    T::store(
239        dx2_ptr.add_offsets(row_offs),
240        dx2,
241        Some(in_row),
242        &[],
243        None,
244        None,
245    );
246}
247
248// ── TripletMarginLoss ─────────────────────────────────────────────────────────
249
250/// Triplet margin loss forward (per-row).
251///
252/// `d(a,p) = sqrt(||a-p||^2 + eps)`,  `d(a,n) = sqrt(||a-n||^2 + eps)`
253/// `out[i] = max(0, d(a,p) - d(a,n) + margin)`
254///
255/// Grid: `[n_rows, 1, 1]`.  `BLOCK_SIZE` must equal `next_power_of_two(n_dim)`.
256#[kernel]
257pub fn triplet_margin_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
258    anchor_ptr: T::Pointer<f32>,
259    positive_ptr: T::Pointer<f32>,
260    negative_ptr: T::Pointer<f32>,
261    out_ptr: T::Pointer<f32>,
262    _n_rows: i32,
263    n_dim: i32,
264    margin: f32,
265    eps: f32,
266) where
267    T::I32Tensor: types::Tensor<i32, 1>,
268    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
269    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
270{
271    let pid = T::program_id(Axis::X);
272    let row_base = pid * n_dim;
273    let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
274    let row_offs: T::I32Tensor = col_offs + row_base;
275    let in_row = col_offs.lt(n_dim);
276    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
277
278    let a = T::load(
279        anchor_ptr.add_offsets(row_offs),
280        Some(in_row),
281        Some(zeros),
282        &[],
283        None,
284        None,
285        None,
286        false,
287    );
288    let p = T::load(
289        positive_ptr.add_offsets(row_offs),
290        Some(in_row),
291        Some(zeros),
292        &[],
293        None,
294        None,
295        None,
296        false,
297    );
298    let n = T::load(
299        negative_ptr.add_offsets(row_offs),
300        Some(in_row),
301        Some(zeros),
302        &[],
303        None,
304        None,
305        None,
306        false,
307    );
308
309    let diff_ap = a - p;
310    let diff_an = a - n;
311
312    // sq + eps: scalar + tensor<1xf32> → tensor<1xf32>
313    let eps_t = T::full::<f32>(&[1], eps);
314    let sq_ap = T::sum(diff_ap * diff_ap, Some(0), true) + eps_t;
315    let sq_an = T::sum(diff_an * diff_an, Some(0), true) + eps_t;
316
317    let d_ap = T::sqrt_rn(sq_ap);
318    let d_an = T::sqrt_rn(sq_an);
319
320    let margin_t = T::full::<f32>(&[1], margin);
321    let zeros1 = T::zeros::<f32>(&[1]);
322    let loss = T::maximum(d_ap - d_an + margin_t, zeros1);
323
324    let out_off: T::I32Tensor = T::arange(0, 1) + pid;
325    T::store(out_ptr.add_offsets(out_off), loss, None, &[], None, None);
326}
327
328/// Triplet margin loss backward (per-row).
329///
330/// When active (d(a,p) - d(a,n) + margin > 0):
331/// ```text
332/// da[k] = dy * ((a[k]-p[k])/d(a,p) - (a[k]-n[k])/d(a,n))
333/// dp[k] = dy * (p[k]-a[k])/d(a,p)
334/// dn[k] = dy * (a[k]-n[k])/d(a,n)
335/// ```
336///
337/// Grid: `[n_rows, 1, 1]`.  `BLOCK_SIZE` must equal `next_power_of_two(n_dim)`.
338#[kernel]
339pub fn triplet_margin_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
340    dy_ptr: T::Pointer<f32>,
341    anchor_ptr: T::Pointer<f32>,
342    positive_ptr: T::Pointer<f32>,
343    negative_ptr: T::Pointer<f32>,
344    da_ptr: T::Pointer<f32>,
345    dp_ptr: T::Pointer<f32>,
346    dn_ptr: T::Pointer<f32>,
347    _n_rows: i32,
348    n_dim: i32,
349    margin: f32,
350    eps: f32,
351) where
352    T::I32Tensor: types::Tensor<i32, 1>,
353    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
354    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
355{
356    let pid = T::program_id(Axis::X);
357    let row_base = pid * n_dim;
358    let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
359    let row_offs: T::I32Tensor = col_offs + row_base;
360    let in_row = col_offs.lt(n_dim);
361    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
362
363    let a = T::load(
364        anchor_ptr.add_offsets(row_offs),
365        Some(in_row),
366        Some(zeros),
367        &[],
368        None,
369        None,
370        None,
371        false,
372    );
373    let p = T::load(
374        positive_ptr.add_offsets(row_offs),
375        Some(in_row),
376        Some(zeros),
377        &[],
378        None,
379        None,
380        None,
381        false,
382    );
383    let n = T::load(
384        negative_ptr.add_offsets(row_offs),
385        Some(in_row),
386        Some(zeros),
387        &[],
388        None,
389        None,
390        None,
391        false,
392    );
393
394    let diff_ap = a - p;
395    let diff_an = a - n;
396
397    let eps_t = T::full::<f32>(&[1], eps);
398    let sq_ap = T::sum(diff_ap * diff_ap, Some(0), true) + eps_t;
399    let sq_an = T::sum(diff_an * diff_an, Some(0), true) + eps_t;
400
401    let one = T::full::<f32>(&[1], 1.0_f32);
402    let d_ap = T::sqrt_rn(sq_ap);
403    let d_an = T::sqrt_rn(sq_an);
404    let inv_d_ap = one / d_ap;
405    let inv_d_an = one / d_an;
406
407    let margin_t = T::full::<f32>(&[1], margin);
408    let zeros1 = T::zeros::<f32>(&[1]);
409    // active: gradient flows only when triplet margin is positive
410    let active = T::gt(d_ap - d_an + margin_t, zeros1);
411
412    // Load dy → tensor<1xf32>
413    let scalar_off: T::I32Tensor = T::arange(0, 1) + pid;
414    let dy: T::Tensor<f32> = T::load(
415        dy_ptr.add_offsets(scalar_off),
416        None,
417        None,
418        &[],
419        None,
420        None,
421        None,
422        false,
423    );
424
425    // Effective dy: zero out gradient for inactive triplets
426    let eff_dy = T::where_(active, dy, zeros1);
427    let neg_eff_dy = T::full::<f32>(&[1], -1.0_f32) * eff_dy;
428
429    // Broadcast tensor<1xf32> to tensor<BLOCK_SIZExf32> for element-wise ops
430    let inv_d_ap_b = T::broadcast_to(inv_d_ap, &[BLOCK_SIZE]);
431    let inv_d_an_b = T::broadcast_to(inv_d_an, &[BLOCK_SIZE]);
432    let eff_dy_b = T::broadcast_to(eff_dy, &[BLOCK_SIZE]);
433    let neg_eff_b = T::broadcast_to(neg_eff_dy, &[BLOCK_SIZE]);
434
435    // Unit direction vectors (diff / distance)
436    let unit_ap = diff_ap * inv_d_ap_b;
437    let unit_an = diff_an * inv_d_an_b;
438
439    let da = eff_dy_b * (unit_ap - unit_an);
440    let dp = neg_eff_b * unit_ap;
441    let dn = eff_dy_b * unit_an;
442
443    T::store(
444        da_ptr.add_offsets(row_offs),
445        da,
446        Some(in_row),
447        &[],
448        None,
449        None,
450    );
451    T::store(
452        dp_ptr.add_offsets(row_offs),
453        dp,
454        Some(in_row),
455        &[],
456        None,
457        None,
458    );
459    T::store(
460        dn_ptr.add_offsets(row_offs),
461        dn,
462        Some(in_row),
463        &[],
464        None,
465        None,
466    );
467}