Skip to main content

teeny_kernels/nn/tensor/
elemwise_unary.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_core::dtype::{Float, Num};
20use teeny_macros::kernel;
21use teeny_triton::triton::{
22    types::{AddOffsets, Comparison},
23    *,
24};
25
26// ── Helper macro for standard unary Float RuntimeOp ──────────────────────────
27//
28// All float-only unary ops share the same RuntimeOp: 1 input, 1 output, n i32.
29// Backward packs: dy, x, dx, n.
30
31macro_rules! impl_float_unary_runtime_op {
32    ($Fwd:ident) => {
33        impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
34            fn n_activation_inputs(&self) -> usize {
35                1
36            }
37            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
38                vec![]
39            }
40            fn pack_args(
41                &self,
42                inputs: &[(teeny_core::model::RawPtr, &[usize])],
43                _: &[teeny_core::model::RawPtr],
44                output: teeny_core::model::RawPtr,
45                output_shape: &[usize],
46                _: i32,
47                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
48            ) {
49                let n: usize = output_shape.iter().product();
50                visitor.visit_ptr(inputs[0].0);
51                visitor.visit_ptr(output);
52                visitor.visit_i32(n as i32);
53            }
54            fn block(&self) -> [u32; 3] {
55                [self.block_size as u32, 1, 1]
56            }
57            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
58                let n: usize = output_shape.iter().product();
59                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
60            }
61            #[cfg(feature = "training")]
62            fn has_backward(&self) -> bool {
63                true
64            }
65            #[cfg(feature = "training")]
66            fn pack_backward_args(
67                &self,
68                inputs: &[(teeny_core::model::RawPtr, &[usize])],
69                _: &[teeny_core::model::RawPtr],
70                _: teeny_core::model::RawPtr,
71                output_shape: &[usize],
72                grad_output: teeny_core::model::RawPtr,
73                _: i32,
74                grad_inputs: &[teeny_core::model::RawPtr],
75                _: &[teeny_core::model::RawPtr],
76                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
77            ) {
78                let n: usize = output_shape.iter().product();
79                visitor.visit_ptr(grad_output);
80                visitor.visit_ptr(inputs[0].0);
81                visitor.visit_ptr(grad_inputs[0]);
82                visitor.visit_i32(n as i32);
83            }
84            #[cfg(feature = "training")]
85            fn backward_block(&self) -> [u32; 3] {
86                [self.block_size as u32, 1, 1]
87            }
88            #[cfg(feature = "training")]
89            fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
90                let n: usize = output_shape.iter().product();
91                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
92            }
93        }
94    };
95}
96
97macro_rules! impl_float_unary_runtime_op_no_bwd {
98    ($Fwd:ident) => {
99        impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
100            fn n_activation_inputs(&self) -> usize {
101                1
102            }
103            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
104                vec![]
105            }
106            fn pack_args(
107                &self,
108                inputs: &[(teeny_core::model::RawPtr, &[usize])],
109                _: &[teeny_core::model::RawPtr],
110                output: teeny_core::model::RawPtr,
111                output_shape: &[usize],
112                _: i32,
113                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
114            ) {
115                let n: usize = output_shape.iter().product();
116                visitor.visit_ptr(inputs[0].0);
117                visitor.visit_ptr(output);
118                visitor.visit_i32(n as i32);
119            }
120            fn block(&self) -> [u32; 3] {
121                [self.block_size as u32, 1, 1]
122            }
123            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
124                let n: usize = output_shape.iter().product();
125                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
126            }
127        }
128    };
129}
130
131macro_rules! impl_num_unary_runtime_op {
132    ($Fwd:ident) => {
133        impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
134            fn n_activation_inputs(&self) -> usize {
135                1
136            }
137            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
138                vec![]
139            }
140            fn pack_args(
141                &self,
142                inputs: &[(teeny_core::model::RawPtr, &[usize])],
143                _: &[teeny_core::model::RawPtr],
144                output: teeny_core::model::RawPtr,
145                output_shape: &[usize],
146                _: i32,
147                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
148            ) {
149                let n: usize = output_shape.iter().product();
150                visitor.visit_ptr(inputs[0].0);
151                visitor.visit_ptr(output);
152                visitor.visit_i32(n as i32);
153            }
154            fn block(&self) -> [u32; 3] {
155                [self.block_size as u32, 1, 1]
156            }
157            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
158                let n: usize = output_shape.iter().product();
159                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
160            }
161        }
162    };
163}
164
165macro_rules! impl_num_unary_runtime_op_with_bwd {
166    ($Fwd:ident) => {
167        impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
168            fn n_activation_inputs(&self) -> usize {
169                1
170            }
171            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
172                vec![]
173            }
174            fn pack_args(
175                &self,
176                inputs: &[(teeny_core::model::RawPtr, &[usize])],
177                _: &[teeny_core::model::RawPtr],
178                output: teeny_core::model::RawPtr,
179                output_shape: &[usize],
180                _: i32,
181                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
182            ) {
183                let n: usize = output_shape.iter().product();
184                visitor.visit_ptr(inputs[0].0);
185                visitor.visit_ptr(output);
186                visitor.visit_i32(n as i32);
187            }
188            fn block(&self) -> [u32; 3] {
189                [self.block_size as u32, 1, 1]
190            }
191            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
192                let n: usize = output_shape.iter().product();
193                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
194            }
195            #[cfg(feature = "training")]
196            fn has_backward(&self) -> bool {
197                true
198            }
199            #[cfg(feature = "training")]
200            fn pack_backward_args(
201                &self,
202                inputs: &[(teeny_core::model::RawPtr, &[usize])],
203                _: &[teeny_core::model::RawPtr],
204                _: teeny_core::model::RawPtr,
205                output_shape: &[usize],
206                grad_output: teeny_core::model::RawPtr,
207                _: i32,
208                grad_inputs: &[teeny_core::model::RawPtr],
209                _: &[teeny_core::model::RawPtr],
210                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
211            ) {
212                let n: usize = output_shape.iter().product();
213                visitor.visit_ptr(grad_output);
214                visitor.visit_ptr(inputs[0].0);
215                visitor.visit_ptr(grad_inputs[0]);
216                visitor.visit_i32(n as i32);
217            }
218            #[cfg(feature = "training")]
219            fn backward_block(&self) -> [u32; 3] {
220                [self.block_size as u32, 1, 1]
221            }
222            #[cfg(feature = "training")]
223            fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
224                let n: usize = output_shape.iter().product();
225                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
226            }
227        }
228    };
229}
230
231// ── Neg backward RuntimeOp (dy only, no saved input) ─────────────────────────
232macro_rules! impl_num_neg_bwd_runtime_op {
233    ($Fwd:ident) => {
234        impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
235            fn n_activation_inputs(&self) -> usize {
236                1
237            }
238            fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
239                vec![]
240            }
241            fn pack_args(
242                &self,
243                inputs: &[(teeny_core::model::RawPtr, &[usize])],
244                _: &[teeny_core::model::RawPtr],
245                output: teeny_core::model::RawPtr,
246                output_shape: &[usize],
247                _: i32,
248                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
249            ) {
250                let n: usize = output_shape.iter().product();
251                visitor.visit_ptr(inputs[0].0);
252                visitor.visit_ptr(output);
253                visitor.visit_i32(n as i32);
254            }
255            fn block(&self) -> [u32; 3] {
256                [self.block_size as u32, 1, 1]
257            }
258            fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
259                let n: usize = output_shape.iter().product();
260                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
261            }
262            #[cfg(feature = "training")]
263            fn has_backward(&self) -> bool {
264                true
265            }
266            #[cfg(feature = "training")]
267            fn pack_backward_args(
268                &self,
269                _: &[(teeny_core::model::RawPtr, &[usize])],
270                _: &[teeny_core::model::RawPtr],
271                _: teeny_core::model::RawPtr,
272                output_shape: &[usize],
273                grad_output: teeny_core::model::RawPtr,
274                _: i32,
275                grad_inputs: &[teeny_core::model::RawPtr],
276                _: &[teeny_core::model::RawPtr],
277                visitor: &mut dyn teeny_core::device::program::ArgVisitor,
278            ) {
279                let n: usize = output_shape.iter().product();
280                visitor.visit_ptr(grad_output);
281                visitor.visit_ptr(grad_inputs[0]);
282                visitor.visit_i32(n as i32);
283            }
284            #[cfg(feature = "training")]
285            fn backward_block(&self) -> [u32; 3] {
286                [self.block_size as u32, 1, 1]
287            }
288            #[cfg(feature = "training")]
289            fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
290                let n: usize = output_shape.iter().product();
291                [n.div_ceil(self.block_size as usize) as u32, 1, 1]
292            }
293        }
294    };
295}
296
297// ── Abs ───────────────────────────────────────────────────────────────────────
298
299/// Forward: y = |x|
300#[kernel]
301pub fn elemwise_abs_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
302    x_ptr: T::Pointer<D>,
303    y_ptr: T::Pointer<D>,
304    n_elements: i32,
305) where
306    T::I32Tensor: types::Tensor<i32, 1>,
307    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
308    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
309{
310    let pid = T::program_id(Axis::X);
311    let block_start = pid * BLOCK_SIZE;
312    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
313    let in_bounds = offsets.lt(n_elements);
314    let x = T::load(
315        x_ptr.add_offsets(offsets),
316        Some(in_bounds),
317        None,
318        &[],
319        None,
320        None,
321        None,
322        false,
323    );
324    let y = T::abs(x);
325    T::store(
326        y_ptr.add_offsets(offsets),
327        y,
328        Some(in_bounds),
329        &[],
330        None,
331        None,
332    );
333}
334
335/// Backward: dx = sign(x) * dy  where sign = 1 if x>0, -1 if x<0, 0 if x==0
336#[kernel]
337pub fn elemwise_abs_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
338    dy_ptr: T::Pointer<D>,
339    x_ptr: T::Pointer<D>,
340    dx_ptr: T::Pointer<D>,
341    n_elements: i32,
342) where
343    T::I32Tensor: types::Tensor<i32, 1>,
344    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
345    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
346{
347    let pid = T::program_id(Axis::X);
348    let block_start = pid * BLOCK_SIZE;
349    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
350    let in_bounds = offsets.lt(n_elements);
351    let dy = T::load(
352        dy_ptr.add_offsets(offsets),
353        Some(in_bounds),
354        None,
355        &[],
356        None,
357        None,
358        None,
359        false,
360    );
361    let x = T::load(
362        x_ptr.add_offsets(offsets),
363        Some(in_bounds),
364        None,
365        &[],
366        None,
367        None,
368        None,
369        false,
370    );
371    let zeros = T::zeros_like(x);
372    let ones = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
373    let neg = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], -1), None, false);
374    let pos_mask = T::gt(x, zeros);
375    let neg_mask = T::lt(x, zeros);
376    let sign = T::where_(pos_mask, ones, T::where_(neg_mask, neg, zeros));
377    let dx = sign * dy;
378    T::store(
379        dx_ptr.add_offsets(offsets),
380        dx,
381        Some(in_bounds),
382        &[],
383        None,
384        None,
385    );
386}
387
388impl_num_unary_runtime_op_with_bwd!(ElemwiseAbsForward);
389
390// ── Neg ───────────────────────────────────────────────────────────────────────
391
392/// Forward: y = -x
393#[kernel]
394pub fn elemwise_neg_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
395    x_ptr: T::Pointer<D>,
396    y_ptr: T::Pointer<D>,
397    n_elements: i32,
398) where
399    T::I32Tensor: types::Tensor<i32, 1>,
400    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
401    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
402{
403    let pid = T::program_id(Axis::X);
404    let block_start = pid * BLOCK_SIZE;
405    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
406    let in_bounds = offsets.lt(n_elements);
407    let x = T::load(
408        x_ptr.add_offsets(offsets),
409        Some(in_bounds),
410        None,
411        &[],
412        None,
413        None,
414        None,
415        false,
416    );
417    let y = -x;
418    T::store(
419        y_ptr.add_offsets(offsets),
420        y,
421        Some(in_bounds),
422        &[],
423        None,
424        None,
425    );
426}
427
428/// Backward: dx = -dy
429#[kernel]
430pub fn elemwise_neg_backward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
431    dy_ptr: T::Pointer<D>,
432    dx_ptr: T::Pointer<D>,
433    n_elements: i32,
434) where
435    T::I32Tensor: types::Tensor<i32, 1>,
436    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
437    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
438{
439    let pid = T::program_id(Axis::X);
440    let block_start = pid * BLOCK_SIZE;
441    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
442    let in_bounds = offsets.lt(n_elements);
443    let dy = T::load(
444        dy_ptr.add_offsets(offsets),
445        Some(in_bounds),
446        None,
447        &[],
448        None,
449        None,
450        None,
451        false,
452    );
453    let dx = -dy;
454    T::store(
455        dx_ptr.add_offsets(offsets),
456        dx,
457        Some(in_bounds),
458        &[],
459        None,
460        None,
461    );
462}
463
464impl_num_neg_bwd_runtime_op!(ElemwiseNegForward);
465
466// ── Sign ──────────────────────────────────────────────────────────────────────
467
468/// Forward: y = 1 if x > 0, -1 if x < 0, 0 if x == 0
469#[kernel]
470pub fn elemwise_sign_forward<T: Triton, D: Num, const BLOCK_SIZE: i32>(
471    x_ptr: T::Pointer<D>,
472    y_ptr: T::Pointer<D>,
473    n_elements: i32,
474) where
475    T::I32Tensor: types::Tensor<i32, 1>,
476    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
477    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
478{
479    let pid = T::program_id(Axis::X);
480    let block_start = pid * BLOCK_SIZE;
481    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
482    let in_bounds = offsets.lt(n_elements);
483    let x = T::load(
484        x_ptr.add_offsets(offsets),
485        Some(in_bounds),
486        None,
487        &[],
488        None,
489        None,
490        None,
491        false,
492    );
493    let zeros = T::zeros_like(x);
494    let ones = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], 1), None, false);
495    let neg = T::cast::<i32, D>(T::full::<i32>(&[BLOCK_SIZE], -1), None, false);
496    let pos_mask = T::gt(x, zeros);
497    let neg_mask = T::lt(x, zeros);
498    let y = T::where_(pos_mask, ones, T::where_(neg_mask, neg, zeros));
499    T::store(
500        y_ptr.add_offsets(offsets),
501        y,
502        Some(in_bounds),
503        &[],
504        None,
505        None,
506    );
507}
508
509impl_num_unary_runtime_op!(ElemwiseSignForward);
510
511// ── IsNaN ─────────────────────────────────────────────────────────────────────
512
513/// Forward: y = 1.0 if x is NaN else 0.0
514#[kernel]
515pub fn elemwise_isnan_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
516    x_ptr: T::Pointer<D>,
517    y_ptr: T::Pointer<D>,
518    n_elements: i32,
519) where
520    T::I32Tensor: types::Tensor<i32, 1>,
521    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
522    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
523{
524    let pid = T::program_id(Axis::X);
525    let block_start = pid * BLOCK_SIZE;
526    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
527    let in_bounds = offsets.lt(n_elements);
528    let x = T::load(
529        x_ptr.add_offsets(offsets),
530        Some(in_bounds),
531        None,
532        &[],
533        None,
534        None,
535        None,
536        false,
537    );
538    // Triton uses ordered comparison: eq(NaN, NaN) = False.
539    // Exploit: where(eq(x, x), 0, 1) yields 1 for NaN, 0 otherwise.
540    let one = T::full::<D>(&[BLOCK_SIZE], D::from_f64(1.0));
541    let zero = T::full::<D>(&[BLOCK_SIZE], D::from_f64(0.0));
542    let is_not_nan = T::eq(x, x); // False for NaN (ordered), True for normal
543    let y = T::where_(is_not_nan, zero, one);
544    T::store(
545        y_ptr.add_offsets(offsets),
546        y,
547        Some(in_bounds),
548        &[],
549        None,
550        None,
551    );
552}
553
554impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for ElemwiseIsnanForward<D> {
555    fn n_activation_inputs(&self) -> usize {
556        1
557    }
558    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
559        vec![]
560    }
561    fn pack_args(
562        &self,
563        inputs: &[(teeny_core::model::RawPtr, &[usize])],
564        _: &[teeny_core::model::RawPtr],
565        output: teeny_core::model::RawPtr,
566        output_shape: &[usize],
567        _: i32,
568        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
569    ) {
570        let n: usize = output_shape.iter().product();
571        visitor.visit_ptr(inputs[0].0);
572        visitor.visit_ptr(output);
573        visitor.visit_i32(n as i32);
574    }
575    fn block(&self) -> [u32; 3] {
576        [self.block_size as u32, 1, 1]
577    }
578    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
579        let n: usize = output_shape.iter().product();
580        [n.div_ceil(self.block_size as usize) as u32, 1, 1]
581    }
582}
583
584// ── Ceil (D: Float) ───────────────────────────────────────────────────────────
585
586/// Forward: y = ceil(x)
587#[kernel]
588pub fn elemwise_ceil_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
589    x_ptr: T::Pointer<D>,
590    y_ptr: T::Pointer<D>,
591    n_elements: i32,
592) where
593    T::I32Tensor: types::Tensor<i32, 1>,
594    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
595    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
596{
597    let pid = T::program_id(Axis::X);
598    let block_start = pid * BLOCK_SIZE;
599    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
600    let in_bounds = offsets.lt(n_elements);
601    let x = T::load(
602        x_ptr.add_offsets(offsets),
603        Some(in_bounds),
604        None,
605        &[],
606        None,
607        None,
608        None,
609        false,
610    );
611    let y = T::ceil(x);
612    T::store(
613        y_ptr.add_offsets(offsets),
614        y,
615        Some(in_bounds),
616        &[],
617        None,
618        None,
619    );
620}
621
622impl_float_unary_runtime_op_no_bwd!(ElemwiseCeilForward);
623
624// ── Floor (D: Float) ──────────────────────────────────────────────────────────
625
626/// Forward: y = floor(x)
627#[kernel]
628pub fn elemwise_floor_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
629    x_ptr: T::Pointer<D>,
630    y_ptr: T::Pointer<D>,
631    n_elements: i32,
632) where
633    T::I32Tensor: types::Tensor<i32, 1>,
634    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
635    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
636{
637    let pid = T::program_id(Axis::X);
638    let block_start = pid * BLOCK_SIZE;
639    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
640    let in_bounds = offsets.lt(n_elements);
641    let x = T::load(
642        x_ptr.add_offsets(offsets),
643        Some(in_bounds),
644        None,
645        &[],
646        None,
647        None,
648        None,
649        false,
650    );
651    let y = T::floor(x);
652    T::store(
653        y_ptr.add_offsets(offsets),
654        y,
655        Some(in_bounds),
656        &[],
657        None,
658        None,
659    );
660}
661
662impl_float_unary_runtime_op_no_bwd!(ElemwiseFloorForward);
663
664// ── Sqrt (D: Float) ───────────────────────────────────────────────────────────
665
666/// Forward: y = sqrt(x)
667#[kernel]
668pub fn elemwise_sqrt_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
669    x_ptr: T::Pointer<D>,
670    y_ptr: T::Pointer<D>,
671    n_elements: i32,
672) where
673    T::I32Tensor: types::Tensor<i32, 1>,
674    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
675    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
676{
677    let pid = T::program_id(Axis::X);
678    let block_start = pid * BLOCK_SIZE;
679    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
680    let in_bounds = offsets.lt(n_elements);
681    let x = T::load(
682        x_ptr.add_offsets(offsets),
683        Some(in_bounds),
684        None,
685        &[],
686        None,
687        None,
688        None,
689        false,
690    );
691    let y = T::sqrt(x);
692    T::store(
693        y_ptr.add_offsets(offsets),
694        y,
695        Some(in_bounds),
696        &[],
697        None,
698        None,
699    );
700}
701
702/// Backward: dx = dy / (2 * sqrt(x))
703#[kernel]
704pub fn elemwise_sqrt_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
705    dy_ptr: T::Pointer<D>,
706    x_ptr: T::Pointer<D>,
707    dx_ptr: T::Pointer<D>,
708    n_elements: i32,
709) where
710    T::I32Tensor: types::Tensor<i32, 1>,
711    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
712    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
713{
714    let pid = T::program_id(Axis::X);
715    let block_start = pid * BLOCK_SIZE;
716    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
717    let in_bounds = offsets.lt(n_elements);
718    let dy = T::load(
719        dy_ptr.add_offsets(offsets),
720        Some(in_bounds),
721        None,
722        &[],
723        None,
724        None,
725        None,
726        false,
727    );
728    let x = T::load(
729        x_ptr.add_offsets(offsets),
730        Some(in_bounds),
731        None,
732        &[],
733        None,
734        None,
735        None,
736        false,
737    );
738    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
739    let dx = dy / (two * T::sqrt(x));
740    T::store(
741        dx_ptr.add_offsets(offsets),
742        dx,
743        Some(in_bounds),
744        &[],
745        None,
746        None,
747    );
748}
749
750impl_float_unary_runtime_op!(ElemwiseSqrtForward);
751
752// ── Reciprocal (D: Float) ─────────────────────────────────────────────────────
753
754/// Forward: y = 1 / x
755#[kernel]
756pub fn elemwise_reciprocal_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
757    x_ptr: T::Pointer<D>,
758    y_ptr: T::Pointer<D>,
759    n_elements: i32,
760) where
761    T::I32Tensor: types::Tensor<i32, 1>,
762    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
763    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
764{
765    let pid = T::program_id(Axis::X);
766    let block_start = pid * BLOCK_SIZE;
767    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
768    let in_bounds = offsets.lt(n_elements);
769    let x = T::load(
770        x_ptr.add_offsets(offsets),
771        Some(in_bounds),
772        None,
773        &[],
774        None,
775        None,
776        None,
777        false,
778    );
779    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
780    let y = one / x;
781    T::store(
782        y_ptr.add_offsets(offsets),
783        y,
784        Some(in_bounds),
785        &[],
786        None,
787        None,
788    );
789}
790
791/// Backward: dx = -dy / x^2
792#[kernel]
793pub fn elemwise_reciprocal_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
794    dy_ptr: T::Pointer<D>,
795    x_ptr: T::Pointer<D>,
796    dx_ptr: T::Pointer<D>,
797    n_elements: i32,
798) where
799    T::I32Tensor: types::Tensor<i32, 1>,
800    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
801    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
802{
803    let pid = T::program_id(Axis::X);
804    let block_start = pid * BLOCK_SIZE;
805    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
806    let in_bounds = offsets.lt(n_elements);
807    let dy = T::load(
808        dy_ptr.add_offsets(offsets),
809        Some(in_bounds),
810        None,
811        &[],
812        None,
813        None,
814        None,
815        false,
816    );
817    let x = T::load(
818        x_ptr.add_offsets(offsets),
819        Some(in_bounds),
820        None,
821        &[],
822        None,
823        None,
824        None,
825        false,
826    );
827    let dx = -(dy / (x * x));
828    T::store(
829        dx_ptr.add_offsets(offsets),
830        dx,
831        Some(in_bounds),
832        &[],
833        None,
834        None,
835    );
836}
837
838impl_float_unary_runtime_op!(ElemwiseReciprocalForward);
839
840// ── Exp (D: Float) ────────────────────────────────────────────────────────────
841
842/// Forward: y = exp(x)
843#[kernel]
844pub fn elemwise_exp_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
845    x_ptr: T::Pointer<D>,
846    y_ptr: T::Pointer<D>,
847    n_elements: i32,
848) where
849    T::I32Tensor: types::Tensor<i32, 1>,
850    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
851    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
852{
853    let pid = T::program_id(Axis::X);
854    let block_start = pid * BLOCK_SIZE;
855    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
856    let in_bounds = offsets.lt(n_elements);
857    let x = T::load(
858        x_ptr.add_offsets(offsets),
859        Some(in_bounds),
860        None,
861        &[],
862        None,
863        None,
864        None,
865        false,
866    );
867    let y = T::exp(x);
868    T::store(
869        y_ptr.add_offsets(offsets),
870        y,
871        Some(in_bounds),
872        &[],
873        None,
874        None,
875    );
876}
877
878/// Backward: dx = exp(x) * dy
879#[kernel]
880pub fn elemwise_exp_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
881    dy_ptr: T::Pointer<D>,
882    x_ptr: T::Pointer<D>,
883    dx_ptr: T::Pointer<D>,
884    n_elements: i32,
885) where
886    T::I32Tensor: types::Tensor<i32, 1>,
887    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
888    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
889{
890    let pid = T::program_id(Axis::X);
891    let block_start = pid * BLOCK_SIZE;
892    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
893    let in_bounds = offsets.lt(n_elements);
894    let dy = T::load(
895        dy_ptr.add_offsets(offsets),
896        Some(in_bounds),
897        None,
898        &[],
899        None,
900        None,
901        None,
902        false,
903    );
904    let x = T::load(
905        x_ptr.add_offsets(offsets),
906        Some(in_bounds),
907        None,
908        &[],
909        None,
910        None,
911        None,
912        false,
913    );
914    let dx = T::exp(x) * dy;
915    T::store(
916        dx_ptr.add_offsets(offsets),
917        dx,
918        Some(in_bounds),
919        &[],
920        None,
921        None,
922    );
923}
924
925impl_float_unary_runtime_op!(ElemwiseExpForward);
926
927// ── Log (D: Float) ────────────────────────────────────────────────────────────
928
929/// Forward: y = log(x)
930#[kernel]
931pub fn elemwise_log_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
932    x_ptr: T::Pointer<D>,
933    y_ptr: T::Pointer<D>,
934    n_elements: i32,
935) where
936    T::I32Tensor: types::Tensor<i32, 1>,
937    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
938    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
939{
940    let pid = T::program_id(Axis::X);
941    let block_start = pid * BLOCK_SIZE;
942    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
943    let in_bounds = offsets.lt(n_elements);
944    let x = T::load(
945        x_ptr.add_offsets(offsets),
946        Some(in_bounds),
947        None,
948        &[],
949        None,
950        None,
951        None,
952        false,
953    );
954    let y = T::log(x);
955    T::store(
956        y_ptr.add_offsets(offsets),
957        y,
958        Some(in_bounds),
959        &[],
960        None,
961        None,
962    );
963}
964
965/// Backward: dx = dy / x
966#[kernel]
967pub fn elemwise_log_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
968    dy_ptr: T::Pointer<D>,
969    x_ptr: T::Pointer<D>,
970    dx_ptr: T::Pointer<D>,
971    n_elements: i32,
972) where
973    T::I32Tensor: types::Tensor<i32, 1>,
974    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
975    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
976{
977    let pid = T::program_id(Axis::X);
978    let block_start = pid * BLOCK_SIZE;
979    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
980    let in_bounds = offsets.lt(n_elements);
981    let dy = T::load(
982        dy_ptr.add_offsets(offsets),
983        Some(in_bounds),
984        None,
985        &[],
986        None,
987        None,
988        None,
989        false,
990    );
991    let x = T::load(
992        x_ptr.add_offsets(offsets),
993        Some(in_bounds),
994        None,
995        &[],
996        None,
997        None,
998        None,
999        false,
1000    );
1001    let dx = dy / x;
1002    T::store(
1003        dx_ptr.add_offsets(offsets),
1004        dx,
1005        Some(in_bounds),
1006        &[],
1007        None,
1008        None,
1009    );
1010}
1011
1012impl_float_unary_runtime_op!(ElemwiseLogForward);
1013
1014// ── Erf (D: Float) ────────────────────────────────────────────────────────────
1015
1016/// Forward: y = erf(x)
1017#[kernel]
1018pub fn elemwise_erf_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1019    x_ptr: T::Pointer<D>,
1020    y_ptr: T::Pointer<D>,
1021    n_elements: i32,
1022) where
1023    T::I32Tensor: types::Tensor<i32, 1>,
1024    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1025    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1026{
1027    let pid = T::program_id(Axis::X);
1028    let block_start = pid * BLOCK_SIZE;
1029    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1030    let in_bounds = offsets.lt(n_elements);
1031    let x = T::load(
1032        x_ptr.add_offsets(offsets),
1033        Some(in_bounds),
1034        None,
1035        &[],
1036        None,
1037        None,
1038        None,
1039        false,
1040    );
1041    let y = T::erf(x);
1042    T::store(
1043        y_ptr.add_offsets(offsets),
1044        y,
1045        Some(in_bounds),
1046        &[],
1047        None,
1048        None,
1049    );
1050}
1051
1052/// Backward: dx = 2/sqrt(pi) * exp(-x^2) * dy  where 2/sqrt(pi) ~= 1.1283791670955126
1053#[kernel]
1054pub fn elemwise_erf_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1055    dy_ptr: T::Pointer<D>,
1056    x_ptr: T::Pointer<D>,
1057    dx_ptr: T::Pointer<D>,
1058    n_elements: i32,
1059) where
1060    T::I32Tensor: types::Tensor<i32, 1>,
1061    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1062    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1063{
1064    let pid = T::program_id(Axis::X);
1065    let block_start = pid * BLOCK_SIZE;
1066    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1067    let in_bounds = offsets.lt(n_elements);
1068    let dy = T::load(
1069        dy_ptr.add_offsets(offsets),
1070        Some(in_bounds),
1071        None,
1072        &[],
1073        None,
1074        None,
1075        None,
1076        false,
1077    );
1078    let x = T::load(
1079        x_ptr.add_offsets(offsets),
1080        Some(in_bounds),
1081        None,
1082        &[],
1083        None,
1084        None,
1085        None,
1086        false,
1087    );
1088    // 2/sqrt(pi) = 1.1283791670955126. A literal (not `f32::consts::FRAC_2_SQRT_PI`) is
1089    // deliberate: this function body is compiled through the `teenyc`/no_core Triton DSL
1090    // frontend, which isn't guaranteed to evaluate arbitrary `std` const paths.
1091    #[allow(clippy::approx_constant)]
1092    let coeff = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.128_379_2_f32), None, false);
1093    let dx = coeff * T::exp(-(x * x)) * dy;
1094    T::store(
1095        dx_ptr.add_offsets(offsets),
1096        dx,
1097        Some(in_bounds),
1098        &[],
1099        None,
1100        None,
1101    );
1102}
1103
1104impl_float_unary_runtime_op!(ElemwiseErfForward);
1105
1106// ── Sin (D: Float) ────────────────────────────────────────────────────────────
1107
1108/// Forward: y = sin(x)
1109#[kernel]
1110pub fn elemwise_sin_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1111    x_ptr: T::Pointer<D>,
1112    y_ptr: T::Pointer<D>,
1113    n_elements: i32,
1114) where
1115    T::I32Tensor: types::Tensor<i32, 1>,
1116    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1117    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1118{
1119    let pid = T::program_id(Axis::X);
1120    let block_start = pid * BLOCK_SIZE;
1121    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1122    let in_bounds = offsets.lt(n_elements);
1123    let x = T::load(
1124        x_ptr.add_offsets(offsets),
1125        Some(in_bounds),
1126        None,
1127        &[],
1128        None,
1129        None,
1130        None,
1131        false,
1132    );
1133    let y = T::sin(x);
1134    T::store(
1135        y_ptr.add_offsets(offsets),
1136        y,
1137        Some(in_bounds),
1138        &[],
1139        None,
1140        None,
1141    );
1142}
1143
1144/// Backward: dx = cos(x) * dy
1145#[kernel]
1146pub fn elemwise_sin_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1147    dy_ptr: T::Pointer<D>,
1148    x_ptr: T::Pointer<D>,
1149    dx_ptr: T::Pointer<D>,
1150    n_elements: i32,
1151) where
1152    T::I32Tensor: types::Tensor<i32, 1>,
1153    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1154    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1155{
1156    let pid = T::program_id(Axis::X);
1157    let block_start = pid * BLOCK_SIZE;
1158    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1159    let in_bounds = offsets.lt(n_elements);
1160    let dy = T::load(
1161        dy_ptr.add_offsets(offsets),
1162        Some(in_bounds),
1163        None,
1164        &[],
1165        None,
1166        None,
1167        None,
1168        false,
1169    );
1170    let x = T::load(
1171        x_ptr.add_offsets(offsets),
1172        Some(in_bounds),
1173        None,
1174        &[],
1175        None,
1176        None,
1177        None,
1178        false,
1179    );
1180    let dx = T::cos(x) * dy;
1181    T::store(
1182        dx_ptr.add_offsets(offsets),
1183        dx,
1184        Some(in_bounds),
1185        &[],
1186        None,
1187        None,
1188    );
1189}
1190
1191impl_float_unary_runtime_op!(ElemwiseSinForward);
1192
1193// ── Cos (D: Float) ────────────────────────────────────────────────────────────
1194
1195/// Forward: y = cos(x)
1196#[kernel]
1197pub fn elemwise_cos_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1198    x_ptr: T::Pointer<D>,
1199    y_ptr: T::Pointer<D>,
1200    n_elements: i32,
1201) where
1202    T::I32Tensor: types::Tensor<i32, 1>,
1203    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1204    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1205{
1206    let pid = T::program_id(Axis::X);
1207    let block_start = pid * BLOCK_SIZE;
1208    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1209    let in_bounds = offsets.lt(n_elements);
1210    let x = T::load(
1211        x_ptr.add_offsets(offsets),
1212        Some(in_bounds),
1213        None,
1214        &[],
1215        None,
1216        None,
1217        None,
1218        false,
1219    );
1220    let y = T::cos(x);
1221    T::store(
1222        y_ptr.add_offsets(offsets),
1223        y,
1224        Some(in_bounds),
1225        &[],
1226        None,
1227        None,
1228    );
1229}
1230
1231/// Backward: dx = -sin(x) * dy
1232#[kernel]
1233pub fn elemwise_cos_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1234    dy_ptr: T::Pointer<D>,
1235    x_ptr: T::Pointer<D>,
1236    dx_ptr: T::Pointer<D>,
1237    n_elements: i32,
1238) where
1239    T::I32Tensor: types::Tensor<i32, 1>,
1240    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1241    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1242{
1243    let pid = T::program_id(Axis::X);
1244    let block_start = pid * BLOCK_SIZE;
1245    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1246    let in_bounds = offsets.lt(n_elements);
1247    let dy = T::load(
1248        dy_ptr.add_offsets(offsets),
1249        Some(in_bounds),
1250        None,
1251        &[],
1252        None,
1253        None,
1254        None,
1255        false,
1256    );
1257    let x = T::load(
1258        x_ptr.add_offsets(offsets),
1259        Some(in_bounds),
1260        None,
1261        &[],
1262        None,
1263        None,
1264        None,
1265        false,
1266    );
1267    let dx = -(T::sin(x) * dy);
1268    T::store(
1269        dx_ptr.add_offsets(offsets),
1270        dx,
1271        Some(in_bounds),
1272        &[],
1273        None,
1274        None,
1275    );
1276}
1277
1278impl_float_unary_runtime_op!(ElemwiseCosForward);
1279
1280// ── Tan (D: Float) ────────────────────────────────────────────────────────────
1281
1282/// Forward: y = tan(x) = sin(x) / cos(x)
1283#[kernel]
1284pub fn elemwise_tan_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1285    x_ptr: T::Pointer<D>,
1286    y_ptr: T::Pointer<D>,
1287    n_elements: i32,
1288) where
1289    T::I32Tensor: types::Tensor<i32, 1>,
1290    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1291    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1292{
1293    let pid = T::program_id(Axis::X);
1294    let block_start = pid * BLOCK_SIZE;
1295    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1296    let in_bounds = offsets.lt(n_elements);
1297    let x = T::load(
1298        x_ptr.add_offsets(offsets),
1299        Some(in_bounds),
1300        None,
1301        &[],
1302        None,
1303        None,
1304        None,
1305        false,
1306    );
1307    let y = T::sin(x) / T::cos(x);
1308    T::store(
1309        y_ptr.add_offsets(offsets),
1310        y,
1311        Some(in_bounds),
1312        &[],
1313        None,
1314        None,
1315    );
1316}
1317
1318/// Backward: dx = (1 + tan^2(x)) * dy
1319#[kernel]
1320pub fn elemwise_tan_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1321    dy_ptr: T::Pointer<D>,
1322    x_ptr: T::Pointer<D>,
1323    dx_ptr: T::Pointer<D>,
1324    n_elements: i32,
1325) where
1326    T::I32Tensor: types::Tensor<i32, 1>,
1327    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1328    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1329{
1330    let pid = T::program_id(Axis::X);
1331    let block_start = pid * BLOCK_SIZE;
1332    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1333    let in_bounds = offsets.lt(n_elements);
1334    let dy = T::load(
1335        dy_ptr.add_offsets(offsets),
1336        Some(in_bounds),
1337        None,
1338        &[],
1339        None,
1340        None,
1341        None,
1342        false,
1343    );
1344    let x = T::load(
1345        x_ptr.add_offsets(offsets),
1346        Some(in_bounds),
1347        None,
1348        &[],
1349        None,
1350        None,
1351        None,
1352        false,
1353    );
1354    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1355    let tan = T::sin(x) / T::cos(x);
1356    let dx = (one + tan * tan) * dy;
1357    T::store(
1358        dx_ptr.add_offsets(offsets),
1359        dx,
1360        Some(in_bounds),
1361        &[],
1362        None,
1363        None,
1364    );
1365}
1366
1367impl_float_unary_runtime_op!(ElemwiseTanForward);
1368
1369// ── Asin (D: Float) ───────────────────────────────────────────────────────────
1370
1371/// Forward: y = asin(x) = atan(x / sqrt(1 - x^2))
1372#[kernel]
1373pub fn elemwise_asin_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1374    x_ptr: T::Pointer<D>,
1375    y_ptr: T::Pointer<D>,
1376    n_elements: i32,
1377) where
1378    T::I32Tensor: types::Tensor<i32, 1>,
1379    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1380    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1381{
1382    let pid = T::program_id(Axis::X);
1383    let block_start = pid * BLOCK_SIZE;
1384    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1385    let in_bounds = offsets.lt(n_elements);
1386    let x = T::load(
1387        x_ptr.add_offsets(offsets),
1388        Some(in_bounds),
1389        None,
1390        &[],
1391        None,
1392        None,
1393        None,
1394        false,
1395    );
1396    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1397    let y = T::atan(x / T::sqrt(one - x * x));
1398    T::store(
1399        y_ptr.add_offsets(offsets),
1400        y,
1401        Some(in_bounds),
1402        &[],
1403        None,
1404        None,
1405    );
1406}
1407
1408/// Backward: dx = dy / sqrt(1 - x^2)
1409#[kernel]
1410pub fn elemwise_asin_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1411    dy_ptr: T::Pointer<D>,
1412    x_ptr: T::Pointer<D>,
1413    dx_ptr: T::Pointer<D>,
1414    n_elements: i32,
1415) where
1416    T::I32Tensor: types::Tensor<i32, 1>,
1417    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1418    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1419{
1420    let pid = T::program_id(Axis::X);
1421    let block_start = pid * BLOCK_SIZE;
1422    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1423    let in_bounds = offsets.lt(n_elements);
1424    let dy = T::load(
1425        dy_ptr.add_offsets(offsets),
1426        Some(in_bounds),
1427        None,
1428        &[],
1429        None,
1430        None,
1431        None,
1432        false,
1433    );
1434    let x = T::load(
1435        x_ptr.add_offsets(offsets),
1436        Some(in_bounds),
1437        None,
1438        &[],
1439        None,
1440        None,
1441        None,
1442        false,
1443    );
1444    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1445    let dx = dy / T::sqrt(one - x * x);
1446    T::store(
1447        dx_ptr.add_offsets(offsets),
1448        dx,
1449        Some(in_bounds),
1450        &[],
1451        None,
1452        None,
1453    );
1454}
1455
1456impl_float_unary_runtime_op!(ElemwiseAsinForward);
1457
1458// ── Acos (D: Float) ───────────────────────────────────────────────────────────
1459
1460/// Forward: y = acos(x) = pi/2 - atan(x / sqrt(1 - x^2))
1461#[kernel]
1462pub fn elemwise_acos_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1463    x_ptr: T::Pointer<D>,
1464    y_ptr: T::Pointer<D>,
1465    n_elements: i32,
1466) where
1467    T::I32Tensor: types::Tensor<i32, 1>,
1468    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1469    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1470{
1471    let pid = T::program_id(Axis::X);
1472    let block_start = pid * BLOCK_SIZE;
1473    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1474    let in_bounds = offsets.lt(n_elements);
1475    let x = T::load(
1476        x_ptr.add_offsets(offsets),
1477        Some(in_bounds),
1478        None,
1479        &[],
1480        None,
1481        None,
1482        None,
1483        false,
1484    );
1485    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1486    let half_pi = T::cast::<f32, D>(
1487        T::full::<f32>(&[BLOCK_SIZE], 1.570_796_4_f32), // π/2
1488        None,
1489        false,
1490    );
1491    let y = half_pi - T::atan(x / T::sqrt(one - x * x));
1492    T::store(
1493        y_ptr.add_offsets(offsets),
1494        y,
1495        Some(in_bounds),
1496        &[],
1497        None,
1498        None,
1499    );
1500}
1501
1502/// Backward: dx = -dy / sqrt(1 - x^2)
1503#[kernel]
1504pub fn elemwise_acos_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1505    dy_ptr: T::Pointer<D>,
1506    x_ptr: T::Pointer<D>,
1507    dx_ptr: T::Pointer<D>,
1508    n_elements: i32,
1509) where
1510    T::I32Tensor: types::Tensor<i32, 1>,
1511    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1512    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1513{
1514    let pid = T::program_id(Axis::X);
1515    let block_start = pid * BLOCK_SIZE;
1516    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1517    let in_bounds = offsets.lt(n_elements);
1518    let dy = T::load(
1519        dy_ptr.add_offsets(offsets),
1520        Some(in_bounds),
1521        None,
1522        &[],
1523        None,
1524        None,
1525        None,
1526        false,
1527    );
1528    let x = T::load(
1529        x_ptr.add_offsets(offsets),
1530        Some(in_bounds),
1531        None,
1532        &[],
1533        None,
1534        None,
1535        None,
1536        false,
1537    );
1538    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1539    let dx = -(dy / T::sqrt(one - x * x));
1540    T::store(
1541        dx_ptr.add_offsets(offsets),
1542        dx,
1543        Some(in_bounds),
1544        &[],
1545        None,
1546        None,
1547    );
1548}
1549
1550impl_float_unary_runtime_op!(ElemwiseAcosForward);
1551
1552// ── Atan (D: Float) ───────────────────────────────────────────────────────────
1553
1554/// Forward: y = atan(x)
1555#[kernel]
1556pub fn elemwise_atan_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1557    x_ptr: T::Pointer<D>,
1558    y_ptr: T::Pointer<D>,
1559    n_elements: i32,
1560) where
1561    T::I32Tensor: types::Tensor<i32, 1>,
1562    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1563    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1564{
1565    let pid = T::program_id(Axis::X);
1566    let block_start = pid * BLOCK_SIZE;
1567    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1568    let in_bounds = offsets.lt(n_elements);
1569    let x = T::load(
1570        x_ptr.add_offsets(offsets),
1571        Some(in_bounds),
1572        None,
1573        &[],
1574        None,
1575        None,
1576        None,
1577        false,
1578    );
1579    let y = T::atan(x);
1580    T::store(
1581        y_ptr.add_offsets(offsets),
1582        y,
1583        Some(in_bounds),
1584        &[],
1585        None,
1586        None,
1587    );
1588}
1589
1590/// Backward: dx = dy / (1 + x^2)
1591#[kernel]
1592pub fn elemwise_atan_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1593    dy_ptr: T::Pointer<D>,
1594    x_ptr: T::Pointer<D>,
1595    dx_ptr: T::Pointer<D>,
1596    n_elements: i32,
1597) where
1598    T::I32Tensor: types::Tensor<i32, 1>,
1599    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1600    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1601{
1602    let pid = T::program_id(Axis::X);
1603    let block_start = pid * BLOCK_SIZE;
1604    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1605    let in_bounds = offsets.lt(n_elements);
1606    let dy = T::load(
1607        dy_ptr.add_offsets(offsets),
1608        Some(in_bounds),
1609        None,
1610        &[],
1611        None,
1612        None,
1613        None,
1614        false,
1615    );
1616    let x = T::load(
1617        x_ptr.add_offsets(offsets),
1618        Some(in_bounds),
1619        None,
1620        &[],
1621        None,
1622        None,
1623        None,
1624        false,
1625    );
1626    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1627    let dx = dy / (one + x * x);
1628    T::store(
1629        dx_ptr.add_offsets(offsets),
1630        dx,
1631        Some(in_bounds),
1632        &[],
1633        None,
1634        None,
1635    );
1636}
1637
1638impl_float_unary_runtime_op!(ElemwiseAtanForward);
1639
1640// ── Sinh (D: Float) ───────────────────────────────────────────────────────────
1641
1642/// Forward: y = sinh(x) = (exp(x) - exp(-x)) / 2
1643#[kernel]
1644pub fn elemwise_sinh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1645    x_ptr: T::Pointer<D>,
1646    y_ptr: T::Pointer<D>,
1647    n_elements: i32,
1648) where
1649    T::I32Tensor: types::Tensor<i32, 1>,
1650    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1651    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1652{
1653    let pid = T::program_id(Axis::X);
1654    let block_start = pid * BLOCK_SIZE;
1655    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1656    let in_bounds = offsets.lt(n_elements);
1657    let x = T::load(
1658        x_ptr.add_offsets(offsets),
1659        Some(in_bounds),
1660        None,
1661        &[],
1662        None,
1663        None,
1664        None,
1665        false,
1666    );
1667    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
1668    let y = (T::exp(x) - T::exp(-x)) / two;
1669    T::store(
1670        y_ptr.add_offsets(offsets),
1671        y,
1672        Some(in_bounds),
1673        &[],
1674        None,
1675        None,
1676    );
1677}
1678
1679/// Backward: dx = cosh(x) * dy
1680#[kernel]
1681pub fn elemwise_sinh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1682    dy_ptr: T::Pointer<D>,
1683    x_ptr: T::Pointer<D>,
1684    dx_ptr: T::Pointer<D>,
1685    n_elements: i32,
1686) where
1687    T::I32Tensor: types::Tensor<i32, 1>,
1688    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1689    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1690{
1691    let pid = T::program_id(Axis::X);
1692    let block_start = pid * BLOCK_SIZE;
1693    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1694    let in_bounds = offsets.lt(n_elements);
1695    let dy = T::load(
1696        dy_ptr.add_offsets(offsets),
1697        Some(in_bounds),
1698        None,
1699        &[],
1700        None,
1701        None,
1702        None,
1703        false,
1704    );
1705    let x = T::load(
1706        x_ptr.add_offsets(offsets),
1707        Some(in_bounds),
1708        None,
1709        &[],
1710        None,
1711        None,
1712        None,
1713        false,
1714    );
1715    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
1716    let cosh_x = (T::exp(x) + T::exp(-x)) / two;
1717    let dx = cosh_x * dy;
1718    T::store(
1719        dx_ptr.add_offsets(offsets),
1720        dx,
1721        Some(in_bounds),
1722        &[],
1723        None,
1724        None,
1725    );
1726}
1727
1728impl_float_unary_runtime_op!(ElemwiseSinhForward);
1729
1730// ── Cosh (D: Float) ───────────────────────────────────────────────────────────
1731
1732/// Forward: y = cosh(x) = (exp(x) + exp(-x)) / 2
1733#[kernel]
1734pub fn elemwise_cosh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1735    x_ptr: T::Pointer<D>,
1736    y_ptr: T::Pointer<D>,
1737    n_elements: i32,
1738) where
1739    T::I32Tensor: types::Tensor<i32, 1>,
1740    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1741    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1742{
1743    let pid = T::program_id(Axis::X);
1744    let block_start = pid * BLOCK_SIZE;
1745    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1746    let in_bounds = offsets.lt(n_elements);
1747    let x = T::load(
1748        x_ptr.add_offsets(offsets),
1749        Some(in_bounds),
1750        None,
1751        &[],
1752        None,
1753        None,
1754        None,
1755        false,
1756    );
1757    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
1758    let y = (T::exp(x) + T::exp(-x)) / two;
1759    T::store(
1760        y_ptr.add_offsets(offsets),
1761        y,
1762        Some(in_bounds),
1763        &[],
1764        None,
1765        None,
1766    );
1767}
1768
1769/// Backward: dx = sinh(x) * dy
1770#[kernel]
1771pub fn elemwise_cosh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1772    dy_ptr: T::Pointer<D>,
1773    x_ptr: T::Pointer<D>,
1774    dx_ptr: T::Pointer<D>,
1775    n_elements: i32,
1776) where
1777    T::I32Tensor: types::Tensor<i32, 1>,
1778    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1779    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1780{
1781    let pid = T::program_id(Axis::X);
1782    let block_start = pid * BLOCK_SIZE;
1783    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1784    let in_bounds = offsets.lt(n_elements);
1785    let dy = T::load(
1786        dy_ptr.add_offsets(offsets),
1787        Some(in_bounds),
1788        None,
1789        &[],
1790        None,
1791        None,
1792        None,
1793        false,
1794    );
1795    let x = T::load(
1796        x_ptr.add_offsets(offsets),
1797        Some(in_bounds),
1798        None,
1799        &[],
1800        None,
1801        None,
1802        None,
1803        false,
1804    );
1805    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
1806    let sinh_x = (T::exp(x) - T::exp(-x)) / two;
1807    let dx = sinh_x * dy;
1808    T::store(
1809        dx_ptr.add_offsets(offsets),
1810        dx,
1811        Some(in_bounds),
1812        &[],
1813        None,
1814        None,
1815    );
1816}
1817
1818impl_float_unary_runtime_op!(ElemwiseCoshForward);
1819
1820// ── Asinh (D: Float) ──────────────────────────────────────────────────────────
1821
1822/// Forward: y = asinh(x) = log(x + sqrt(x^2 + 1))
1823#[kernel]
1824pub fn elemwise_asinh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1825    x_ptr: T::Pointer<D>,
1826    y_ptr: T::Pointer<D>,
1827    n_elements: i32,
1828) where
1829    T::I32Tensor: types::Tensor<i32, 1>,
1830    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1831    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1832{
1833    let pid = T::program_id(Axis::X);
1834    let block_start = pid * BLOCK_SIZE;
1835    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1836    let in_bounds = offsets.lt(n_elements);
1837    let x = T::load(
1838        x_ptr.add_offsets(offsets),
1839        Some(in_bounds),
1840        None,
1841        &[],
1842        None,
1843        None,
1844        None,
1845        false,
1846    );
1847    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1848    let y = T::log(x + T::sqrt(x * x + one));
1849    T::store(
1850        y_ptr.add_offsets(offsets),
1851        y,
1852        Some(in_bounds),
1853        &[],
1854        None,
1855        None,
1856    );
1857}
1858
1859/// Backward: dx = dy / sqrt(x^2 + 1)
1860#[kernel]
1861pub fn elemwise_asinh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1862    dy_ptr: T::Pointer<D>,
1863    x_ptr: T::Pointer<D>,
1864    dx_ptr: T::Pointer<D>,
1865    n_elements: i32,
1866) where
1867    T::I32Tensor: types::Tensor<i32, 1>,
1868    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1869    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1870{
1871    let pid = T::program_id(Axis::X);
1872    let block_start = pid * BLOCK_SIZE;
1873    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1874    let in_bounds = offsets.lt(n_elements);
1875    let dy = T::load(
1876        dy_ptr.add_offsets(offsets),
1877        Some(in_bounds),
1878        None,
1879        &[],
1880        None,
1881        None,
1882        None,
1883        false,
1884    );
1885    let x = T::load(
1886        x_ptr.add_offsets(offsets),
1887        Some(in_bounds),
1888        None,
1889        &[],
1890        None,
1891        None,
1892        None,
1893        false,
1894    );
1895    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1896    let dx = dy / T::sqrt(x * x + one);
1897    T::store(
1898        dx_ptr.add_offsets(offsets),
1899        dx,
1900        Some(in_bounds),
1901        &[],
1902        None,
1903        None,
1904    );
1905}
1906
1907impl_float_unary_runtime_op!(ElemwiseAsinhForward);
1908
1909// ── Acosh (D: Float) ──────────────────────────────────────────────────────────
1910
1911/// Forward: y = acosh(x) = log(x + sqrt(x^2 - 1))
1912#[kernel]
1913pub fn elemwise_acosh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1914    x_ptr: T::Pointer<D>,
1915    y_ptr: T::Pointer<D>,
1916    n_elements: i32,
1917) where
1918    T::I32Tensor: types::Tensor<i32, 1>,
1919    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1920    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1921{
1922    let pid = T::program_id(Axis::X);
1923    let block_start = pid * BLOCK_SIZE;
1924    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1925    let in_bounds = offsets.lt(n_elements);
1926    let x = T::load(
1927        x_ptr.add_offsets(offsets),
1928        Some(in_bounds),
1929        None,
1930        &[],
1931        None,
1932        None,
1933        None,
1934        false,
1935    );
1936    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1937    let y = T::log(x + T::sqrt(x * x - one));
1938    T::store(
1939        y_ptr.add_offsets(offsets),
1940        y,
1941        Some(in_bounds),
1942        &[],
1943        None,
1944        None,
1945    );
1946}
1947
1948/// Backward: dx = dy / sqrt(x^2 - 1)
1949#[kernel]
1950pub fn elemwise_acosh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
1951    dy_ptr: T::Pointer<D>,
1952    x_ptr: T::Pointer<D>,
1953    dx_ptr: T::Pointer<D>,
1954    n_elements: i32,
1955) where
1956    T::I32Tensor: types::Tensor<i32, 1>,
1957    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
1958    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
1959{
1960    let pid = T::program_id(Axis::X);
1961    let block_start = pid * BLOCK_SIZE;
1962    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
1963    let in_bounds = offsets.lt(n_elements);
1964    let dy = T::load(
1965        dy_ptr.add_offsets(offsets),
1966        Some(in_bounds),
1967        None,
1968        &[],
1969        None,
1970        None,
1971        None,
1972        false,
1973    );
1974    let x = T::load(
1975        x_ptr.add_offsets(offsets),
1976        Some(in_bounds),
1977        None,
1978        &[],
1979        None,
1980        None,
1981        None,
1982        false,
1983    );
1984    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
1985    let dx = dy / T::sqrt(x * x - one);
1986    T::store(
1987        dx_ptr.add_offsets(offsets),
1988        dx,
1989        Some(in_bounds),
1990        &[],
1991        None,
1992        None,
1993    );
1994}
1995
1996impl_float_unary_runtime_op!(ElemwiseAcoshForward);
1997
1998// ── Atanh (D: Float) ──────────────────────────────────────────────────────────
1999
2000/// Forward: y = atanh(x) = log((1+x)/(1-x)) / 2
2001#[kernel]
2002pub fn elemwise_atanh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
2003    x_ptr: T::Pointer<D>,
2004    y_ptr: T::Pointer<D>,
2005    n_elements: i32,
2006) where
2007    T::I32Tensor: types::Tensor<i32, 1>,
2008    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
2009    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
2010{
2011    let pid = T::program_id(Axis::X);
2012    let block_start = pid * BLOCK_SIZE;
2013    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
2014    let in_bounds = offsets.lt(n_elements);
2015    let x = T::load(
2016        x_ptr.add_offsets(offsets),
2017        Some(in_bounds),
2018        None,
2019        &[],
2020        None,
2021        None,
2022        None,
2023        false,
2024    );
2025    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
2026    let two = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 2.0_f32), None, false);
2027    let y = T::log((one + x) / (one - x)) / two;
2028    T::store(
2029        y_ptr.add_offsets(offsets),
2030        y,
2031        Some(in_bounds),
2032        &[],
2033        None,
2034        None,
2035    );
2036}
2037
2038/// Backward: dx = dy / (1 - x^2)
2039#[kernel]
2040pub fn elemwise_atanh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
2041    dy_ptr: T::Pointer<D>,
2042    x_ptr: T::Pointer<D>,
2043    dx_ptr: T::Pointer<D>,
2044    n_elements: i32,
2045) where
2046    T::I32Tensor: types::Tensor<i32, 1>,
2047    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
2048    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
2049{
2050    let pid = T::program_id(Axis::X);
2051    let block_start = pid * BLOCK_SIZE;
2052    let offsets = T::arange(0, BLOCK_SIZE) + block_start;
2053    let in_bounds = offsets.lt(n_elements);
2054    let dy = T::load(
2055        dy_ptr.add_offsets(offsets),
2056        Some(in_bounds),
2057        None,
2058        &[],
2059        None,
2060        None,
2061        None,
2062        false,
2063    );
2064    let x = T::load(
2065        x_ptr.add_offsets(offsets),
2066        Some(in_bounds),
2067        None,
2068        &[],
2069        None,
2070        None,
2071        None,
2072        false,
2073    );
2074    let one = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_SIZE], 1.0_f32), None, false);
2075    let dx = dy / (one - x * x);
2076    T::store(
2077        dx_ptr.add_offsets(offsets),
2078        dx,
2079        Some(in_bounds),
2080        &[],
2081        None,
2082        None,
2083    );
2084}
2085
2086impl_float_unary_runtime_op!(ElemwiseAtanhForward);