1#![allow(non_snake_case)]
25
26use teeny_core::dtype::{Float, Num};
27use teeny_macros::kernel;
28use teeny_triton::triton::{
29 types::{AddOffsets, Comparison},
30 *,
31};
32
33macro_rules! impl_reduce_num_runtime_op {
38 ($Fwd:ident) => {
39 impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
40 fn n_activation_inputs(&self) -> usize {
41 1
42 }
43 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
44 vec![]
45 }
46 fn pack_args(
47 &self,
48 inputs: &[(teeny_core::model::RawPtr, &[usize])],
49 _: &[teeny_core::model::RawPtr],
50 output: teeny_core::model::RawPtr,
51 output_shape: &[usize],
52 _: i32,
53 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
54 ) {
55 let n_outer: usize = output_shape.iter().product::<usize>().max(1);
59 let n_total: usize = inputs[0].1.iter().product();
60 let n_inner: usize = if n_outer > 0 {
61 n_total / n_outer
62 } else {
63 n_total
64 };
65 visitor.visit_ptr(inputs[0].0);
66 visitor.visit_ptr(output);
67 visitor.visit_i32(n_inner as i32);
68 visitor.visit_i32(n_outer as i32);
69 }
70 fn block(&self) -> [u32; 3] {
71 [self.block_inner as u32, 1, 1]
72 }
73 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
74 let n_outer: usize = output_shape.iter().product::<usize>().max(1);
75 [n_outer as u32, 1, 1]
76 }
77 }
78 };
79}
80
81macro_rules! impl_reduce_float_runtime_op {
82 ($Fwd:ident) => {
83 impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
84 fn n_activation_inputs(&self) -> usize {
85 1
86 }
87 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
88 vec![]
89 }
90 fn pack_args(
91 &self,
92 inputs: &[(teeny_core::model::RawPtr, &[usize])],
93 _: &[teeny_core::model::RawPtr],
94 output: teeny_core::model::RawPtr,
95 output_shape: &[usize],
96 _: i32,
97 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
98 ) {
99 let n_outer: usize = output_shape.iter().product::<usize>().max(1);
100 let n_total: usize = inputs[0].1.iter().product();
101 let n_inner: usize = if n_outer > 0 {
102 n_total / n_outer
103 } else {
104 n_total
105 };
106 visitor.visit_ptr(inputs[0].0);
107 visitor.visit_ptr(output);
108 visitor.visit_i32(n_inner as i32);
109 visitor.visit_i32(n_outer as i32);
110 }
111 fn block(&self) -> [u32; 3] {
112 [self.block_inner as u32, 1, 1]
113 }
114 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
115 let n_outer: usize = output_shape.iter().product::<usize>().max(1);
116 [n_outer as u32, 1, 1]
117 }
118 }
119 };
120}
121
122#[kernel]
127pub fn reduce_sum_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
128 x_ptr: T::Pointer<D>,
129 y_ptr: T::Pointer<D>,
130 n_inner: i32,
131 n_outer: i32,
132) where
133 T::I32Tensor: types::Tensor<i32, 1>,
134 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
135 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
136{
137 let row = T::program_id(Axis::X);
138 if row >= n_outer {
139 return;
140 }
141 let col_offsets = T::arange(0, BLOCK_INNER);
142 let offsets = col_offsets + row * n_inner;
143 let mask = col_offsets.lt(n_inner);
144 let x = T::load(
145 x_ptr.add_offsets(offsets),
146 Some(mask),
147 Some(T::zeros::<D>(&[BLOCK_INNER])),
148 &[],
149 None,
150 None,
151 None,
152 false,
153 );
154 let sum = T::sum(x, Some(0), true); let row_offsets = T::arange(0, 1) + row;
156 T::store(y_ptr.add_offsets(row_offsets), sum, None, &[], None, None);
157}
158
159impl_reduce_num_runtime_op!(ReduceSumForward);
162
163#[kernel]
167pub fn reduce_mean_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
168 x_ptr: T::Pointer<D>,
169 y_ptr: T::Pointer<D>,
170 n_inner: i32,
171 n_outer: i32,
172) where
173 T::I32Tensor: types::Tensor<i32, 1>,
174 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
175 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
176{
177 let row = T::program_id(Axis::X);
178 if row >= n_outer {
179 return;
180 }
181 let col_offsets = T::arange(0, BLOCK_INNER);
182 let offsets = col_offsets + row * n_inner;
183 let mask = col_offsets.lt(n_inner);
184 let x = T::load(
185 x_ptr.add_offsets(offsets),
186 Some(mask),
187 Some(T::zeros::<D>(&[BLOCK_INNER])),
188 &[],
189 None,
190 None,
191 None,
192 false,
193 );
194 let sum = T::sum(x, Some(0), true);
195 let n_f = T::cast::<i32, D>(T::full::<i32>(&[1], n_inner), None, false);
196 let mean = sum / n_f;
197 let row_offsets = T::arange(0, 1) + row;
198 T::store(y_ptr.add_offsets(row_offsets), mean, None, &[], None, None);
199}
200
201impl_reduce_float_runtime_op!(ReduceMeanForward);
202
203#[kernel]
207pub fn reduce_max_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
208 x_ptr: T::Pointer<D>,
209 y_ptr: T::Pointer<D>,
210 n_inner: i32,
211 n_outer: i32,
212) where
213 T::I32Tensor: types::Tensor<i32, 1>,
214 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
215 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
216{
217 let row = T::program_id(Axis::X);
218 if row >= n_outer {
219 return;
220 }
221 let col_offsets = T::arange(0, BLOCK_INNER);
222 let offsets = col_offsets + row * n_inner;
223 let mask = col_offsets.lt(n_inner);
224 let neg_inf = T::cast::<f32, D>(
226 T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
227 None,
228 false,
229 );
230 let x = T::load(
231 x_ptr.add_offsets(offsets),
232 Some(mask),
233 Some(neg_inf),
234 &[],
235 None,
236 None,
237 None,
238 false,
239 );
240 let val = T::max(x, Some(0), true);
241 let row_offsets = T::arange(0, 1) + row;
242 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
243}
244
245impl_reduce_num_runtime_op!(ReduceMaxForward);
246
247#[kernel]
251pub fn reduce_min_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
252 x_ptr: T::Pointer<D>,
253 y_ptr: T::Pointer<D>,
254 n_inner: i32,
255 n_outer: i32,
256) where
257 T::I32Tensor: types::Tensor<i32, 1>,
258 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
259 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
260{
261 let row = T::program_id(Axis::X);
262 if row >= n_outer {
263 return;
264 }
265 let col_offsets = T::arange(0, BLOCK_INNER);
266 let offsets = col_offsets + row * n_inner;
267 let mask = col_offsets.lt(n_inner);
268 let pos_inf = T::cast::<f32, D>(
269 T::full::<f32>(&[BLOCK_INNER], 3.4028235e38_f32),
270 None,
271 false,
272 );
273 let x = T::load(
274 x_ptr.add_offsets(offsets),
275 Some(mask),
276 Some(pos_inf),
277 &[],
278 None,
279 None,
280 None,
281 false,
282 );
283 let val = T::min(x, Some(0), true);
284 let row_offsets = T::arange(0, 1) + row;
285 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
286}
287
288impl_reduce_num_runtime_op!(ReduceMinForward);
289
290#[kernel]
294pub fn reduce_l1_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
295 x_ptr: T::Pointer<D>,
296 y_ptr: T::Pointer<D>,
297 n_inner: i32,
298 n_outer: i32,
299) where
300 T::I32Tensor: types::Tensor<i32, 1>,
301 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
302 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
303{
304 let row = T::program_id(Axis::X);
305 if row >= n_outer {
306 return;
307 }
308 let col_offsets = T::arange(0, BLOCK_INNER);
309 let offsets = col_offsets + row * n_inner;
310 let mask = col_offsets.lt(n_inner);
311 let x = T::load(
312 x_ptr.add_offsets(offsets),
313 Some(mask),
314 Some(T::zeros::<D>(&[BLOCK_INNER])),
315 &[],
316 None,
317 None,
318 None,
319 false,
320 );
321 let val = T::sum(T::abs(x), Some(0), true);
322 let row_offsets = T::arange(0, 1) + row;
323 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
324}
325
326impl_reduce_num_runtime_op!(ReduceL1Forward);
327
328#[kernel]
332pub fn reduce_l2_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
333 x_ptr: T::Pointer<D>,
334 y_ptr: T::Pointer<D>,
335 n_inner: i32,
336 n_outer: i32,
337) where
338 T::I32Tensor: types::Tensor<i32, 1>,
339 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
340 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
341{
342 let row = T::program_id(Axis::X);
343 if row >= n_outer {
344 return;
345 }
346 let col_offsets = T::arange(0, BLOCK_INNER);
347 let offsets = col_offsets + row * n_inner;
348 let mask = col_offsets.lt(n_inner);
349 let x = T::load(
350 x_ptr.add_offsets(offsets),
351 Some(mask),
352 Some(T::zeros::<D>(&[BLOCK_INNER])),
353 &[],
354 None,
355 None,
356 None,
357 false,
358 );
359 let sum_sq = T::sum(x * x, Some(0), true);
360 let val = T::sqrt(sum_sq);
361 let row_offsets = T::arange(0, 1) + row;
362 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
363}
364
365impl_reduce_float_runtime_op!(ReduceL2Forward);
366
367#[kernel]
371pub fn reduce_sum_square_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
372 x_ptr: T::Pointer<D>,
373 y_ptr: T::Pointer<D>,
374 n_inner: i32,
375 n_outer: i32,
376) where
377 T::I32Tensor: types::Tensor<i32, 1>,
378 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
379 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
380{
381 let row = T::program_id(Axis::X);
382 if row >= n_outer {
383 return;
384 }
385 let col_offsets = T::arange(0, BLOCK_INNER);
386 let offsets = col_offsets + row * n_inner;
387 let mask = col_offsets.lt(n_inner);
388 let x = T::load(
389 x_ptr.add_offsets(offsets),
390 Some(mask),
391 Some(T::zeros::<D>(&[BLOCK_INNER])),
392 &[],
393 None,
394 None,
395 None,
396 false,
397 );
398 let val = T::sum(x * x, Some(0), true);
399 let row_offsets = T::arange(0, 1) + row;
400 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
401}
402
403impl_reduce_num_runtime_op!(ReduceSumSquareForward);
404
405#[kernel]
409pub fn reduce_log_sum_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
410 x_ptr: T::Pointer<D>,
411 y_ptr: T::Pointer<D>,
412 n_inner: i32,
413 n_outer: i32,
414) where
415 T::I32Tensor: types::Tensor<i32, 1>,
416 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
417 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
418{
419 let row = T::program_id(Axis::X);
420 if row >= n_outer {
421 return;
422 }
423 let col_offsets = T::arange(0, BLOCK_INNER);
424 let offsets = col_offsets + row * n_inner;
425 let mask = col_offsets.lt(n_inner);
426 let x = T::load(
427 x_ptr.add_offsets(offsets),
428 Some(mask),
429 Some(T::zeros::<D>(&[BLOCK_INNER])),
430 &[],
431 None,
432 None,
433 None,
434 false,
435 );
436 let sum = T::sum(x, Some(0), true);
437 let val = T::log(sum);
438 let row_offsets = T::arange(0, 1) + row;
439 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
440}
441
442impl_reduce_float_runtime_op!(ReduceLogSumForward);
443
444#[kernel]
448pub fn reduce_log_sum_exp_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
449 x_ptr: T::Pointer<D>,
450 y_ptr: T::Pointer<D>,
451 n_inner: i32,
452 n_outer: i32,
453) where
454 T::I32Tensor: types::Tensor<i32, 1>,
455 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
456 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
457{
458 let row = T::program_id(Axis::X);
459 if row >= n_outer {
460 return;
461 }
462 let col_offsets = T::arange(0, BLOCK_INNER);
463 let offsets = col_offsets + row * n_inner;
464 let mask = col_offsets.lt(n_inner);
465 let neg_inf = T::cast::<f32, D>(
466 T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
467 None,
468 false,
469 );
470 let x = T::load(
471 x_ptr.add_offsets(offsets),
472 Some(mask),
473 Some(neg_inf),
474 &[],
475 None,
476 None,
477 None,
478 false,
479 );
480 let m = T::max(x, Some(0), true); let fill = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_INNER], 0.0_f32), None, false);
484 let x_adj = T::load(
485 x_ptr.add_offsets(offsets),
486 Some(mask),
487 Some(fill),
488 &[],
489 None,
490 None,
491 None,
492 false,
493 );
494 let sum_exp = T::sum(T::exp(x_adj - m), Some(0), true);
495 let val = m + T::log(sum_exp);
496 let row_offsets = T::arange(0, 1) + row;
497 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
498}
499
500impl_reduce_float_runtime_op!(ReduceLogSumExpForward);
501
502#[kernel]
508pub fn reduce_prod_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
509 x_ptr: T::Pointer<D>,
510 y_ptr: T::Pointer<D>,
511 n_inner: i32,
512 n_outer: i32,
513) where
514 T::I32Tensor: types::Tensor<i32, 1>,
515 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
516 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
517{
518 let row = T::program_id(Axis::X);
519 if row >= n_outer {
520 return;
521 }
522 let col_offsets = T::arange(0, BLOCK_INNER);
523 let offsets = col_offsets + row * n_inner;
524 let mask = col_offsets.lt(n_inner);
525 let one_fill = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_INNER], 1.0_f32), None, false);
527 let x = T::load(
528 x_ptr.add_offsets(offsets),
529 Some(mask),
530 Some(one_fill),
531 &[],
532 None,
533 None,
534 None,
535 false,
536 );
537 let val = T::exp(T::sum(T::log(x), Some(0), true));
539 let row_offsets = T::arange(0, 1) + row;
540 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
541}
542
543impl_reduce_float_runtime_op!(ReduceProdForward);
544
545#[kernel]
550pub fn cum_sum_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
551 x_ptr: T::Pointer<D>,
552 y_ptr: T::Pointer<D>,
553 n_inner: i32,
554 n_outer: i32,
555) where
556 T::I32Tensor: types::Tensor<i32, 1>,
557 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
558 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
559{
560 let row = T::program_id(Axis::X);
561 if row >= n_outer {
562 return;
563 }
564 let col_offsets = T::arange(0, BLOCK_INNER);
565 let offsets = col_offsets + row * n_inner;
566 let mask = col_offsets.lt(n_inner);
567 let x = T::load(
568 x_ptr.add_offsets(offsets),
569 Some(mask),
570 Some(T::zeros::<D>(&[BLOCK_INNER])),
571 &[],
572 None,
573 None,
574 None,
575 false,
576 );
577 let y = T::cumsum(x, 0, false);
579 T::store(y_ptr.add_offsets(offsets), y, Some(mask), &[], None, None);
580}
581
582impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for CumSumForward<D> {
583 fn n_activation_inputs(&self) -> usize {
584 1
585 }
586 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
587 vec![]
588 }
589 fn pack_args(
590 &self,
591 inputs: &[(teeny_core::model::RawPtr, &[usize])],
592 _: &[teeny_core::model::RawPtr],
593 output: teeny_core::model::RawPtr,
594 output_shape: &[usize],
595 _: i32,
596 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
597 ) {
598 let n_total: usize = output_shape.iter().product();
600 let n_inner = output_shape.last().copied().unwrap_or(1);
601 let n_outer = n_total / n_inner;
602 visitor.visit_ptr(inputs[0].0);
603 visitor.visit_ptr(output);
604 visitor.visit_i32(n_inner as i32);
605 visitor.visit_i32(n_outer as i32);
606 }
607 fn block(&self) -> [u32; 3] {
608 [self.block_inner as u32, 1, 1]
609 }
610 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
611 let n_total: usize = output_shape.iter().product();
612 let n_inner = output_shape.last().copied().unwrap_or(1);
613 let n_outer = n_total / n_inner;
614 [n_outer as u32, 1, 1]
615 }
616}
617
618#[kernel]
622pub fn cum_prod_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
623 x_ptr: T::Pointer<D>,
624 y_ptr: T::Pointer<D>,
625 n_inner: i32,
626 n_outer: i32,
627) where
628 T::I32Tensor: types::Tensor<i32, 1>,
629 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
630 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
631{
632 let row = T::program_id(Axis::X);
633 if row >= n_outer {
634 return;
635 }
636 let col_offsets = T::arange(0, BLOCK_INNER);
637 let offsets = col_offsets + row * n_inner;
638 let mask = col_offsets.lt(n_inner);
639 let x = T::load(
640 x_ptr.add_offsets(offsets),
641 Some(mask),
642 Some(T::zeros::<D>(&[BLOCK_INNER])),
643 &[],
644 None,
645 None,
646 None,
647 false,
648 );
649 let y = T::cumprod(x, 0, false);
650 T::store(y_ptr.add_offsets(offsets), y, Some(mask), &[], None, None);
651}
652
653impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for CumProdForward<D> {
654 fn n_activation_inputs(&self) -> usize {
655 1
656 }
657 fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
658 vec![]
659 }
660 fn pack_args(
661 &self,
662 inputs: &[(teeny_core::model::RawPtr, &[usize])],
663 _: &[teeny_core::model::RawPtr],
664 output: teeny_core::model::RawPtr,
665 output_shape: &[usize],
666 _: i32,
667 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
668 ) {
669 let n_total: usize = output_shape.iter().product();
670 let n_inner = output_shape.last().copied().unwrap_or(1);
671 let n_outer = n_total / n_inner;
672 visitor.visit_ptr(inputs[0].0);
673 visitor.visit_ptr(output);
674 visitor.visit_i32(n_inner as i32);
675 visitor.visit_i32(n_outer as i32);
676 }
677 fn block(&self) -> [u32; 3] {
678 [self.block_inner as u32, 1, 1]
679 }
680 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
681 let n_total: usize = output_shape.iter().product();
682 let n_inner = output_shape.last().copied().unwrap_or(1);
683 let n_outer = n_total / n_inner;
684 [n_outer as u32, 1, 1]
685 }
686}
687
688#[kernel]
699pub fn global_avg_pool_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
700 x_ptr: T::Pointer<D>,
701 y_ptr: T::Pointer<D>,
702 n_inner: i32,
703 n_outer: i32,
704) where
705 T::I32Tensor: types::Tensor<i32, 1>,
706 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
707 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
708{
709 let row = T::program_id(Axis::X);
710 if row >= n_outer {
711 return;
712 }
713 let col_offsets = T::arange(0, BLOCK_INNER);
714 let offsets = col_offsets + row * n_inner;
715 let mask = col_offsets.lt(n_inner);
716 let x = T::load(
717 x_ptr.add_offsets(offsets),
718 Some(mask),
719 Some(T::zeros::<D>(&[BLOCK_INNER])),
720 &[],
721 None,
722 None,
723 None,
724 false,
725 );
726 let sum = T::sum(x, Some(0), true);
727 let n_f = T::cast::<i32, D>(T::full::<i32>(&[1], n_inner), None, false);
728 let mean = sum / n_f;
729 let row_offsets = T::arange(0, 1) + row;
730 T::store(y_ptr.add_offsets(row_offsets), mean, None, &[], None, None);
731}
732
733impl_reduce_float_runtime_op!(GlobalAvgPoolForward);
734
735#[kernel]
739pub fn global_max_pool_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
740 x_ptr: T::Pointer<D>,
741 y_ptr: T::Pointer<D>,
742 n_inner: i32,
743 n_outer: i32,
744) where
745 T::I32Tensor: types::Tensor<i32, 1>,
746 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
747 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
748{
749 let row = T::program_id(Axis::X);
750 if row >= n_outer {
751 return;
752 }
753 let col_offsets = T::arange(0, BLOCK_INNER);
754 let offsets = col_offsets + row * n_inner;
755 let mask = col_offsets.lt(n_inner);
756 let neg_inf = T::cast::<f32, D>(
757 T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
758 None,
759 false,
760 );
761 let x = T::load(
762 x_ptr.add_offsets(offsets),
763 Some(mask),
764 Some(neg_inf),
765 &[],
766 None,
767 None,
768 None,
769 false,
770 );
771 let val = T::max(x, Some(0), true);
772 let row_offsets = T::arange(0, 1) + row;
773 T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
774}
775
776impl_reduce_float_runtime_op!(GlobalMaxPoolForward);