ruprim/reduce/routines/
unit.rs1use 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}