Skip to main content

teeny_kernels/nn/loss/
elementwise.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// ── L1Loss ────────────────────────────────────────────────────────────────────
26
27/// Element-wise L1 (MAE) loss forward: `out = |x - y|`.
28///
29/// Grid: `[ceil(n / BLOCK_SIZE), 1, 1]`, block `[128, 1, 1]`.
30#[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/// Element-wise L1 backward: `dx = dy * sign(x - y)`.
80///
81/// `sign(0) = 0` by convention (no gradient at the kink).
82#[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// ── MSELoss ───────────────────────────────────────────────────────────────────
149
150/// Element-wise MSE loss forward: `out = (x - y)^2`.
151#[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/// Element-wise MSE backward: `dx = 2 * (x - y) * dy`.
202#[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// ── HuberLoss ─────────────────────────────────────────────────────────────────
264
265/// Element-wise Huber loss forward.
266///
267/// ```text
268/// out = 0.5 * diff^2              if |diff| <= delta
269///     = delta * (|diff| - 0.5*delta)   otherwise
270/// ```
271#[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    // quadratic: 0.5 * diff^2
316    let quad = half * diff * diff;
317    // linear: delta * (|diff| - 0.5 * delta)
318    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/// Element-wise Huber loss backward.
333///
334/// ```text
335/// dx = diff * dy            if |diff| <= delta
336///    = delta * sign(diff) * dy    otherwise
337/// ```
338#[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    // quadratic gradient: diff; linear gradient: delta * sign(diff)
400    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// ── SmoothL1Loss ──────────────────────────────────────────────────────────────
413
414/// Element-wise SmoothL1 (Huber variant) forward.
415///
416/// PyTorch convention with `beta`:
417/// ```text
418/// out = 0.5 * diff^2 / beta     if |diff| < beta
419///     = |diff| - 0.5 * beta     otherwise
420/// ```
421#[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    // quadratic: 0.5 * diff^2 / beta
466    let quad = half * diff * diff / beta_t;
467    // linear: |diff| - 0.5 * beta
468    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/// Element-wise SmoothL1 backward.
483///
484/// ```text
485/// dx = diff / beta * dy       if |diff| < beta
486///    = sign(diff) * dy        otherwise
487/// ```
488#[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    // quadratic gradient: diff / beta; linear gradient: sign(diff)
550    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}