Skip to main content

teeny_kernels/nn/loss/
bce.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// ── BCELoss ───────────────────────────────────────────────────────────────────
26
27/// Element-wise binary cross-entropy forward.
28///
29/// Assumes `input` is already a probability (sigmoid output), i.e. in `(0, 1)`.
30///
31/// ```text
32/// out = -(target * log(input) + (1 - target) * log(1 - input))
33/// ```
34///
35/// Grid: `[ceil(n / BLOCK_SIZE), 1, 1]`, block `[128, 1, 1]`.
36#[kernel]
37pub fn bce_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
38    input_ptr: T::Pointer<f32>,
39    target_ptr: T::Pointer<f32>,
40    out_ptr: T::Pointer<f32>,
41    n_elements: i32,
42) where
43    T::I32Tensor: types::Tensor<i32, 1>,
44    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
45    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
46{
47    let pid = T::program_id(Axis::X);
48    let block_start = pid * BLOCK_SIZE;
49    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
50    let in_bounds = offsets.lt(n_elements);
51    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
52
53    let inp = T::load(
54        input_ptr.add_offsets(offsets),
55        Some(in_bounds),
56        Some(zeros),
57        &[],
58        None,
59        None,
60        None,
61        false,
62    );
63    let tgt = T::load(
64        target_ptr.add_offsets(offsets),
65        Some(in_bounds),
66        Some(zeros),
67        &[],
68        None,
69        None,
70        None,
71        false,
72    );
73
74    let one = T::full(&[BLOCK_SIZE], 1.0_f32);
75    // Clamp to (eps, 1-eps) to avoid log(0)
76    let eps = T::full(&[BLOCK_SIZE], 1e-7_f32);
77    let one_minus_eps = T::full(&[BLOCK_SIZE], 1.0_f32 - 1e-7_f32);
78    let inp_c = T::clamp(inp, eps, one_minus_eps);
79
80    let loss = T::full(&[BLOCK_SIZE], -1.0_f32)
81        * (tgt * T::log(inp_c) + (one - tgt) * T::log(one - inp_c));
82    T::store(
83        out_ptr.add_offsets(offsets),
84        loss,
85        Some(in_bounds),
86        &[],
87        None,
88        None,
89    );
90}
91
92/// Element-wise BCE backward.
93///
94/// ```text
95/// dx = -(target / input - (1 - target) / (1 - input))
96/// ```
97#[kernel]
98pub fn bce_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
99    dy_ptr: T::Pointer<f32>,
100    input_ptr: T::Pointer<f32>,
101    target_ptr: T::Pointer<f32>,
102    dx_ptr: T::Pointer<f32>,
103    n_elements: i32,
104) where
105    T::I32Tensor: types::Tensor<i32, 1>,
106    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
107    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
108{
109    let pid = T::program_id(Axis::X);
110    let block_start = pid * BLOCK_SIZE;
111    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
112    let in_bounds = offsets.lt(n_elements);
113    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
114
115    let dy = T::load(
116        dy_ptr.add_offsets(offsets),
117        Some(in_bounds),
118        Some(zeros),
119        &[],
120        None,
121        None,
122        None,
123        false,
124    );
125    let inp = T::load(
126        input_ptr.add_offsets(offsets),
127        Some(in_bounds),
128        Some(zeros),
129        &[],
130        None,
131        None,
132        None,
133        false,
134    );
135    let tgt = T::load(
136        target_ptr.add_offsets(offsets),
137        Some(in_bounds),
138        Some(zeros),
139        &[],
140        None,
141        None,
142        None,
143        false,
144    );
145
146    let one = T::full(&[BLOCK_SIZE], 1.0_f32);
147    let eps = T::full(&[BLOCK_SIZE], 1e-7_f32);
148    let one_minus_eps = T::full(&[BLOCK_SIZE], 1.0_f32 - 1e-7_f32);
149    let inp_c = T::clamp(inp, eps, one_minus_eps);
150
151    // dx = -(t/x - (1-t)/(1-x))
152    let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
153    let dx_raw = neg_one * (tgt / inp_c - (one - tgt) / (one - inp_c));
154    let dx = dx_raw * dy;
155    T::store(
156        dx_ptr.add_offsets(offsets),
157        dx,
158        Some(in_bounds),
159        &[],
160        None,
161        None,
162    );
163}
164
165// ── BCEWithLogitsLoss ─────────────────────────────────────────────────────────
166
167/// Element-wise BCE-with-logits forward.
168///
169/// Numerically stable implementation:
170/// ```text
171/// out = max(x, 0) - x*t + log(1 + exp(-|x|))
172/// ```
173#[kernel]
174pub fn bce_with_logits_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
175    input_ptr: T::Pointer<f32>,
176    target_ptr: T::Pointer<f32>,
177    out_ptr: T::Pointer<f32>,
178    n_elements: i32,
179) where
180    T::I32Tensor: types::Tensor<i32, 1>,
181    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
182    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
183{
184    let pid = T::program_id(Axis::X);
185    let block_start = pid * BLOCK_SIZE;
186    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
187    let in_bounds = offsets.lt(n_elements);
188    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
189
190    let inp = T::load(
191        input_ptr.add_offsets(offsets),
192        Some(in_bounds),
193        Some(zeros),
194        &[],
195        None,
196        None,
197        None,
198        false,
199    );
200    let tgt = T::load(
201        target_ptr.add_offsets(offsets),
202        Some(in_bounds),
203        Some(zeros),
204        &[],
205        None,
206        None,
207        None,
208        false,
209    );
210
211    let one = T::full(&[BLOCK_SIZE], 1.0_f32);
212    // Numerically stable: max(x,0) - x*t + log(1+exp(-|x|))
213    let relu_x = T::maximum(inp, zeros);
214    let neg_abs_x = T::full(&[BLOCK_SIZE], -1.0_f32) * T::abs(inp);
215    let loss = relu_x - inp * tgt + T::log(one + T::exp(neg_abs_x));
216    T::store(
217        out_ptr.add_offsets(offsets),
218        loss,
219        Some(in_bounds),
220        &[],
221        None,
222        None,
223    );
224}
225
226/// Element-wise BCE-with-logits backward.
227///
228/// `dx = sigmoid(x) - target`
229#[kernel]
230pub fn bce_with_logits_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
231    dy_ptr: T::Pointer<f32>,
232    input_ptr: T::Pointer<f32>,
233    target_ptr: T::Pointer<f32>,
234    dx_ptr: T::Pointer<f32>,
235    n_elements: i32,
236) where
237    T::I32Tensor: types::Tensor<i32, 1>,
238    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
239    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
240{
241    let pid = T::program_id(Axis::X);
242    let block_start = pid * BLOCK_SIZE;
243    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
244    let in_bounds = offsets.lt(n_elements);
245    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
246
247    let dy = T::load(
248        dy_ptr.add_offsets(offsets),
249        Some(in_bounds),
250        Some(zeros),
251        &[],
252        None,
253        None,
254        None,
255        false,
256    );
257    let inp = T::load(
258        input_ptr.add_offsets(offsets),
259        Some(in_bounds),
260        Some(zeros),
261        &[],
262        None,
263        None,
264        None,
265        false,
266    );
267    let tgt = T::load(
268        target_ptr.add_offsets(offsets),
269        Some(in_bounds),
270        Some(zeros),
271        &[],
272        None,
273        None,
274        None,
275        false,
276    );
277
278    let one = T::full(&[BLOCK_SIZE], 1.0_f32);
279    let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
280    // sigmoid(x) = 1 / (1 + exp(-x))
281    let sig = one / (one + T::exp(neg_one * inp));
282    let dx = (sig - tgt) * dy;
283    T::store(
284        dx_ptr.add_offsets(offsets),
285        dx,
286        Some(in_bounds),
287        &[],
288        None,
289        None,
290    );
291}
292
293// ── SoftMarginLoss ────────────────────────────────────────────────────────────
294
295/// Element-wise soft margin loss forward.
296///
297/// ```text
298/// out = log(1 + exp(-target * input))
299/// ```
300///
301/// Numerically stable via: `log(1 + exp(-t*x)) = max(-t*x, 0) + log(1 + exp(-|t*x|))`
302#[kernel]
303pub fn soft_margin_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
304    input_ptr: T::Pointer<f32>,
305    target_ptr: T::Pointer<f32>,
306    out_ptr: T::Pointer<f32>,
307    n_elements: i32,
308) where
309    T::I32Tensor: types::Tensor<i32, 1>,
310    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
311    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
312{
313    let pid = T::program_id(Axis::X);
314    let block_start = pid * BLOCK_SIZE;
315    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
316    let in_bounds = offsets.lt(n_elements);
317    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
318
319    let inp = T::load(
320        input_ptr.add_offsets(offsets),
321        Some(in_bounds),
322        Some(zeros),
323        &[],
324        None,
325        None,
326        None,
327        false,
328    );
329    let tgt = T::load(
330        target_ptr.add_offsets(offsets),
331        Some(in_bounds),
332        Some(zeros),
333        &[],
334        None,
335        None,
336        None,
337        false,
338    );
339
340    let one = T::full(&[BLOCK_SIZE], 1.0_f32);
341    // log(1 + exp(-t*x)) — numerically stable softplus of (-t*x)
342    let neg_tx = T::full(&[BLOCK_SIZE], -1.0_f32) * tgt * inp;
343    let loss = T::maximum(neg_tx, zeros)
344        + T::log(one + T::exp(T::full(&[BLOCK_SIZE], -1.0_f32) * T::abs(tgt * inp)));
345    T::store(
346        out_ptr.add_offsets(offsets),
347        loss,
348        Some(in_bounds),
349        &[],
350        None,
351        None,
352    );
353}
354
355/// Element-wise soft margin loss backward.
356///
357/// ```text
358/// dx = -target * sigmoid(-target * input)
359/// ```
360#[kernel]
361pub fn soft_margin_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
362    dy_ptr: T::Pointer<f32>,
363    input_ptr: T::Pointer<f32>,
364    target_ptr: T::Pointer<f32>,
365    dx_ptr: T::Pointer<f32>,
366    n_elements: i32,
367) where
368    T::I32Tensor: types::Tensor<i32, 1>,
369    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
370    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
371{
372    let pid = T::program_id(Axis::X);
373    let block_start = pid * BLOCK_SIZE;
374    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
375    let in_bounds = offsets.lt(n_elements);
376    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
377
378    let dy = T::load(
379        dy_ptr.add_offsets(offsets),
380        Some(in_bounds),
381        Some(zeros),
382        &[],
383        None,
384        None,
385        None,
386        false,
387    );
388    let inp = T::load(
389        input_ptr.add_offsets(offsets),
390        Some(in_bounds),
391        Some(zeros),
392        &[],
393        None,
394        None,
395        None,
396        false,
397    );
398    let tgt = T::load(
399        target_ptr.add_offsets(offsets),
400        Some(in_bounds),
401        Some(zeros),
402        &[],
403        None,
404        None,
405        None,
406        false,
407    );
408
409    let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
410    let one = T::full(&[BLOCK_SIZE], 1.0_f32);
411    // sigmoid(-t*x) = 1 / (1 + exp(t*x))
412    let neg_tx = neg_one * tgt * inp;
413    let sig_neg_tx = one / (one + T::exp(neg_one * neg_tx));
414    // dx = -t * sigmoid(-t*x) * dy
415    let dx = neg_one * tgt * sig_neg_tx * dy;
416    T::store(
417        dx_ptr.add_offsets(offsets),
418        dx,
419        Some(in_bounds),
420        &[],
421        None,
422        None,
423    );
424}
425
426// ── KLDivLoss ─────────────────────────────────────────────────────────────────
427
428/// Element-wise KL-divergence forward.
429///
430/// PyTorch convention: `input` is log-probability, `target` is probability.
431///
432/// ```text
433/// out = target * (log(target) - input)
434/// ```
435///
436/// Masked: `out = 0` where `target <= 0`.
437#[kernel]
438pub fn kl_div_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
439    input_ptr: T::Pointer<f32>,
440    target_ptr: T::Pointer<f32>,
441    out_ptr: T::Pointer<f32>,
442    n_elements: i32,
443) where
444    T::I32Tensor: types::Tensor<i32, 1>,
445    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
446    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
447{
448    let pid = T::program_id(Axis::X);
449    let block_start = pid * BLOCK_SIZE;
450    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
451    let in_bounds = offsets.lt(n_elements);
452    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
453
454    let inp = T::load(
455        input_ptr.add_offsets(offsets),
456        Some(in_bounds),
457        Some(zeros),
458        &[],
459        None,
460        None,
461        None,
462        false,
463    );
464    let tgt = T::load(
465        target_ptr.add_offsets(offsets),
466        Some(in_bounds),
467        Some(zeros),
468        &[],
469        None,
470        None,
471        None,
472        false,
473    );
474
475    // out = target * (log(target) - input), masked to 0 where target <= 0
476    let eps = T::full(&[BLOCK_SIZE], 1e-10_f32);
477    let tgt_safe = T::maximum(tgt, eps);
478    let loss_raw = tgt * (T::log(tgt_safe) - inp);
479    let positive = T::gt(tgt, zeros);
480    let loss = T::where_(positive, loss_raw, zeros);
481    T::store(
482        out_ptr.add_offsets(offsets),
483        loss,
484        Some(in_bounds),
485        &[],
486        None,
487        None,
488    );
489}
490
491/// Element-wise KL-divergence backward w.r.t. log-input.
492///
493/// `dx = -target`
494#[kernel]
495pub fn kl_div_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
496    dy_ptr: T::Pointer<f32>,
497    target_ptr: T::Pointer<f32>,
498    dx_ptr: T::Pointer<f32>,
499    n_elements: i32,
500) where
501    T::I32Tensor: types::Tensor<i32, 1>,
502    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
503    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
504{
505    let pid = T::program_id(Axis::X);
506    let block_start = pid * BLOCK_SIZE;
507    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
508    let in_bounds = offsets.lt(n_elements);
509    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
510
511    let dy = T::load(
512        dy_ptr.add_offsets(offsets),
513        Some(in_bounds),
514        Some(zeros),
515        &[],
516        None,
517        None,
518        None,
519        false,
520    );
521    let tgt = T::load(
522        target_ptr.add_offsets(offsets),
523        Some(in_bounds),
524        Some(zeros),
525        &[],
526        None,
527        None,
528        None,
529        false,
530    );
531
532    let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
533    let dx = neg_one * tgt * dy;
534    T::store(
535        dx_ptr.add_offsets(offsets),
536        dx,
537        Some(in_bounds),
538        &[],
539        None,
540        None,
541    );
542}
543
544// ── PoissonNLLLoss ────────────────────────────────────────────────────────────
545
546/// Element-wise Poisson NLL loss forward with `log_input=True` (default).
547///
548/// ```text
549/// out = exp(input) - target * input
550/// ```
551#[kernel]
552pub fn poisson_nll_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
553    input_ptr: T::Pointer<f32>,
554    target_ptr: T::Pointer<f32>,
555    out_ptr: T::Pointer<f32>,
556    n_elements: i32,
557) where
558    T::I32Tensor: types::Tensor<i32, 1>,
559    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
560    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
561{
562    let pid = T::program_id(Axis::X);
563    let block_start = pid * BLOCK_SIZE;
564    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
565    let in_bounds = offsets.lt(n_elements);
566    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
567
568    let inp = T::load(
569        input_ptr.add_offsets(offsets),
570        Some(in_bounds),
571        Some(zeros),
572        &[],
573        None,
574        None,
575        None,
576        false,
577    );
578    let tgt = T::load(
579        target_ptr.add_offsets(offsets),
580        Some(in_bounds),
581        Some(zeros),
582        &[],
583        None,
584        None,
585        None,
586        false,
587    );
588
589    // loss = exp(input) - target * input  (log_input mode)
590    let loss = T::exp(inp) - tgt * inp;
591    T::store(
592        out_ptr.add_offsets(offsets),
593        loss,
594        Some(in_bounds),
595        &[],
596        None,
597        None,
598    );
599}
600
601/// Element-wise Poisson NLL loss backward (log_input=True).
602///
603/// `dx = (exp(input) - target) * dy`
604#[kernel]
605pub fn poisson_nll_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
606    dy_ptr: T::Pointer<f32>,
607    input_ptr: T::Pointer<f32>,
608    target_ptr: T::Pointer<f32>,
609    dx_ptr: T::Pointer<f32>,
610    n_elements: i32,
611) where
612    T::I32Tensor: types::Tensor<i32, 1>,
613    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
614    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
615{
616    let pid = T::program_id(Axis::X);
617    let block_start = pid * BLOCK_SIZE;
618    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
619    let in_bounds = offsets.lt(n_elements);
620    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
621
622    let dy = T::load(
623        dy_ptr.add_offsets(offsets),
624        Some(in_bounds),
625        Some(zeros),
626        &[],
627        None,
628        None,
629        None,
630        false,
631    );
632    let inp = T::load(
633        input_ptr.add_offsets(offsets),
634        Some(in_bounds),
635        Some(zeros),
636        &[],
637        None,
638        None,
639        None,
640        false,
641    );
642    let tgt = T::load(
643        target_ptr.add_offsets(offsets),
644        Some(in_bounds),
645        Some(zeros),
646        &[],
647        None,
648        None,
649        None,
650        false,
651    );
652
653    let dx = (T::exp(inp) - tgt) * dy;
654    T::store(
655        dx_ptr.add_offsets(offsets),
656        dx,
657        Some(in_bounds),
658        &[],
659        None,
660        None,
661    );
662}
663
664// ── GaussianNLLLoss ───────────────────────────────────────────────────────────
665
666/// Element-wise Gaussian NLL loss forward.
667///
668/// ```text
669/// out = 0.5 * (log(var) + (input - target)^2 / var)
670/// ```
671///
672/// `var` is clamped to `eps_var` from below for numerical stability.
673#[kernel]
674pub fn gaussian_nll_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
675    input_ptr: T::Pointer<f32>,
676    target_ptr: T::Pointer<f32>,
677    var_ptr: T::Pointer<f32>,
678    out_ptr: T::Pointer<f32>,
679    n_elements: i32,
680    eps_var: f32,
681) where
682    T::I32Tensor: types::Tensor<i32, 1>,
683    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
684    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
685{
686    let pid = T::program_id(Axis::X);
687    let block_start = pid * BLOCK_SIZE;
688    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
689    let in_bounds = offsets.lt(n_elements);
690    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
691
692    let inp = T::load(
693        input_ptr.add_offsets(offsets),
694        Some(in_bounds),
695        Some(zeros),
696        &[],
697        None,
698        None,
699        None,
700        false,
701    );
702    let tgt = T::load(
703        target_ptr.add_offsets(offsets),
704        Some(in_bounds),
705        Some(zeros),
706        &[],
707        None,
708        None,
709        None,
710        false,
711    );
712    let var = T::load(
713        var_ptr.add_offsets(offsets),
714        Some(in_bounds),
715        Some(zeros),
716        &[],
717        None,
718        None,
719        None,
720        false,
721    );
722
723    let eps_t = T::full(&[BLOCK_SIZE], eps_var);
724    let half = T::full(&[BLOCK_SIZE], 0.5_f32);
725    let var_c = T::maximum(var, eps_t);
726    let diff = inp - tgt;
727    let loss = half * (T::log(var_c) + diff * diff / var_c);
728    T::store(
729        out_ptr.add_offsets(offsets),
730        loss,
731        Some(in_bounds),
732        &[],
733        None,
734        None,
735    );
736}
737
738/// Element-wise Gaussian NLL backward w.r.t. input (mean prediction).
739///
740/// `dx = (input - target) / var * dy`
741#[kernel]
742pub fn gaussian_nll_loss_backward_input<T: Triton, const BLOCK_SIZE: i32>(
743    dy_ptr: T::Pointer<f32>,
744    input_ptr: T::Pointer<f32>,
745    target_ptr: T::Pointer<f32>,
746    var_ptr: T::Pointer<f32>,
747    dx_ptr: T::Pointer<f32>,
748    n_elements: i32,
749    eps_var: f32,
750) where
751    T::I32Tensor: types::Tensor<i32, 1>,
752    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
753    T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
754{
755    let pid = T::program_id(Axis::X);
756    let block_start = pid * BLOCK_SIZE;
757    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
758    let in_bounds = offsets.lt(n_elements);
759    let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
760
761    let dy = T::load(
762        dy_ptr.add_offsets(offsets),
763        Some(in_bounds),
764        Some(zeros),
765        &[],
766        None,
767        None,
768        None,
769        false,
770    );
771    let inp = T::load(
772        input_ptr.add_offsets(offsets),
773        Some(in_bounds),
774        Some(zeros),
775        &[],
776        None,
777        None,
778        None,
779        false,
780    );
781    let tgt = T::load(
782        target_ptr.add_offsets(offsets),
783        Some(in_bounds),
784        Some(zeros),
785        &[],
786        None,
787        None,
788        None,
789        false,
790    );
791    let var = T::load(
792        var_ptr.add_offsets(offsets),
793        Some(in_bounds),
794        Some(zeros),
795        &[],
796        None,
797        None,
798        None,
799        false,
800    );
801
802    let eps_t = T::full(&[BLOCK_SIZE], eps_var);
803    let var_c = T::maximum(var, eps_t);
804    let dx = (inp - tgt) / var_c * dy;
805    T::store(
806        dx_ptr.add_offsets(offsets),
807        dx,
808        Some(in_bounds),
809        &[],
810        None,
811        None,
812    );
813}