1use ruda_kernel::dsl as kernel_dsl;
2use super::{
3 ArgMax, ArgMin, ArgTopK, Max, MaxAbs, Mean, Min, Prod, ReduceFamily, ReduceInstruction,
4 ReduceRequirements, SharedAccumulator, Sum,
5};
6use crate::reduce::components::instructions::{
7 Accumulator, AccumulatorFormat, Item, SharedAccumulatorKind, TopK,
8};
9use crate::reduce::{
10 ReduceDtypes,
11 components::{
12 instructions::{ReduceStep, Value},
13 precision::ReducePrecision,
14 },
15};
16use ruda_kernel::dsl::ir::ElemType;
17use ruda_kernel::dsl::ir::FloatKind;
18use ruda_kernel::dsl::ir::IntKind;
19use ruda_kernel::dsl::ir::UIntKind;
20use ruda_kernel::dsl::prelude::*;
21use serde::{Deserialize, Serialize};
22
23#[derive(Debug, RudaType, Clone)]
24pub enum ReduceOperation {
25 Sum(Sum),
26 Prod(Prod),
27 Mean(Mean),
28 MaxAbs(MaxAbs),
29 ArgMax(ArgMax),
30 ArgMin(ArgMin),
31 Max(Max),
32 Min(Min),
33 ArgTopK(ArgTopK),
34 TopK(TopK),
35}
36
37#[derive_ruda_comptime]
38#[derive(Serialize, Deserialize)]
39pub enum ReduceOperationConfig {
40 Sum,
41 Prod,
42 Mean,
43 MaxAbs,
44 ArgMax,
45 ArgMin,
46 Max,
47 Min,
48 ArgTopK(usize),
49 TopK(usize),
50}
51
52impl ReduceOperationConfig {
53 pub fn precision(&self, input: ElemType, output: Option<ElemType>) -> ReduceDtypes {
55 match self {
56 ReduceOperationConfig::Sum
57 | ReduceOperationConfig::Prod
58 | ReduceOperationConfig::Mean => {}
59 ReduceOperationConfig::MaxAbs
61 | ReduceOperationConfig::Max
62 | ReduceOperationConfig::TopK(_)
63 | ReduceOperationConfig::Min => {
64 return ReduceDtypes {
65 input: input.into(),
66 output: input.into(),
67 accumulation: input.into(),
68 };
69 }
70 ReduceOperationConfig::ArgMax
71 | ReduceOperationConfig::ArgMin
72 | ReduceOperationConfig::ArgTopK(_) => {
73 return ReduceDtypes {
74 input: input.into(),
75 output: output
76 .expect("ArgMax, ArgMin and ArgTopK must specify output type")
77 .into(),
78 accumulation: input.into(),
79 };
80 }
81 };
82
83 match input {
84 ElemType::Float(kind) => {
85 let acc = match kind {
86 FloatKind::F64 => f64::as_type_native_unchecked(),
87 _ => f32::as_type_native_unchecked(),
88 };
89
90 ReduceDtypes {
91 input: input.into(),
92 output: input.into(),
93 accumulation: acc.storage_type(),
94 }
95 }
96 ElemType::Int(kind) => {
97 let acc = match kind {
98 IntKind::I64 => i64::as_type_native_unchecked(),
99 _ => i32::as_type_native_unchecked(),
100 };
101
102 ReduceDtypes {
103 input: input.into(),
104 output: input.into(),
105 accumulation: acc.storage_type(),
106 }
107 }
108 ElemType::UInt(kind) => {
109 let acc = match kind {
110 UIntKind::U64 => u64::as_type_native_unchecked(),
111 _ => u32::as_type_native_unchecked(),
112 };
113
114 ReduceDtypes {
115 input: input.into(),
116 output: input.into(),
117 accumulation: acc.storage_type(),
118 }
119 }
120 ElemType::Bool => panic!("Can't reduce on booleans"),
121 }
122 }
123}
124
125impl ReduceFamily for ReduceOperation {
126 type Instruction<P: ReducePrecision> = Self;
127 type Config = ReduceOperationConfig;
128}
129
130#[derive(RudaType)]
131pub struct DynamicSharedAccumulator<P: ReducePrecision> {
132 pub elements: SharedAccumulatorKind<Vector<P::EA, P::SI>>,
133 pub args: SharedAccumulatorKind<Vector<u32, P::SI>>,
134}
135
136#[derive(RudaType)]
137pub struct DynamicAccumulator<P: ReducePrecision> {
138 pub elements: Value<Vector<P::EA, P::SI>>,
139 pub args: Value<Vector<u32, P::SI>>,
140}
141
142#[ruda]
143impl<P: ReducePrecision, I: ReduceInstruction<P>> SharedAccumulator<P, I>
144 for DynamicSharedAccumulator<P>
145{
146 fn allocate(#[comptime] length: usize, #[comptime] coordinate: bool, inst: &I) -> Self {
147 let format = I::accumulator_format(inst);
148 match comptime!(format) {
149 AccumulatorFormat::Single => {
150 let elements = SharedMemory::new(length);
151 let args = if coordinate {
153 let args = SharedMemory::new(length);
154 SharedAccumulatorKind::new_Single(args)
155 } else {
156 SharedAccumulatorKind::new_None()
157 };
158 DynamicSharedAccumulator::<P> {
159 elements: SharedAccumulatorKind::new_Single(elements),
160 args,
161 }
162 }
163 AccumulatorFormat::Multiple(len) => {
164 let mut elements = Sequence::new();
165 #[unroll]
166 for _ in 0..len {
167 elements.push(SharedMemory::new(length));
168 }
169
170 if comptime!(!coordinate) {
171 DynamicSharedAccumulator::<P> {
172 elements: SharedAccumulatorKind::new_Multiple(elements),
173 args: SharedAccumulatorKind::new_None(),
174 }
175 } else {
176 let mut args = Sequence::new();
177 #[unroll]
178 for _ in 0..len {
179 args.push(SharedMemory::new(length));
180 }
181 DynamicSharedAccumulator::<P> {
182 elements: SharedAccumulatorKind::new_Multiple(elements),
183 args: SharedAccumulatorKind::new_Multiple(args),
184 }
185 }
186 }
187 }
188 }
189
190 fn read(accumulator: &Self, index: usize) -> Accumulator<P> {
191 let elements = accumulator.elements.get(index);
192 let args = accumulator.args.get(index);
193
194 Accumulator::<P> { elements, args }
195 }
196
197 fn write(accumulator: &mut Self, index: usize, item: Accumulator<P>) {
198 accumulator.elements.set(index, item.elements);
199 accumulator.args.set(index, item.args);
200 }
201}
202
203#[ruda]
204impl<P: ReducePrecision> ReduceInstruction<P> for ReduceOperation {
205 type SharedAccumulator = DynamicSharedAccumulator<P>;
206 type Config = ReduceOperationConfig;
207
208 fn requirements(this: &Self) -> ReduceRequirements {
209 match this {
210 ReduceOperation::Sum(sum) => <Sum as ReduceInstruction<P>>::requirements(sum),
211 ReduceOperation::Prod(prod) => <Prod as ReduceInstruction<P>>::requirements(prod),
212 ReduceOperation::Mean(mean) => <Mean as ReduceInstruction<P>>::requirements(mean),
213 ReduceOperation::MaxAbs(max_abs) => {
214 <MaxAbs as ReduceInstruction<P>>::requirements(max_abs)
215 }
216 ReduceOperation::ArgMax(arg_max) => {
217 <ArgMax as ReduceInstruction<P>>::requirements(arg_max)
218 }
219 ReduceOperation::ArgMin(arg_min) => {
220 <ArgMin as ReduceInstruction<P>>::requirements(arg_min)
221 }
222 ReduceOperation::ArgTopK(arg_topk) => {
223 <ArgTopK as ReduceInstruction<P>>::requirements(arg_topk)
224 }
225 ReduceOperation::TopK(topk) => <TopK as ReduceInstruction<P>>::requirements(topk),
226 ReduceOperation::Max(max) => <Max as ReduceInstruction<P>>::requirements(max),
227 ReduceOperation::Min(min) => <Min as ReduceInstruction<P>>::requirements(min),
228 }
229 }
230
231 fn accumulator_format(this: &Self) -> comptime_type!(AccumulatorFormat) {
232 match this {
233 ReduceOperation::Sum(sum) => <Sum as ReduceInstruction<P>>::accumulator_format(sum),
234 ReduceOperation::Prod(prod) => <Prod as ReduceInstruction<P>>::accumulator_format(prod),
235 ReduceOperation::Mean(mean) => <Mean as ReduceInstruction<P>>::accumulator_format(mean),
236 ReduceOperation::MaxAbs(maxabs) => {
237 <MaxAbs as ReduceInstruction<P>>::accumulator_format(maxabs)
238 }
239 ReduceOperation::ArgMax(argmax) => {
240 <ArgMax as ReduceInstruction<P>>::accumulator_format(argmax)
241 }
242 ReduceOperation::ArgMin(argmin) => {
243 <ArgMin as ReduceInstruction<P>>::accumulator_format(argmin)
244 }
245 ReduceOperation::ArgTopK(args) => {
246 <ArgTopK as ReduceInstruction<P>>::accumulator_format(args)
247 }
248 ReduceOperation::Max(max) => <Max as ReduceInstruction<P>>::accumulator_format(max),
249 ReduceOperation::Min(min) => <Min as ReduceInstruction<P>>::accumulator_format(min),
250 ReduceOperation::TopK(topk) => <TopK as ReduceInstruction<P>>::accumulator_format(topk),
251 }
252 }
253
254 fn from_config(#[comptime] config: Self::Config) -> Self {
255 match config {
256 ReduceOperationConfig::Sum => ReduceOperation::new_Sum(Sum {}),
257 ReduceOperationConfig::Prod => ReduceOperation::new_Prod(Prod {}),
258 ReduceOperationConfig::Mean => ReduceOperation::new_Mean(Mean { sum: Sum {} }),
259 ReduceOperationConfig::MaxAbs => ReduceOperation::new_MaxAbs(MaxAbs {}),
260 ReduceOperationConfig::ArgMax => ReduceOperation::new_ArgMax(ArgMax {}),
261 ReduceOperationConfig::ArgMin => ReduceOperation::new_ArgMin(ArgMin {}),
262 ReduceOperationConfig::ArgTopK(k) => ReduceOperation::new_ArgTopK(ArgTopK { k }),
263 ReduceOperationConfig::Max => ReduceOperation::new_Max(Max {}),
264 ReduceOperationConfig::Min => ReduceOperation::new_Min(Min {}),
265 ReduceOperationConfig::TopK(k) => ReduceOperation::new_TopK(TopK { k }),
266 }
267 }
268
269 fn null_input(this: &Self) -> Vector<P::EI, P::SI> {
270 match this {
271 ReduceOperation::Sum(sum) => <Sum as ReduceInstruction<P>>::null_input(sum),
272 ReduceOperation::Prod(prod) => <Prod as ReduceInstruction<P>>::null_input(prod),
273 ReduceOperation::Mean(mean) => <Mean as ReduceInstruction<P>>::null_input(mean),
274 ReduceOperation::MaxAbs(maxabs) => <MaxAbs as ReduceInstruction<P>>::null_input(maxabs),
275 ReduceOperation::ArgMax(argmax) => <ArgMax as ReduceInstruction<P>>::null_input(argmax),
276 ReduceOperation::ArgMin(argmin) => <ArgMin as ReduceInstruction<P>>::null_input(argmin),
277 ReduceOperation::ArgTopK(args) => <ArgTopK as ReduceInstruction<P>>::null_input(args),
278 ReduceOperation::Max(max) => <Max as ReduceInstruction<P>>::null_input(max),
279 ReduceOperation::Min(min) => <Min as ReduceInstruction<P>>::null_input(min),
280 ReduceOperation::TopK(topk) => <TopK as ReduceInstruction<P>>::null_input(topk),
281 }
282 }
283
284 fn null_accumulator(this: &Self) -> Accumulator<P> {
285 match this {
286 ReduceOperation::Sum(sum) => <Sum as ReduceInstruction<P>>::null_accumulator(sum),
287 ReduceOperation::Mean(sum) => <Mean as ReduceInstruction<P>>::null_accumulator(sum),
288 ReduceOperation::Prod(prod) => <Prod as ReduceInstruction<P>>::null_accumulator(prod),
289 ReduceOperation::MaxAbs(maxabs) => {
290 <MaxAbs as ReduceInstruction<P>>::null_accumulator(maxabs)
291 }
292 ReduceOperation::ArgMax(argmax) => {
293 <ArgMax as ReduceInstruction<P>>::null_accumulator(argmax)
294 }
295 ReduceOperation::ArgMin(argmin) => {
296 <ArgMin as ReduceInstruction<P>>::null_accumulator(argmin)
297 }
298 ReduceOperation::ArgTopK(args) => {
299 <ArgTopK as ReduceInstruction<P>>::null_accumulator(args)
300 }
301 ReduceOperation::Max(max) => <Max as ReduceInstruction<P>>::null_accumulator(max),
302 ReduceOperation::Min(min) => <Min as ReduceInstruction<P>>::null_accumulator(min),
303 ReduceOperation::TopK(topk) => <TopK as ReduceInstruction<P>>::null_accumulator(topk),
304 }
305 }
306
307 fn reduce(
308 this: &Self,
309 accumulator: &mut Accumulator<P>,
310 item: Item<P>,
311 #[comptime] reduce_step: ReduceStep,
312 ) {
313 match this {
314 ReduceOperation::Sum(sum) => {
315 <Sum as ReduceInstruction<P>>::reduce(sum, accumulator, item, reduce_step)
316 }
317 ReduceOperation::Prod(sum) => {
318 <Prod as ReduceInstruction<P>>::reduce(sum, accumulator, item, reduce_step)
319 }
320 ReduceOperation::Mean(sum) => {
321 <Mean as ReduceInstruction<P>>::reduce(sum, accumulator, item, reduce_step)
322 }
323 ReduceOperation::MaxAbs(maxabs) => {
324 <MaxAbs as ReduceInstruction<P>>::reduce(maxabs, accumulator, item, reduce_step)
325 }
326 ReduceOperation::ArgMax(argmax) => {
327 <ArgMax as ReduceInstruction<P>>::reduce(argmax, accumulator, item, reduce_step)
328 }
329 ReduceOperation::ArgMin(argmin) => {
330 <ArgMin as ReduceInstruction<P>>::reduce(argmin, accumulator, item, reduce_step)
331 }
332 ReduceOperation::ArgTopK(argtopk) => {
333 <ArgTopK as ReduceInstruction<P>>::reduce(argtopk, accumulator, item, reduce_step)
334 }
335 ReduceOperation::Max(max) => {
336 <Max as ReduceInstruction<P>>::reduce(max, accumulator, item, reduce_step)
337 }
338 ReduceOperation::Min(min) => {
339 <Min as ReduceInstruction<P>>::reduce(min, accumulator, item, reduce_step)
340 }
341 ReduceOperation::TopK(topk) => {
342 <TopK as ReduceInstruction<P>>::reduce(topk, accumulator, item, reduce_step)
343 }
344 }
345 }
346
347 fn plane_reduce_inplace(this: &Self, accumulator: &mut Accumulator<P>) {
348 match this {
349 ReduceOperation::Sum(sum) => {
350 <Sum as ReduceInstruction<P>>::plane_reduce_inplace(sum, accumulator)
351 }
352 ReduceOperation::Prod(prod) => {
353 <Prod as ReduceInstruction<P>>::plane_reduce_inplace(prod, accumulator)
354 }
355 ReduceOperation::Mean(mean) => {
356 <Mean as ReduceInstruction<P>>::plane_reduce_inplace(mean, accumulator)
357 }
358 ReduceOperation::MaxAbs(max_abs) => {
359 <MaxAbs as ReduceInstruction<P>>::plane_reduce_inplace(max_abs, accumulator)
360 }
361 ReduceOperation::ArgMax(arg_max) => {
362 <ArgMax as ReduceInstruction<P>>::plane_reduce_inplace(arg_max, accumulator)
363 }
364 ReduceOperation::ArgMin(arg_min) => {
365 <ArgMin as ReduceInstruction<P>>::plane_reduce_inplace(arg_min, accumulator)
366 }
367 ReduceOperation::Max(max) => {
368 <Max as ReduceInstruction<P>>::plane_reduce_inplace(max, accumulator)
369 }
370 ReduceOperation::Min(min) => {
371 <Min as ReduceInstruction<P>>::plane_reduce_inplace(min, accumulator)
372 }
373 ReduceOperation::ArgTopK(argtopk) => {
374 <ArgTopK as ReduceInstruction<P>>::plane_reduce_inplace(argtopk, accumulator)
375 }
376 ReduceOperation::TopK(topk) => {
377 <TopK as ReduceInstruction<P>>::plane_reduce_inplace(topk, accumulator)
378 }
379 }
380 }
381
382 fn fuse_accumulators(this: &Self, accumulator: &mut Accumulator<P>, other: &Accumulator<P>) {
383 match this {
384 ReduceOperation::Sum(sum) => {
385 <Sum as ReduceInstruction<P>>::fuse_accumulators(sum, accumulator, other)
386 }
387 ReduceOperation::Prod(prod) => {
388 <Prod as ReduceInstruction<P>>::fuse_accumulators(prod, accumulator, other)
389 }
390 ReduceOperation::Mean(mean) => {
391 <Mean as ReduceInstruction<P>>::fuse_accumulators(mean, accumulator, other)
392 }
393 ReduceOperation::MaxAbs(maxabs) => {
394 <MaxAbs as ReduceInstruction<P>>::fuse_accumulators(maxabs, accumulator, other)
395 }
396 ReduceOperation::ArgMax(argmax) => {
397 <ArgMax as ReduceInstruction<P>>::fuse_accumulators(argmax, accumulator, other)
398 }
399 ReduceOperation::ArgMin(argmin) => {
400 <ArgMin as ReduceInstruction<P>>::fuse_accumulators(argmin, accumulator, other)
401 }
402 ReduceOperation::ArgTopK(argtopk) => {
403 <ArgTopK as ReduceInstruction<P>>::fuse_accumulators(argtopk, accumulator, other)
404 }
405 ReduceOperation::Max(max) => {
406 <Max as ReduceInstruction<P>>::fuse_accumulators(max, accumulator, other)
407 }
408 ReduceOperation::Min(min) => {
409 <Min as ReduceInstruction<P>>::fuse_accumulators(min, accumulator, other)
410 }
411 ReduceOperation::TopK(topk) => {
412 <TopK as ReduceInstruction<P>>::fuse_accumulators(topk, accumulator, other)
413 }
414 }
415 }
416
417 fn to_output_parallel<Out: Numeric>(
418 this: &Self,
419 accumulator: Accumulator<P>,
420 shape_axis_reduce: usize,
421 ) -> Value<Out> {
422 match this {
423 ReduceOperation::Sum(sum) => <Sum as ReduceInstruction<P>>::to_output_parallel::<Out>(
424 sum,
425 accumulator,
426 shape_axis_reduce,
427 ),
428 ReduceOperation::Prod(prod) => {
429 <Prod as ReduceInstruction<P>>::to_output_parallel::<Out>(
430 prod,
431 accumulator,
432 shape_axis_reduce,
433 )
434 }
435 ReduceOperation::Mean(mean) => {
436 <Mean as ReduceInstruction<P>>::to_output_parallel::<Out>(
437 mean,
438 accumulator,
439 shape_axis_reduce,
440 )
441 }
442 ReduceOperation::MaxAbs(maxabs) => {
443 <MaxAbs as ReduceInstruction<P>>::to_output_parallel::<Out>(
444 maxabs,
445 accumulator,
446 shape_axis_reduce,
447 )
448 }
449 ReduceOperation::ArgMax(argmax) => {
450 <ArgMax as ReduceInstruction<P>>::to_output_parallel::<Out>(
451 argmax,
452 accumulator,
453 shape_axis_reduce,
454 )
455 }
456 ReduceOperation::ArgMin(argmin) => {
457 <ArgMin as ReduceInstruction<P>>::to_output_parallel::<Out>(
458 argmin,
459 accumulator,
460 shape_axis_reduce,
461 )
462 }
463 ReduceOperation::ArgTopK(argtopk) => {
464 <ArgTopK as ReduceInstruction<P>>::to_output_parallel::<Out>(
465 argtopk,
466 accumulator,
467 shape_axis_reduce,
468 )
469 }
470 ReduceOperation::Max(max) => <Max as ReduceInstruction<P>>::to_output_parallel::<Out>(
471 max,
472 accumulator,
473 shape_axis_reduce,
474 ),
475 ReduceOperation::Min(min) => <Min as ReduceInstruction<P>>::to_output_parallel::<Out>(
476 min,
477 accumulator,
478 shape_axis_reduce,
479 ),
480 ReduceOperation::TopK(topk) => {
481 <TopK as ReduceInstruction<P>>::to_output_parallel::<Out>(
482 topk,
483 accumulator,
484 shape_axis_reduce,
485 )
486 }
487 }
488 }
489
490 fn to_output_perpendicular<Out: Numeric>(
491 this: &Self,
492 accumulator: Accumulator<P>,
493 shape_axis_reduce: usize,
494 ) -> Value<Vector<Out, P::SI>> {
495 match this {
496 ReduceOperation::Sum(sum) => <Sum as ReduceInstruction<P>>::to_output_perpendicular::<
497 Out,
498 >(sum, accumulator, shape_axis_reduce),
499 ReduceOperation::Prod(prod) => {
500 <Prod as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
501 prod,
502 accumulator,
503 shape_axis_reduce,
504 )
505 }
506 ReduceOperation::Mean(mean) => {
507 <Mean as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
508 mean,
509 accumulator,
510 shape_axis_reduce,
511 )
512 }
513 ReduceOperation::MaxAbs(maxabs) => {
514 <MaxAbs as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
515 maxabs,
516 accumulator,
517 shape_axis_reduce,
518 )
519 }
520 ReduceOperation::ArgMax(args) => {
521 <ArgMax as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
522 args,
523 accumulator,
524 shape_axis_reduce,
525 )
526 }
527 ReduceOperation::ArgTopK(argtopk) => {
528 <ArgTopK as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
529 argtopk,
530 accumulator,
531 shape_axis_reduce,
532 )
533 }
534 ReduceOperation::ArgMin(args) => {
535 <ArgMin as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
536 args,
537 accumulator,
538 shape_axis_reduce,
539 )
540 }
541 ReduceOperation::Max(max) => <Max as ReduceInstruction<P>>::to_output_perpendicular::<
542 Out,
543 >(max, accumulator, shape_axis_reduce),
544 ReduceOperation::Min(min) => <Min as ReduceInstruction<P>>::to_output_perpendicular::<
545 Out,
546 >(min, accumulator, shape_axis_reduce),
547 ReduceOperation::TopK(topk) => {
548 <TopK as ReduceInstruction<P>>::to_output_perpendicular::<Out>(
549 topk,
550 accumulator,
551 shape_axis_reduce,
552 )
553 }
554 }
555 }
556}