ruprim/reduce/components/global/
unit.rs1use ruda_kernel::dsl as kernel_dsl;
2use crate::reduce::{
3 BoundChecks, ReduceInstruction, ReducePrecision, VectorizationMode,
4 components::{
5 args::NumericVector,
6 global::idle_check,
7 instructions::{Accumulator, ReduceStep, reduce_inplace},
8 readers::{Reader, unit::UnitReader},
9 writers::Writer,
10 },
11 routines::UnitReduceBlueprint,
12};
13use ruda_kernel::dsl::prelude::*;
14use ruda_kernel::library::tensor::r#virtual::VirtualTensor;
15
16#[derive(RudaType)]
17pub struct GlobalFullUnitReduce;
18
19#[ruda]
20impl GlobalFullUnitReduce {
21 pub fn execute<P: ReducePrecision, Out: NumericVector, I: ReduceInstruction<P>>(
22 input: &VirtualTensor<P::EI, P::SI>,
23 output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
24 reduce_axis: usize,
25 out_vec_axis: usize,
26 inst: &I,
27 #[comptime] vectorization_mode: VectorizationMode,
28 #[comptime] blueprint: UnitReduceBlueprint,
29 ) {
30 let acc_format = I::accumulator_format(inst);
31 let write_index = ABSOLUTE_POS;
32
33 let mut writer = Writer::<Out>::new::<P>(
34 input,
35 output,
36 reduce_axis,
37 out_vec_axis,
38 write_index,
39 vectorization_mode,
40 acc_format,
41 );
42
43 let write_count = writer.write_count();
44 let reduce_index_start = write_index * write_count;
45
46 let idle = idle_check::<P, Out>(
47 input,
48 output,
49 reduce_axis,
50 reduce_index_start,
51 vectorization_mode,
52 blueprint.unit_idle,
53 );
54
55 for b in 0..write_count {
56 let reduce_index = reduce_index_start + b;
57 let accumulator = Self::reduce_single::<P, Out, I>(
58 input,
59 output,
60 reduce_axis,
61 reduce_index,
62 inst,
63 idle,
64 vectorization_mode,
65 );
66 writer.write::<P, I>(b, accumulator, inst);
67 }
68
69 writer.commit();
70 }
71
72 #[allow(clippy::too_many_arguments)]
73 pub fn reduce_single<P: ReducePrecision, Out: NumericVector, I: ReduceInstruction<P>>(
74 input: &VirtualTensor<P::EI, P::SI>,
75 output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
76 reduce_axis: usize,
77 reduce_index: usize,
78 inst: &I,
79 idle: ComptimeOption<bool>,
80 #[comptime] vectorization_mode: VectorizationMode,
81 ) -> Accumulator<P> {
82 let reader = Reader::<P>::new::<I, Out>(
83 input,
84 output,
85 inst,
86 reduce_axis,
87 reduce_index,
88 idle,
89 comptime!(BoundChecks::None),
90 vectorization_mode,
91 false,
92 );
93 let reader = UnitReader::<P>::new(reader);
94
95 let mut accumulator = I::null_accumulator(inst);
96
97 for i in 0..reader.length() {
98 let item = reader.read(i);
99 reduce_inplace::<P, I>(inst, &mut accumulator, item, ReduceStep::Identity);
100 }
101
102 accumulator
103 }
104}