Skip to main content

ruprim/reduce/components/writers/
base.rs

1use ruda_kernel::dsl as kernel_dsl;
2use crate::reduce::{
3    ReduceInstruction, ReducePrecision, VectorizationMode,
4    components::{
5        args::NumericVector,
6        instructions::{Accumulator, AccumulatorFormat},
7        writers::{parallel::ParallelWriter, perpendicular::PerpendicularWriter},
8    },
9};
10use ruda_kernel::dsl::prelude::*;
11use ruda_kernel::library::tensor::r#virtual::VirtualTensor;
12
13#[derive(RudaType)]
14/// Abstract how data is written to global memory.
15///
16/// Depending on the problem kind, writes might be buffered to optimize vectorization, only
17/// happening when [Writer::commit()] is called.
18pub enum Writer<Out: NumericVector> {
19    Parallel(ParallelWriter<Out>),
20    Perpendicular(PerpendicularWriter<Out>),
21}
22
23#[ruda]
24impl<Out: NumericVector> Writer<Out> {
25    pub fn new<P: ReducePrecision>(
26        input: &VirtualTensor<P::EI, P::SI>,
27        output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
28        reduce_axis: usize,
29        out_vec_axis: usize,
30        write_index: usize,
31        #[comptime] vectorization_mode: VectorizationMode,
32        #[comptime] acc_format: AccumulatorFormat,
33    ) -> Writer<Out> {
34        match vectorization_mode {
35            VectorizationMode::Parallel => {
36                Writer::<Out>::new_Parallel(ParallelWriter::<Out>::new::<P>(
37                    input,
38                    output,
39                    reduce_axis,
40                    out_vec_axis,
41                    write_index,
42                    acc_format,
43                ))
44            }
45            VectorizationMode::Perpendicular => {
46                Writer::<Out>::new_Perpendicular(PerpendicularWriter::<Out>::new::<P>(
47                    input,
48                    output,
49                    reduce_axis,
50                    out_vec_axis,
51                    write_index,
52                    acc_format,
53                ))
54            }
55        }
56    }
57
58    pub fn write<P: ReducePrecision, I: ReduceInstruction<P>>(
59        &mut self,
60        local_index: usize,
61        accumulator: Accumulator<P>,
62        inst: &I,
63    ) {
64        match self {
65            Writer::Parallel(writer) => writer.write::<P, I>(local_index, accumulator, inst),
66            Writer::Perpendicular(writer) => writer.write::<P, I>(local_index, accumulator, inst),
67        }
68    }
69
70    pub fn commit_required(&self) -> comptime_type!(bool) {
71        match self {
72            Writer::Parallel(writer) => writer.commit_required(),
73            Writer::Perpendicular(writer) => writer.commit_required(),
74        }
75    }
76
77    pub fn commit(&mut self) {
78        match self {
79            Writer::Parallel(writer) => writer.commit(),
80            Writer::Perpendicular(writer) => writer.commit(),
81        }
82    }
83
84    pub fn write_count(&self) -> comptime_type!(VectorSize) {
85        match self {
86            Writer::Parallel(writer) => writer.write_count(),
87            Writer::Perpendicular(writer) => writer.write_count(),
88        }
89    }
90}