Skip to main content

ruprim/reduce/components/global/
unit.rs

1use 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}