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#[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 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}