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 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 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 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#[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 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 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 let coeff = T::where_(y_is_pos, neg_dy, T::where_(cos_gt_margin, dy, zeros1));
216
217 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 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#[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 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#[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 let active = T::gt(d_ap - d_an + margin_t, zeros1);
411
412 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 let eff_dy = T::where_(active, dy, zeros1);
427 let neg_eff_dy = T::full::<f32>(&[1], -1.0_f32) * eff_dy;
428
429 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 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}