1#![allow(non_snake_case)]
21
22use teeny_macros::kernel;
23use teeny_triton::triton::{
24 types::{AddOffsets, Comparison},
25 *,
26};
27
28#[kernel]
32pub fn swish_forward<T: Triton, const BLOCK_SIZE: i32>(
33 x_ptr: T::Pointer<f32>,
34 y_ptr: T::Pointer<f32>,
35 n_elements: i32,
36) where
37 T::I32Tensor: types::Tensor<i32, 1>,
38 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
39 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
40{
41 let pid = T::program_id(Axis::X);
42 let block_start = pid * BLOCK_SIZE;
43 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
44 let in_bounds = offsets.lt(n_elements);
45 let x = T::load(
46 x_ptr.add_offsets(offsets),
47 Some(in_bounds),
48 None,
49 &[],
50 None,
51 None,
52 None,
53 false,
54 );
55 let one = T::full::<f32>(&[BLOCK_SIZE], 1.0_f32);
56 let neg1 = T::full::<f32>(&[BLOCK_SIZE], -1.0_f32);
57 let sig = one / (one + T::exp(neg1 * x));
58 let y = x * sig;
59 T::store(
60 y_ptr.add_offsets(offsets),
61 y,
62 Some(in_bounds),
63 &[],
64 None,
65 None,
66 );
67}
68
69#[kernel]
72pub fn swish_backward<T: Triton, const BLOCK_SIZE: i32>(
73 dy_ptr: T::Pointer<f32>,
74 x_ptr: T::Pointer<f32>,
75 dx_ptr: T::Pointer<f32>,
76 n_elements: i32,
77) where
78 T::I32Tensor: types::Tensor<i32, 1>,
79 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
80 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
81{
82 let pid = T::program_id(Axis::X);
83 let block_start = pid * BLOCK_SIZE;
84 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
85 let in_bounds = offsets.lt(n_elements);
86 let dy = T::load(
87 dy_ptr.add_offsets(offsets),
88 Some(in_bounds),
89 None,
90 &[],
91 None,
92 None,
93 None,
94 false,
95 );
96 let x = T::load(
97 x_ptr.add_offsets(offsets),
98 Some(in_bounds),
99 None,
100 &[],
101 None,
102 None,
103 None,
104 false,
105 );
106 let one = T::full::<f32>(&[BLOCK_SIZE], 1.0_f32);
107 let neg1 = T::full::<f32>(&[BLOCK_SIZE], -1.0_f32);
108 let sig = one / (one + T::exp(neg1 * x));
109 let dx = (sig + x * sig * (one - sig)) * dy;
110 T::store(
111 dx_ptr.add_offsets(offsets),
112 dx,
113 Some(in_bounds),
114 &[],
115 None,
116 None,
117 );
118}
119
120impl teeny_core::model::RuntimeOp for SwishForward {
121 fn n_activation_inputs(&self) -> usize {
122 1
123 }
124 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
125 vec![]
126 }
127 fn pack_args(
128 &self,
129 inputs: &[(teeny_core::model::RawPtr, &[usize])],
130 _: &[teeny_core::model::RawPtr],
131 output: teeny_core::model::RawPtr,
132 output_shape: &[usize],
133 _: i32,
134 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
135 ) {
136 let n: usize = output_shape.iter().product();
137 visitor.visit_ptr(inputs[0].0);
138 visitor.visit_ptr(output);
139 visitor.visit_i32(n as i32);
140 }
141 fn block(&self) -> [u32; 3] {
142 [self.block_size as u32, 1, 1]
143 }
144 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
145 let n: usize = output_shape.iter().product();
146 [n.div_ceil(self.block_size as usize) as u32, 1, 1]
147 }
148 #[cfg(feature = "training")]
149 fn has_backward(&self) -> bool {
150 true
151 }
152 #[cfg(feature = "training")]
153 fn pack_backward_args(
154 &self,
155 inputs: &[(teeny_core::model::RawPtr, &[usize])],
156 _: &[teeny_core::model::RawPtr],
157 _: teeny_core::model::RawPtr,
158 output_shape: &[usize],
159 grad_output: teeny_core::model::RawPtr,
160 _: i32,
161 grad_inputs: &[teeny_core::model::RawPtr],
162 _: &[teeny_core::model::RawPtr],
163 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
164 ) {
165 let n: usize = output_shape.iter().product();
166 visitor.visit_ptr(grad_output);
167 visitor.visit_ptr(inputs[0].0);
168 visitor.visit_ptr(grad_inputs[0]);
169 visitor.visit_i32(n as i32);
170 }
171 #[cfg(feature = "training")]
172 fn backward_block(&self) -> [u32; 3] {
173 [self.block_size as u32, 1, 1]
174 }
175 #[cfg(feature = "training")]
176 fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
177 let n: usize = output_shape.iter().product();
178 [n.div_ceil(self.block_size as usize) as u32, 1, 1]
179 }
180}
181
182#[kernel]
187pub fn prelu_forward<T: Triton, const BLOCK_SIZE: i32>(
188 x_ptr: T::Pointer<f32>,
189 slope_ptr: T::Pointer<f32>,
190 y_ptr: T::Pointer<f32>,
191 n_elements: i32,
192) where
193 T::I32Tensor: types::Tensor<i32, 1>,
194 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
195 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
196{
197 let pid = T::program_id(Axis::X);
198 let block_start = pid * BLOCK_SIZE;
199 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
200 let in_bounds = offsets.lt(n_elements);
201 let x = T::load(
202 x_ptr.add_offsets(offsets),
203 Some(in_bounds),
204 None,
205 &[],
206 None,
207 None,
208 None,
209 false,
210 );
211 let slope = T::load(
212 slope_ptr.add_offsets(offsets),
213 Some(in_bounds),
214 None,
215 &[],
216 None,
217 None,
218 None,
219 false,
220 );
221 let zero = T::zeros_like(x);
222 let pos = T::maximum(x, zero);
223 let neg = slope * T::minimum(x, zero);
224 let y = pos + neg;
225 T::store(
226 y_ptr.add_offsets(offsets),
227 y,
228 Some(in_bounds),
229 &[],
230 None,
231 None,
232 );
233}
234
235#[kernel]
238pub fn prelu_backward<T: Triton, const BLOCK_SIZE: i32>(
239 dy_ptr: T::Pointer<f32>,
240 x_ptr: T::Pointer<f32>,
241 slope_ptr: T::Pointer<f32>,
242 dx_ptr: T::Pointer<f32>,
243 dslope_ptr: T::Pointer<f32>,
244 n_elements: i32,
245) where
246 T::I32Tensor: types::Tensor<i32, 1>,
247 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
248 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
249{
250 let pid = T::program_id(Axis::X);
251 let block_start = pid * BLOCK_SIZE;
252 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
253 let in_bounds = offsets.lt(n_elements);
254 let dy = T::load(
255 dy_ptr.add_offsets(offsets),
256 Some(in_bounds),
257 None,
258 &[],
259 None,
260 None,
261 None,
262 false,
263 );
264 let x = T::load(
265 x_ptr.add_offsets(offsets),
266 Some(in_bounds),
267 None,
268 &[],
269 None,
270 None,
271 None,
272 false,
273 );
274 let slope = T::load(
275 slope_ptr.add_offsets(offsets),
276 Some(in_bounds),
277 None,
278 &[],
279 None,
280 None,
281 None,
282 false,
283 );
284 let zero = T::zeros_like(x);
285 let x_pos = T::ge(x, zero);
286 let dx = T::where_(x_pos, dy, slope * dy);
287 let dslope = T::where_(x_pos, zero, x * dy);
288 T::store(
289 dx_ptr.add_offsets(offsets),
290 dx,
291 Some(in_bounds),
292 &[],
293 None,
294 None,
295 );
296 T::store(
297 dslope_ptr.add_offsets(offsets),
298 dslope,
299 Some(in_bounds),
300 &[],
301 None,
302 None,
303 );
304}
305
306impl teeny_core::model::RuntimeOp for PreluForward {
307 fn n_activation_inputs(&self) -> usize {
308 2
309 }
310 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
311 vec![]
312 }
313 fn pack_args(
314 &self,
315 inputs: &[(teeny_core::model::RawPtr, &[usize])],
316 _: &[teeny_core::model::RawPtr],
317 output: teeny_core::model::RawPtr,
318 output_shape: &[usize],
319 _: i32,
320 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
321 ) {
322 let n: usize = output_shape.iter().product();
323 visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(inputs[1].0); visitor.visit_ptr(output);
326 visitor.visit_i32(n as i32);
327 }
328 fn block(&self) -> [u32; 3] {
329 [self.block_size as u32, 1, 1]
330 }
331 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
332 let n: usize = output_shape.iter().product();
333 [n.div_ceil(self.block_size as usize) as u32, 1, 1]
334 }
335 #[cfg(feature = "training")]
336 fn has_backward(&self) -> bool {
337 true
338 }
339 #[cfg(feature = "training")]
340 fn pack_backward_args(
341 &self,
342 inputs: &[(teeny_core::model::RawPtr, &[usize])],
343 _: &[teeny_core::model::RawPtr],
344 _: teeny_core::model::RawPtr,
345 output_shape: &[usize],
346 grad_output: teeny_core::model::RawPtr,
347 _: i32,
348 grad_inputs: &[teeny_core::model::RawPtr],
349 _: &[teeny_core::model::RawPtr],
350 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
351 ) {
352 let n: usize = output_shape.iter().product();
353 visitor.visit_ptr(grad_output);
354 visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(inputs[1].0); visitor.visit_ptr(grad_inputs[0]); visitor.visit_ptr(grad_inputs[1]); visitor.visit_i32(n as i32);
359 }
360 #[cfg(feature = "training")]
361 fn backward_block(&self) -> [u32; 3] {
362 [self.block_size as u32, 1, 1]
363 }
364 #[cfg(feature = "training")]
365 fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
366 let n: usize = output_shape.iter().product();
367 [n.div_ceil(self.block_size as usize) as u32, 1, 1]
368 }
369}
370
371#[kernel]
375pub fn thresholded_relu_forward<T: Triton, const BLOCK_SIZE: i32>(
376 x_ptr: T::Pointer<f32>,
377 y_ptr: T::Pointer<f32>,
378 n_elements: i32,
379 alpha: f32,
380) where
381 T::I32Tensor: types::Tensor<i32, 1>,
382 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
383 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
384{
385 let pid = T::program_id(Axis::X);
386 let block_start = pid * BLOCK_SIZE;
387 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
388 let in_bounds = offsets.lt(n_elements);
389 let x = T::load(
390 x_ptr.add_offsets(offsets),
391 Some(in_bounds),
392 None,
393 &[],
394 None,
395 None,
396 None,
397 false,
398 );
399 let alpha_t = T::full::<f32>(&[BLOCK_SIZE], alpha);
400 let above = T::gt(x, alpha_t);
401 let y = T::where_(above, x, T::zeros_like(x));
402 T::store(
403 y_ptr.add_offsets(offsets),
404 y,
405 Some(in_bounds),
406 &[],
407 None,
408 None,
409 );
410}
411
412#[kernel]
414pub fn thresholded_relu_backward<T: Triton, const BLOCK_SIZE: i32>(
415 dy_ptr: T::Pointer<f32>,
416 x_ptr: T::Pointer<f32>,
417 dx_ptr: T::Pointer<f32>,
418 n_elements: i32,
419 alpha: f32,
420) where
421 T::I32Tensor: types::Tensor<i32, 1>,
422 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
423 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
424{
425 let pid = T::program_id(Axis::X);
426 let block_start = pid * BLOCK_SIZE;
427 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
428 let in_bounds = offsets.lt(n_elements);
429 let dy = T::load(
430 dy_ptr.add_offsets(offsets),
431 Some(in_bounds),
432 None,
433 &[],
434 None,
435 None,
436 None,
437 false,
438 );
439 let x = T::load(
440 x_ptr.add_offsets(offsets),
441 Some(in_bounds),
442 None,
443 &[],
444 None,
445 None,
446 None,
447 false,
448 );
449 let alpha_t = T::full::<f32>(&[BLOCK_SIZE], alpha);
450 let above = T::gt(x, alpha_t);
451 let dx = T::where_(above, dy, T::zeros_like(dy));
452 T::store(
453 dx_ptr.add_offsets(offsets),
454 dx,
455 Some(in_bounds),
456 &[],
457 None,
458 None,
459 );
460}
461
462pub struct ThresholdedReluRuntimeOp {
464 pub kernel: ThresholdedReluForward,
465 pub backward_kernel: ThresholdedReluBackward,
466 pub alpha: f32,
467}
468
469impl ThresholdedReluRuntimeOp {
470 pub fn new(block_size: i32, alpha: f32) -> Self {
471 Self {
472 kernel: ThresholdedReluForward::new(block_size),
473 backward_kernel: ThresholdedReluBackward::new(block_size),
474 alpha,
475 }
476 }
477 pub fn forward_source(&self) -> &str {
478 &self.kernel.source
479 }
480 pub fn backward_source(&self) -> &str {
481 &self.backward_kernel.source
482 }
483 pub fn kernel_name(&self) -> &str {
484 self.kernel.name
485 }
486}
487
488impl teeny_core::model::RuntimeOp for ThresholdedReluRuntimeOp {
489 fn n_activation_inputs(&self) -> usize {
490 1
491 }
492 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
493 vec![]
494 }
495 fn pack_args(
496 &self,
497 inputs: &[(teeny_core::model::RawPtr, &[usize])],
498 _: &[teeny_core::model::RawPtr],
499 output: teeny_core::model::RawPtr,
500 output_shape: &[usize],
501 _: i32,
502 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
503 ) {
504 let n: usize = output_shape.iter().product();
505 visitor.visit_ptr(inputs[0].0);
506 visitor.visit_ptr(output);
507 visitor.visit_i32(n as i32);
508 visitor.visit_f32(self.alpha);
509 }
510 fn block(&self) -> [u32; 3] {
511 [self.kernel.block_size as u32, 1, 1]
512 }
513 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
514 let n: usize = output_shape.iter().product();
515 [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
516 }
517 #[cfg(feature = "training")]
518 fn has_backward(&self) -> bool {
519 true
520 }
521 #[cfg(feature = "training")]
522 fn pack_backward_args(
523 &self,
524 inputs: &[(teeny_core::model::RawPtr, &[usize])],
525 _: &[teeny_core::model::RawPtr],
526 _: teeny_core::model::RawPtr,
527 output_shape: &[usize],
528 grad_output: teeny_core::model::RawPtr,
529 _: i32,
530 grad_inputs: &[teeny_core::model::RawPtr],
531 _: &[teeny_core::model::RawPtr],
532 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
533 ) {
534 let n: usize = output_shape.iter().product();
535 visitor.visit_ptr(grad_output);
536 visitor.visit_ptr(inputs[0].0);
537 visitor.visit_ptr(grad_inputs[0]);
538 visitor.visit_i32(n as i32);
539 visitor.visit_f32(self.alpha);
540 }
541 #[cfg(feature = "training")]
542 fn backward_block(&self) -> [u32; 3] {
543 [self.kernel.block_size as u32, 1, 1]
544 }
545 #[cfg(feature = "training")]
546 fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
547 let n: usize = output_shape.iter().product();
548 [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
549 }
550}
551
552#[kernel]
556pub fn shrink_forward<T: Triton, const BLOCK_SIZE: i32>(
557 x_ptr: T::Pointer<f32>,
558 y_ptr: T::Pointer<f32>,
559 n_elements: i32,
560 lambd: f32,
561 bias: f32,
562) where
563 T::I32Tensor: types::Tensor<i32, 1>,
564 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
565 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
566{
567 let pid = T::program_id(Axis::X);
568 let block_start = pid * BLOCK_SIZE;
569 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
570 let in_bounds = offsets.lt(n_elements);
571 let x = T::load(
572 x_ptr.add_offsets(offsets),
573 Some(in_bounds),
574 None,
575 &[],
576 None,
577 None,
578 None,
579 false,
580 );
581 let lam = T::full::<f32>(&[BLOCK_SIZE], lambd);
582 let neg_lam = T::full::<f32>(&[BLOCK_SIZE], -lambd);
583 let b = T::full::<f32>(&[BLOCK_SIZE], bias);
584 let x_gt = T::gt(x, lam);
585 let x_lt = T::lt(x, neg_lam);
586 let y_upper = x - b;
587 let y_lower = x + b;
588 let y_mid = T::where_(x_lt, y_lower, T::zeros_like(x));
589 let y = T::where_(x_gt, y_upper, y_mid);
590 T::store(
591 y_ptr.add_offsets(offsets),
592 y,
593 Some(in_bounds),
594 &[],
595 None,
596 None,
597 );
598}
599
600#[kernel]
602pub fn shrink_backward<T: Triton, const BLOCK_SIZE: i32>(
603 dy_ptr: T::Pointer<f32>,
604 x_ptr: T::Pointer<f32>,
605 dx_ptr: T::Pointer<f32>,
606 n_elements: i32,
607 lambd: f32,
608) where
609 T::I32Tensor: types::Tensor<i32, 1>,
610 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
611 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
612{
613 let pid = T::program_id(Axis::X);
614 let block_start = pid * BLOCK_SIZE;
615 let offsets = T::arange(0, BLOCK_SIZE) + block_start;
616 let in_bounds = offsets.lt(n_elements);
617 let dy = T::load(
618 dy_ptr.add_offsets(offsets),
619 Some(in_bounds),
620 None,
621 &[],
622 None,
623 None,
624 None,
625 false,
626 );
627 let x = T::load(
628 x_ptr.add_offsets(offsets),
629 Some(in_bounds),
630 None,
631 &[],
632 None,
633 None,
634 None,
635 false,
636 );
637 let lam = T::full::<f32>(&[BLOCK_SIZE], lambd);
638 let outside = T::gt(T::abs(x), lam);
639 let dx = T::where_(outside, dy, T::zeros_like(dy));
640 T::store(
641 dx_ptr.add_offsets(offsets),
642 dx,
643 Some(in_bounds),
644 &[],
645 None,
646 None,
647 );
648}
649
650pub struct ShrinkRuntimeOp {
652 pub kernel: ShrinkForward,
653 pub backward_kernel: ShrinkBackward,
654 pub lambd: f32,
655 pub bias: f32,
656}
657
658impl ShrinkRuntimeOp {
659 pub fn new(block_size: i32, lambd: f32, bias: f32) -> Self {
660 Self {
661 kernel: ShrinkForward::new(block_size),
662 backward_kernel: ShrinkBackward::new(block_size),
663 lambd,
664 bias,
665 }
666 }
667 pub fn forward_source(&self) -> &str {
668 &self.kernel.source
669 }
670 pub fn backward_source(&self) -> &str {
671 &self.backward_kernel.source
672 }
673 pub fn kernel_name(&self) -> &str {
674 self.kernel.name
675 }
676}
677
678impl teeny_core::model::RuntimeOp for ShrinkRuntimeOp {
679 fn n_activation_inputs(&self) -> usize {
680 1
681 }
682 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
683 vec![]
684 }
685 fn pack_args(
686 &self,
687 inputs: &[(teeny_core::model::RawPtr, &[usize])],
688 _: &[teeny_core::model::RawPtr],
689 output: teeny_core::model::RawPtr,
690 output_shape: &[usize],
691 _: i32,
692 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
693 ) {
694 let n: usize = output_shape.iter().product();
695 visitor.visit_ptr(inputs[0].0);
696 visitor.visit_ptr(output);
697 visitor.visit_i32(n as i32);
698 visitor.visit_f32(self.lambd);
699 visitor.visit_f32(self.bias);
700 }
701 fn block(&self) -> [u32; 3] {
702 [self.kernel.block_size as u32, 1, 1]
703 }
704 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
705 let n: usize = output_shape.iter().product();
706 [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
707 }
708 #[cfg(feature = "training")]
709 fn has_backward(&self) -> bool {
710 true
711 }
712 #[cfg(feature = "training")]
713 fn pack_backward_args(
714 &self,
715 inputs: &[(teeny_core::model::RawPtr, &[usize])],
716 _: &[teeny_core::model::RawPtr],
717 _: teeny_core::model::RawPtr,
718 output_shape: &[usize],
719 grad_output: teeny_core::model::RawPtr,
720 _: i32,
721 grad_inputs: &[teeny_core::model::RawPtr],
722 _: &[teeny_core::model::RawPtr],
723 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
724 ) {
725 let n: usize = output_shape.iter().product();
726 visitor.visit_ptr(grad_output);
727 visitor.visit_ptr(inputs[0].0);
728 visitor.visit_ptr(grad_inputs[0]);
729 visitor.visit_i32(n as i32);
730 visitor.visit_f32(self.lambd);
731 }
732 #[cfg(feature = "training")]
733 fn backward_block(&self) -> [u32; 3] {
734 [self.kernel.block_size as u32, 1, 1]
735 }
736 #[cfg(feature = "training")]
737 fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
738 let n: usize = output_shape.iter().product();
739 [n.div_ceil(self.kernel.block_size as usize) as u32, 1, 1]
740 }
741}
742
743#[kernel]
749pub fn log_softmax_forward<T: Triton, const BLOCK_SIZE: i32>(
750 x_ptr: T::Pointer<f32>,
751 y_ptr: T::Pointer<f32>,
752 _n_rows: i32,
753 n_cols: i32,
754) where
755 T::I32Tensor: types::Tensor<i32, 1>,
756 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
757 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
758{
759 let pid = T::program_id(Axis::X);
760 let row_offset = pid * n_cols;
761 let col_offsets = T::arange(0, BLOCK_SIZE);
762 let offsets = col_offsets + row_offset;
763 let x = T::load(
764 x_ptr.add_offsets(offsets),
765 None,
766 None,
767 &[],
768 None,
769 None,
770 None,
771 false,
772 );
773 let m = T::max(x, Some(0), true); let x_m = x - m;
775 let log_sum = T::log(T::sum(T::exp(x_m), Some(0), true));
776 let y = x_m - log_sum;
777 T::store(y_ptr.add_offsets(offsets), y, None, &[], None, None);
778}
779
780#[kernel]
782pub fn log_softmax_backward<T: Triton, const BLOCK_SIZE: i32>(
783 dy_ptr: T::Pointer<f32>,
784 y_ptr: T::Pointer<f32>,
785 dx_ptr: T::Pointer<f32>,
786 _n_rows: i32,
787 n_cols: i32,
788) where
789 T::I32Tensor: types::Tensor<i32, 1>,
790 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
791 T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
792{
793 let pid = T::program_id(Axis::X);
794 let row_offset = pid * n_cols;
795 let col_offsets = T::arange(0, BLOCK_SIZE);
796 let offsets = col_offsets + row_offset;
797 let dy = T::load(
798 dy_ptr.add_offsets(offsets),
799 None,
800 None,
801 &[],
802 None,
803 None,
804 None,
805 false,
806 );
807 let y = T::load(
808 y_ptr.add_offsets(offsets),
809 None,
810 None,
811 &[],
812 None,
813 None,
814 None,
815 false,
816 );
817 let sm = T::exp(y);
819 let sum_dy = T::sum(dy, Some(0), true);
820 let dx = dy - sm * sum_dy;
821 T::store(dx_ptr.add_offsets(offsets), dx, None, &[], None, None);
822}
823
824impl teeny_core::model::RuntimeOp for LogSoftmaxForward {
825 fn n_activation_inputs(&self) -> usize {
826 1
827 }
828 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
829 vec![]
830 }
831 fn pack_args(
832 &self,
833 inputs: &[(teeny_core::model::RawPtr, &[usize])],
834 _: &[teeny_core::model::RawPtr],
835 output: teeny_core::model::RawPtr,
836 output_shape: &[usize],
837 _: i32,
838 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
839 ) {
840 let n_rows = output_shape.first().copied().unwrap_or(1) as i32;
841 let n_cols = output_shape.last().copied().unwrap_or(1) as i32;
842 visitor.visit_ptr(inputs[0].0);
843 visitor.visit_ptr(output);
844 visitor.visit_i32(n_rows);
845 visitor.visit_i32(n_cols);
846 }
847 fn block(&self) -> [u32; 3] {
848 [self.block_size as u32, 1, 1]
849 }
850 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
851 [output_shape.first().copied().unwrap_or(1) as u32, 1, 1]
852 }
853 #[cfg(feature = "training")]
854 fn has_backward(&self) -> bool {
855 true
856 }
857 #[cfg(feature = "training")]
858 fn pack_backward_args(
859 &self,
860 _: &[(teeny_core::model::RawPtr, &[usize])],
861 _: &[teeny_core::model::RawPtr],
862 output: teeny_core::model::RawPtr,
863 output_shape: &[usize],
864 grad_output: teeny_core::model::RawPtr,
865 _: i32,
866 grad_inputs: &[teeny_core::model::RawPtr],
867 _: &[teeny_core::model::RawPtr],
868 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
869 ) {
870 let n_rows = output_shape.first().copied().unwrap_or(1) as i32;
871 let n_cols = output_shape.last().copied().unwrap_or(1) as i32;
872 visitor.visit_ptr(grad_output);
873 visitor.visit_ptr(output);
874 visitor.visit_ptr(grad_inputs[0]);
875 visitor.visit_i32(n_rows);
876 visitor.visit_i32(n_cols);
877 }
878 #[cfg(feature = "training")]
879 fn backward_block(&self) -> [u32; 3] {
880 [self.block_size as u32, 1, 1]
881 }
882 #[cfg(feature = "training")]
883 fn backward_grid(&self, _: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
884 [output_shape.first().copied().unwrap_or(1) as u32, 1, 1]
885 }
886}