Skip to main content

ruprim/reduce/components/instructions/
mixed.rs

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    /// Computes the best case precision for the given config.
54    pub fn precision(&self, input: ElemType, output: Option<ElemType>) -> ReduceDtypes {
55        match self {
56            ReduceOperationConfig::Sum
57            | ReduceOperationConfig::Prod
58            | ReduceOperationConfig::Mean => {}
59            // No benefit to mixed precision accumulation.
60            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                // TODO how to put multiple?
152                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}