1#![allow(non_snake_case)]
18
19use teeny_macros::kernel;
20use teeny_triton::triton::{
21 types::{AddOffsets, Comparison},
22 *,
23};
24
25#[kernel]
31pub fn l1_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
32 x_ptr: T::Pointer<f32>,
33 y_ptr: T::Pointer<f32>,
34 out_ptr: T::Pointer<f32>,
35 n_elements: i32,
36) where
37 T::I32Tensor: types::Tensor<i32, 1>,
38 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
39 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
40{
41 let pid = T::program_id(Axis::X);
42 let block_start = pid * BLOCK_SIZE;
43 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
44 let in_bounds = offsets.lt(n_elements);
45 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
46
47 let x = T::load(
48 x_ptr.add_offsets(offsets),
49 Some(in_bounds),
50 Some(zeros),
51 &[],
52 None,
53 None,
54 None,
55 false,
56 );
57 let y = T::load(
58 y_ptr.add_offsets(offsets),
59 Some(in_bounds),
60 Some(zeros),
61 &[],
62 None,
63 None,
64 None,
65 false,
66 );
67
68 let loss = T::abs(x - y);
69 T::store(
70 out_ptr.add_offsets(offsets),
71 loss,
72 Some(in_bounds),
73 &[],
74 None,
75 None,
76 );
77}
78
79#[kernel]
83pub fn l1_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
84 dy_ptr: T::Pointer<f32>,
85 x_ptr: T::Pointer<f32>,
86 y_ptr: T::Pointer<f32>,
87 dx_ptr: T::Pointer<f32>,
88 n_elements: i32,
89) where
90 T::I32Tensor: types::Tensor<i32, 1>,
91 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
92 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
93{
94 let pid = T::program_id(Axis::X);
95 let block_start = pid * BLOCK_SIZE;
96 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
97 let in_bounds = offsets.lt(n_elements);
98 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
99
100 let dy = T::load(
101 dy_ptr.add_offsets(offsets),
102 Some(in_bounds),
103 Some(zeros),
104 &[],
105 None,
106 None,
107 None,
108 false,
109 );
110 let x = T::load(
111 x_ptr.add_offsets(offsets),
112 Some(in_bounds),
113 Some(zeros),
114 &[],
115 None,
116 None,
117 None,
118 false,
119 );
120 let y = T::load(
121 y_ptr.add_offsets(offsets),
122 Some(in_bounds),
123 Some(zeros),
124 &[],
125 None,
126 None,
127 None,
128 false,
129 );
130
131 let diff = x - y;
132 let ones = T::full(&[BLOCK_SIZE], 1.0_f32);
133 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
134 let pos = T::gt(diff, zeros);
135 let neg = T::lt(diff, zeros);
136 let sign = T::where_(pos, ones, T::where_(neg, neg_one, zeros));
137 let dx = dy * sign;
138 T::store(
139 dx_ptr.add_offsets(offsets),
140 dx,
141 Some(in_bounds),
142 &[],
143 None,
144 None,
145 );
146}
147
148#[kernel]
152pub fn mse_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
153 x_ptr: T::Pointer<f32>,
154 y_ptr: T::Pointer<f32>,
155 out_ptr: T::Pointer<f32>,
156 n_elements: i32,
157) where
158 T::I32Tensor: types::Tensor<i32, 1>,
159 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
160 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
161{
162 let pid = T::program_id(Axis::X);
163 let block_start = pid * BLOCK_SIZE;
164 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
165 let in_bounds = offsets.lt(n_elements);
166 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
167
168 let x = T::load(
169 x_ptr.add_offsets(offsets),
170 Some(in_bounds),
171 Some(zeros),
172 &[],
173 None,
174 None,
175 None,
176 false,
177 );
178 let y = T::load(
179 y_ptr.add_offsets(offsets),
180 Some(in_bounds),
181 Some(zeros),
182 &[],
183 None,
184 None,
185 None,
186 false,
187 );
188
189 let diff = x - y;
190 let loss = diff * diff;
191 T::store(
192 out_ptr.add_offsets(offsets),
193 loss,
194 Some(in_bounds),
195 &[],
196 None,
197 None,
198 );
199}
200
201#[kernel]
203pub fn mse_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
204 dy_ptr: T::Pointer<f32>,
205 x_ptr: T::Pointer<f32>,
206 y_ptr: T::Pointer<f32>,
207 dx_ptr: T::Pointer<f32>,
208 n_elements: i32,
209) where
210 T::I32Tensor: types::Tensor<i32, 1>,
211 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
212 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
213{
214 let pid = T::program_id(Axis::X);
215 let block_start = pid * BLOCK_SIZE;
216 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
217 let in_bounds = offsets.lt(n_elements);
218 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
219
220 let dy = T::load(
221 dy_ptr.add_offsets(offsets),
222 Some(in_bounds),
223 Some(zeros),
224 &[],
225 None,
226 None,
227 None,
228 false,
229 );
230 let x = T::load(
231 x_ptr.add_offsets(offsets),
232 Some(in_bounds),
233 Some(zeros),
234 &[],
235 None,
236 None,
237 None,
238 false,
239 );
240 let y = T::load(
241 y_ptr.add_offsets(offsets),
242 Some(in_bounds),
243 Some(zeros),
244 &[],
245 None,
246 None,
247 None,
248 false,
249 );
250
251 let two = T::full(&[BLOCK_SIZE], 2.0_f32);
252 let dx = two * (x - y) * dy;
253 T::store(
254 dx_ptr.add_offsets(offsets),
255 dx,
256 Some(in_bounds),
257 &[],
258 None,
259 None,
260 );
261}
262
263#[kernel]
272pub fn huber_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
273 x_ptr: T::Pointer<f32>,
274 y_ptr: T::Pointer<f32>,
275 out_ptr: T::Pointer<f32>,
276 n_elements: i32,
277 delta: f32,
278) where
279 T::I32Tensor: types::Tensor<i32, 1>,
280 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
281 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
282{
283 let pid = T::program_id(Axis::X);
284 let block_start = pid * BLOCK_SIZE;
285 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
286 let in_bounds = offsets.lt(n_elements);
287 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
288
289 let x = T::load(
290 x_ptr.add_offsets(offsets),
291 Some(in_bounds),
292 Some(zeros),
293 &[],
294 None,
295 None,
296 None,
297 false,
298 );
299 let y = T::load(
300 y_ptr.add_offsets(offsets),
301 Some(in_bounds),
302 Some(zeros),
303 &[],
304 None,
305 None,
306 None,
307 false,
308 );
309
310 let diff = x - y;
311 let abs_diff = T::abs(diff);
312 let delta_t = T::full(&[BLOCK_SIZE], delta);
313 let half = T::full(&[BLOCK_SIZE], 0.5_f32);
314
315 let quad = half * diff * diff;
317 let lin = delta_t * (abs_diff - half * delta_t);
319
320 let in_quad = T::le(abs_diff, delta_t);
321 let loss = T::where_(in_quad, quad, lin);
322 T::store(
323 out_ptr.add_offsets(offsets),
324 loss,
325 Some(in_bounds),
326 &[],
327 None,
328 None,
329 );
330}
331
332#[kernel]
339pub fn huber_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
340 dy_ptr: T::Pointer<f32>,
341 x_ptr: T::Pointer<f32>,
342 y_ptr: T::Pointer<f32>,
343 dx_ptr: T::Pointer<f32>,
344 n_elements: i32,
345 delta: f32,
346) where
347 T::I32Tensor: types::Tensor<i32, 1>,
348 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
349 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
350{
351 let pid = T::program_id(Axis::X);
352 let block_start = pid * BLOCK_SIZE;
353 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
354 let in_bounds = offsets.lt(n_elements);
355 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
356
357 let dy = T::load(
358 dy_ptr.add_offsets(offsets),
359 Some(in_bounds),
360 Some(zeros),
361 &[],
362 None,
363 None,
364 None,
365 false,
366 );
367 let x = T::load(
368 x_ptr.add_offsets(offsets),
369 Some(in_bounds),
370 Some(zeros),
371 &[],
372 None,
373 None,
374 None,
375 false,
376 );
377 let y = T::load(
378 y_ptr.add_offsets(offsets),
379 Some(in_bounds),
380 Some(zeros),
381 &[],
382 None,
383 None,
384 None,
385 false,
386 );
387
388 let diff = x - y;
389 let abs_diff = T::abs(diff);
390 let delta_t = T::full(&[BLOCK_SIZE], delta);
391 let ones = T::full(&[BLOCK_SIZE], 1.0_f32);
392 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
393
394 let pos = T::gt(diff, zeros);
395 let neg = T::lt(diff, zeros);
396 let sign = T::where_(pos, ones, T::where_(neg, neg_one, zeros));
397
398 let in_quad = T::le(abs_diff, delta_t);
399 let grad = T::where_(in_quad, diff, delta_t * sign);
401 let dx = grad * dy;
402 T::store(
403 dx_ptr.add_offsets(offsets),
404 dx,
405 Some(in_bounds),
406 &[],
407 None,
408 None,
409 );
410}
411
412#[kernel]
422pub fn smooth_l1_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
423 x_ptr: T::Pointer<f32>,
424 y_ptr: T::Pointer<f32>,
425 out_ptr: T::Pointer<f32>,
426 n_elements: i32,
427 beta: f32,
428) where
429 T::I32Tensor: types::Tensor<i32, 1>,
430 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
431 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
432{
433 let pid = T::program_id(Axis::X);
434 let block_start = pid * BLOCK_SIZE;
435 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
436 let in_bounds = offsets.lt(n_elements);
437 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
438
439 let x = T::load(
440 x_ptr.add_offsets(offsets),
441 Some(in_bounds),
442 Some(zeros),
443 &[],
444 None,
445 None,
446 None,
447 false,
448 );
449 let y = T::load(
450 y_ptr.add_offsets(offsets),
451 Some(in_bounds),
452 Some(zeros),
453 &[],
454 None,
455 None,
456 None,
457 false,
458 );
459
460 let diff = x - y;
461 let abs_diff = T::abs(diff);
462 let beta_t = T::full(&[BLOCK_SIZE], beta);
463 let half = T::full(&[BLOCK_SIZE], 0.5_f32);
464
465 let quad = half * diff * diff / beta_t;
467 let lin = abs_diff - half * beta_t;
469
470 let in_quad = T::lt(abs_diff, beta_t);
471 let loss = T::where_(in_quad, quad, lin);
472 T::store(
473 out_ptr.add_offsets(offsets),
474 loss,
475 Some(in_bounds),
476 &[],
477 None,
478 None,
479 );
480}
481
482#[kernel]
489pub fn smooth_l1_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
490 dy_ptr: T::Pointer<f32>,
491 x_ptr: T::Pointer<f32>,
492 y_ptr: T::Pointer<f32>,
493 dx_ptr: T::Pointer<f32>,
494 n_elements: i32,
495 beta: f32,
496) where
497 T::I32Tensor: types::Tensor<i32, 1>,
498 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
499 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
500{
501 let pid = T::program_id(Axis::X);
502 let block_start = pid * BLOCK_SIZE;
503 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
504 let in_bounds = offsets.lt(n_elements);
505 let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
506
507 let dy = T::load(
508 dy_ptr.add_offsets(offsets),
509 Some(in_bounds),
510 Some(zeros),
511 &[],
512 None,
513 None,
514 None,
515 false,
516 );
517 let x = T::load(
518 x_ptr.add_offsets(offsets),
519 Some(in_bounds),
520 Some(zeros),
521 &[],
522 None,
523 None,
524 None,
525 false,
526 );
527 let y = T::load(
528 y_ptr.add_offsets(offsets),
529 Some(in_bounds),
530 Some(zeros),
531 &[],
532 None,
533 None,
534 None,
535 false,
536 );
537
538 let diff = x - y;
539 let abs_diff = T::abs(diff);
540 let beta_t = T::full(&[BLOCK_SIZE], beta);
541 let ones = T::full(&[BLOCK_SIZE], 1.0_f32);
542 let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
543
544 let pos = T::gt(diff, zeros);
545 let neg = T::lt(diff, zeros);
546 let sign = T::where_(pos, ones, T::where_(neg, neg_one, zeros));
547
548 let in_quad = T::lt(abs_diff, beta_t);
549 let grad = T::where_(in_quad, diff / beta_t, sign);
551 let dx = grad * dy;
552 T::store(
553 dx_ptr.add_offsets(offsets),
554 dx,
555 Some(in_bounds),
556 &[],
557 None,
558 None,
559 );
560}