Skip to main content

ruprim/reduce/launch/
base.rs

1use ruda_kernel::dsl as kernel_dsl;
2use crate::reduce::{
3    ReduceError, ReducePrecision, VectorizationMode,
4    components::{
5        args::{NumericVector, ReduceArgs, TensorArgs, init_tensors},
6        global::{
7            ruda::GlobalFullRudaReduce, plane::GlobalFullPlaneReduce, unit::GlobalFullUnitReduce,
8        },
9        instructions::*,
10    },
11    launch::{ReduceStrategy, RoutineStrategy, generate_vector_size_with_dtypes},
12    output_vectorization_axis,
13    routines::{
14        GlobalReduceBlueprint, ReduceBlueprint, ReduceProblem, ReduceVectorSettings, Routine,
15        ruda::RudaRoutine, plane::PlaneRoutine, unit::UnitRoutine,
16    },
17};
18use ruda_kernel::dsl::prelude::*;
19use ruda_kernel::library::tensor::r#virtual::VirtualTensor;
20
21#[derive(Clone, Copy, Debug)]
22pub struct ReduceDtypes {
23    pub input: StorageType,
24    pub output: StorageType,
25    pub accumulation: StorageType,
26}
27
28/// Launch a reduce kernel. This function assumes that all parameters are already validated.
29/// See the main entrypoint `reduce` in `lib.rs` for an example how to call this function
30/// with the appropriate assumptions.
31#[allow(clippy::too_many_arguments)]
32pub(crate) fn launch_reduce<Run: Runtime>(
33    client: &ComputeClient<Run>,
34    input: TensorBinding<Run>,
35    output: TensorBinding<Run>,
36    reduce_axis: usize,
37    strategy: ReduceStrategy,
38    dtypes: ReduceDtypes,
39    inst: ReduceOperationConfig,
40) -> Result<(), ReduceError> {
41    if output.shape.contains(&0) {
42        return Ok(());
43    }
44    let address_type = input
45        .required_address_type(dtypes.input.size())
46        .max(output.required_address_type(dtypes.output.size()));
47
48    // Number of distinct reductions = product of non-reduce input dims.
49    let reduce_len = input.shape[reduce_axis];
50    let input_elems: usize = input.shape.iter().copied().product();
51    let reduce_count = input_elems / reduce_len;
52
53    let problem = ReduceProblem {
54        reduce_len,
55        reduce_count,
56        axis: reduce_axis,
57        dtypes,
58        address_type,
59    };
60    let vectorization_mode = match input.strides[reduce_axis] {
61        1 => VectorizationMode::Parallel,
62        _ => VectorizationMode::Perpendicular,
63    };
64
65    let out_vec_axis = output_vectorization_axis(&input.strides, reduce_axis, vectorization_mode);
66
67    let (vector_size_input, vector_size_output) = generate_vector_size_with_dtypes::<Run>(
68        client,
69        &input,
70        &output,
71        reduce_axis,
72        [problem.dtypes.input, problem.dtypes.output],
73        vectorization_mode,
74        &strategy.vectorization,
75    );
76    let settings = ReduceVectorSettings {
77        vectorization_mode,
78        vector_size_input,
79        vector_size_output,
80    };
81
82    let (blueprint, settings) = match strategy.routine {
83        RoutineStrategy::Unit(strategy) => {
84            let routine = UnitRoutine;
85            routine.prepare(client, problem, settings, strategy)?
86        }
87        RoutineStrategy::Plane(strategy) => {
88            let routine = PlaneRoutine;
89            routine.prepare(client, problem, settings, strategy)?
90        }
91        RoutineStrategy::Ruda(strategy) => {
92            let routine = RudaRoutine;
93            routine.prepare(client, problem, settings, strategy)?
94        }
95    };
96
97    unsafe {
98        reduce_kernel::launch_unchecked::<TensorArgs, Run>(
99            client,
100            settings.ruda_count,
101            settings.ruda_dim,
102            settings.address_type,
103            settings.vector.vector_size_input,
104            settings.vector.vector_size_output,
105            input.into_tensor_arg(),
106            output.into_tensor_arg(),
107            reduce_axis,
108            out_vec_axis,
109            blueprint,
110            inst,
111            dtypes.input,
112            dtypes.output,
113            dtypes.accumulation,
114        )
115    };
116
117    Ok(())
118}
119
120#[ruda(launch_unchecked, address_type = "dynamic")]
121pub fn reduce_kernel<
122    In: Numeric,
123    InSize: Size,
124    Out: Numeric,
125    OutSize: Size,
126    Acc: Numeric,
127    RA: ReduceArgs,
128>(
129    input: &RA::Input<In, InSize>,
130    output: &mut RA::Output<Out, OutSize>,
131    reduce_axis: usize,
132    out_vec_axis: usize,
133    #[comptime] blueprint: ReduceBlueprint,
134    #[comptime] config: ReduceOperationConfig,
135    #[define(In)] _input_dtype: StorageType,
136    #[define(Out)] _output_dtype: StorageType,
137    #[define(Acc)] _acc_dtype: StorageType,
138) {
139    let (input, mut output) = init_tensors::<RA, In, InSize, Out, OutSize>(input, output);
140    reduce_kernel_virtual::<In, InSize, Out, OutSize, Acc>(
141        &input,
142        &mut output,
143        reduce_axis,
144        out_vec_axis,
145        blueprint,
146        config,
147    );
148}
149
150#[ruda]
151pub fn reduce_kernel_virtual<
152    In: Numeric,
153    InSize: Size,
154    Out: Numeric,
155    OutSize: Size,
156    Acc: Numeric,
157>(
158    input: &VirtualTensor<In, InSize>,
159    output: &mut VirtualTensor<Out, OutSize, ReadWrite>,
160    reduce_axis: usize,
161    out_vec_axis: usize,
162    #[comptime] blueprint: ReduceBlueprint,
163    #[comptime] config: ReduceOperationConfig,
164) {
165    reduce_kernel_inner::<(In, InSize, Acc), (Out, OutSize), ReduceOperation>(
166        input,
167        output,
168        reduce_axis,
169        out_vec_axis,
170        blueprint,
171        config,
172    )
173}
174
175#[ruda]
176fn reduce_kernel_inner<P: ReducePrecision, Out: NumericVector, R: ReduceFamily>(
177    input: &VirtualTensor<P::EI, P::SI>,
178    output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
179    reduce_axis: usize,
180    out_vec_axis: usize,
181    #[comptime] blueprint: ReduceBlueprint,
182    #[comptime] config: R::Config,
183) {
184    let inst = &R::Instruction::<P>::from_config(config);
185
186    match blueprint.global {
187        GlobalReduceBlueprint::Ruda(ruda) => {
188            GlobalFullRudaReduce::execute::<P, Out, R::Instruction<P>>(
189                input,
190                output,
191                reduce_axis,
192                out_vec_axis,
193                inst,
194                blueprint.vectorization_mode,
195                ruda,
196            )
197        }
198        GlobalReduceBlueprint::Plane(plane) => {
199            GlobalFullPlaneReduce::execute::<P, Out, R::Instruction<P>>(
200                input,
201                output,
202                reduce_axis,
203                out_vec_axis,
204                inst,
205                blueprint.vectorization_mode,
206                plane,
207            )
208        }
209        GlobalReduceBlueprint::Unit(unit) => {
210            GlobalFullUnitReduce::execute::<P, Out, R::Instruction<P>>(
211                input,
212                output,
213                reduce_axis,
214                out_vec_axis,
215                inst,
216                blueprint.vectorization_mode,
217                unit,
218            )
219        }
220    };
221}