Skip to main content

ruprim/reduce/routines/
unit.rs

1use ruda_kernel::dsl as kernel_dsl;
2use super::{
3    GlobalReduceBlueprint, ReduceBlueprint, ReduceLaunchSettings, ReduceProblem,
4    ReduceVectorSettings,
5};
6use crate::reduce::{
7    IdleMode, ReduceError, VectorizationMode,
8    launch::calculate_plane_count_per_ruda,
9    routines::{BlueprintStrategy, Routine, UnitReduceBlueprint},
10};
11use ruda_kernel::dsl::RudaCount;
12use ruda_kernel::dsl::RudaDim;
13use ruda_kernel::dsl::Runtime;
14use ruda_kernel::dsl::client::ComputeClient;
15use ruda_kernel::tiling::ruda_count::ruda_count_spread_with_total;
16
17#[derive(Debug, Clone)]
18pub struct UnitRoutine;
19
20#[derive(Debug, Clone)]
21pub struct UnitStrategy;
22
23impl Routine for UnitRoutine {
24    type Strategy = UnitStrategy;
25    type Blueprint = UnitReduceBlueprint;
26
27    fn prepare<R: Runtime>(
28        &self,
29        client: &ruda_kernel::dsl::prelude::ComputeClient<R>,
30        problem: ReduceProblem,
31        settings: ReduceVectorSettings,
32        strategy: BlueprintStrategy<Self>,
33    ) -> Result<(ReduceBlueprint, ReduceLaunchSettings), ReduceError> {
34        let address_type = problem.address_type;
35        let (blueprint, ruda_dim, ruda_count) = match strategy {
36            BlueprintStrategy::Forced(blueprint, ruda_dim) => {
37                super::validate_ruda_dim(client, ruda_dim)?;
38                let working_units = working_units(&settings, &problem);
39                let num_units_in_ruda = ruda_dim.num_elems();
40                let working_rudas = working_units.div_ceil(num_units_in_ruda as usize);
41
42                let (ruda_count, launched_rudas) =
43                    ruda_count_spread_with_total(client, working_rudas);
44
45                let unit_idle = !working_units.is_multiple_of(num_units_in_ruda as usize)
46                    || working_rudas != launched_rudas;
47                if unit_idle && !blueprint.unit_idle.is_enabled() {
48                    return Err(ReduceError::Validation {
49                        details: "Too many units launched for the problem causing OOD, but `unit_idle` is off.",
50                    });
51                }
52
53                let blueprint = ReduceBlueprint {
54                    vectorization_mode: settings.vectorization_mode,
55                    global: GlobalReduceBlueprint::Unit(blueprint),
56                };
57
58                (blueprint, ruda_dim, ruda_count)
59            }
60            BlueprintStrategy::Inferred(_) => {
61                let (blueprint, ruda_dim, ruda_count) =
62                    generate_blueprint::<R>(client, problem, &settings)?;
63                (blueprint, ruda_dim, ruda_count)
64            }
65        };
66
67        let launch = ReduceLaunchSettings {
68            ruda_dim,
69            ruda_count,
70            vector: settings,
71            address_type,
72        };
73
74        Ok((blueprint, launch))
75    }
76}
77
78fn generate_blueprint<R: Runtime>(
79    client: &ComputeClient<R>,
80    problem: ReduceProblem,
81    settings: &ReduceVectorSettings,
82) -> Result<(ReduceBlueprint, RudaDim, RudaCount), ReduceError> {
83    let properties = &client.properties().hardware;
84    let plane_size = properties.plane_size_max;
85    let working_units = working_units(settings, &problem);
86    let plane_count = calculate_plane_count_per_ruda(working_units, plane_size, properties);
87
88    let ruda_dim = RudaDim::new_2d(plane_size, plane_count);
89    let num_units_in_ruda = ruda_dim.num_elems();
90
91    let working_rudas = working_units.div_ceil(num_units_in_ruda as usize);
92    let (ruda_count, ruda_launched) = ruda_count_spread_with_total(client, working_rudas);
93    let unit_idle =
94        !working_units.is_multiple_of(num_units_in_ruda as usize) || ruda_launched != working_rudas;
95
96    let unit_idle = match unit_idle {
97        true => IdleMode::Terminate,
98        false => IdleMode::None,
99    };
100    let blueprint = ReduceBlueprint {
101        vectorization_mode: settings.vectorization_mode,
102        global: GlobalReduceBlueprint::Unit(UnitReduceBlueprint { unit_idle }),
103    };
104
105    Ok((blueprint, ruda_dim, ruda_count))
106}
107
108fn working_units(settings: &ReduceVectorSettings, problem: &ReduceProblem) -> usize {
109    match settings.vectorization_mode {
110        VectorizationMode::Parallel => problem.reduce_count / settings.vector_size_output,
111        VectorizationMode::Perpendicular => problem.reduce_count / settings.vector_size_input,
112    }
113}