ruprim/reduce/components/writers/
base.rs1use 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)]
14pub 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}