Skip to main content

ruprim/reduce/components/readers/
base.rs

1use ruda_kernel::dsl as kernel_dsl;
2use crate::reduce::{
3    BoundChecks, ReduceInstruction, ReducePrecision, VectorizationMode,
4    components::{
5        args::NumericVector,
6        instructions::{ReduceRequirements, Value},
7        readers::{parallel::ParallelReader, perpendicular::PerpendicularReader},
8    },
9};
10use ruda_kernel::dsl::prelude::*;
11use ruda_kernel::library::tensor::r#virtual::VirtualTensor;
12
13#[derive(RudaType)]
14pub enum Reader<P: ReducePrecision> {
15    Parallel(ParallelReader<P>),
16    Perpendicular(PerpendicularReader<P>),
17}
18
19#[ruda]
20impl<P: ReducePrecision> Reader<P> {
21    #[allow(clippy::too_many_arguments)]
22    pub fn new<I: ReduceInstruction<P>, Out: NumericVector>(
23        input: &VirtualTensor<P::EI, P::SI>,
24        output: &mut VirtualTensor<Out::T, Out::N, ReadWrite>,
25        inst: &I,
26        reduce_axis: usize,
27        reduce_index: usize,
28        idle: ComptimeOption<bool>,
29        #[comptime] bound_checks: BoundChecks,
30        #[comptime] vectorization_mode: VectorizationMode,
31        #[comptime] plane_dim_ceil: bool,
32    ) -> Reader<P> {
33        let effective_plane_dim = if plane_dim_ceil {
34            min(RUDA_DIM_X, PLANE_DIM)
35        } else {
36            RUDA_DIM_X
37        };
38        match vectorization_mode {
39            VectorizationMode::Parallel => {
40                Reader::<P>::new_Parallel(ParallelReader::<P>::new::<I, Out>(
41                    input,
42                    output,
43                    inst,
44                    reduce_axis,
45                    reduce_index,
46                    idle,
47                    effective_plane_dim,
48                    bound_checks,
49                ))
50            }
51            VectorizationMode::Perpendicular => {
52                Reader::<P>::new_Perpendicular(PerpendicularReader::<P>::new::<I, Out>(
53                    input,
54                    output,
55                    inst,
56                    reduce_axis,
57                    reduce_index,
58                    idle,
59                    effective_plane_dim,
60                    bound_checks,
61                ))
62            }
63        }
64    }
65}
66
67#[ruda]
68pub fn new_coordinates<N: Size>(
69    coordinate: usize,
70    requirements: ReduceRequirements,
71    #[comptime] vectorization_mode: VectorizationMode,
72) -> Value<Vector<u32, N>> {
73    if requirements.coordinates.comptime() {
74        // TODO: Make this generic to allow 64-bit coordinate output.
75        // Can't directly use `usize` for the buffer, since its size isn't defined beyond the
76        // kernel boundary.
77        Value::new_single(fill_coordinate_vector(
78            coordinate as u32,
79            vectorization_mode,
80        ))
81    } else {
82        Value::new_None()
83    }
84}
85
86// If vectorization mode is parallel, fill a vector with `x, x+1, ... x+ vector_size - 1` where `x = first`.
87// If vectorization mode is perpendicular, fill a vector with `x, x, ... x` where `x = first`.
88#[ruda]
89pub(crate) fn fill_coordinate_vector<N: Size>(
90    first: u32,
91    #[comptime] vectorization_mode: VectorizationMode,
92) -> Vector<u32, N> {
93    match vectorization_mode {
94        VectorizationMode::Parallel => {
95            let mut coordinates = Vector::empty();
96            #[unroll]
97            for j in 0..N::value() {
98                coordinates[j] = first + j as u32;
99            }
100            coordinates
101        }
102        VectorizationMode::Perpendicular => Vector::empty().fill(first),
103    }
104}