1#![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
26macro_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
231macro_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#[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#[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#[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#[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#[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#[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 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); 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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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 #[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#[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#[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#[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#[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#[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#[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#[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#[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#[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), 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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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#[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);