ruprim/reduce/components/readers/
base.rs1use 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 Value::new_single(fill_coordinate_vector(
78 coordinate as u32,
79 vectorization_mode,
80 ))
81 } else {
82 Value::new_None()
83 }
84}
85
86#[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}