ruda_runtime/runtime/tune/
local.rs1use super::{AutotuneKey, AutotuneOutput, TunableSet, TuneInputs, Tuner};
2use crate::runtime::{client::ComputeClient, backend::Runtime, tune::TuneCacheResult};
3use alloc::string::ToString;
4use alloc::sync::Arc;
5use core::{
6 any::{Any, TypeId},
7 fmt::Display,
8 hash::Hash,
9};
10use hashbrown::HashMap;
11use spin::RwLock;
12
13type Sets = RwLock<Option<HashMap<TypeId, Arc<dyn Any + Send + Sync>>>>;
14
15pub struct LocalTuner<AK: AutotuneKey, ID> {
18 state: RwLock<Option<HashMap<ID, Arc<Tuner<AK>>>>>,
19 name: &'static str,
20 sets: Sets,
21 device_sets: RwLock<Option<HashMap<ID, Arc<Sets>>>>,
22}
23
24#[macro_export]
26macro_rules! local_tuner {
27 ($name:expr) => {
28 LocalTuner::new(concat!(module_path!(), "-", $name));
29 };
30 () => {
31 LocalTuner::new(module_path!());
32 };
33}
34
35pub use local_tuner;
36
37impl<AK, ID> LocalTuner<AK, ID>
38where
39 AK: AutotuneKey + 'static,
40 ID: Hash + PartialEq + Eq + Clone + Display,
41{
42 pub const fn new(name: &'static str) -> Self {
44 Self {
45 state: RwLock::new(None),
46 name,
47 sets: RwLock::new(None),
48 device_sets: RwLock::new(None),
49 }
50 }
51
52 pub fn init<I, Out, F>(&self, init_set: F) -> Arc<TunableSet<AK, I, Out>>
57 where
58 F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
59 I: TuneInputs,
60 Out: AutotuneOutput,
61 {
62 Self::init_set(&self.sets, init_set)
63 }
64
65 pub fn init_for_device<I, Out, F>(&self, id: &ID, init_set: F) -> Arc<TunableSet<AK, I, Out>>
67 where
68 F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
69 I: TuneInputs,
70 Out: AutotuneOutput,
71 {
72 let existing = self.device_sets.read().as_ref().and_then(|devices| devices.get(id)).cloned();
73 let sets = existing.unwrap_or_else(|| {
74 let mut devices = self.device_sets.write();
75 let devices = devices.get_or_insert_with(HashMap::new);
76 match devices.get(id) {
77 Some(sets) => sets.clone(),
78 None => devices.entry(id.clone())
79 .or_insert_with(|| Arc::new(RwLock::new(None)))
80 .clone(),
81 }
82 });
83 Self::init_set(&sets, init_set)
84 }
85
86 fn init_set<I, Out, F>(
87 sets: &Sets,
88 init_set: F,
89 ) -> Arc<TunableSet<AK, I, Out>>
90 where
91 F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
92 I: TuneInputs,
93 Out: AutotuneOutput,
94 {
95 let key = TypeId::of::<F>();
96 let read = sets.read();
97
98 static DOWNCAST_ERROR: &str = "Local tuner only support one set of tunable that must work on the same input and output declared with the init function.";
99
100 if let Some(sets) = read.as_ref()
101 && let Some(set) = sets.get(&key)
102 {
103 return set.clone().downcast().expect(DOWNCAST_ERROR);
104 };
105
106 core::mem::drop(read);
107
108 let mut sets = sets.write();
109
110 if let Some(sets) = sets.as_ref()
111 && let Some(set) = sets.get(&key)
112 {
113 return set.clone().downcast().expect(DOWNCAST_ERROR);
114 };
115
116 let content = Arc::new(init_set());
117
118 if let Some(sets) = sets.as_mut() {
119 sets.insert(key, content.clone());
120 } else {
121 let mut map = HashMap::<TypeId, Arc<dyn Any + Send + Sync>>::new();
122 map.insert(key, content.clone());
123 *sets = Some(map);
124 };
125
126 content
127 }
128
129 pub fn clear(&self) {
131 if let Some(s) = self.state.write().as_mut() {
132 s.clear()
133 }
134 }
135
136 #[cfg(feature = "runtime-autotune-checks")]
137 fn checks<'a, I: TuneInputs, Out: AutotuneOutput>(
138 &self,
139 operations: &TunableSet<AK, I, Out>,
140 inputs: &<I as TuneInputs>::At<'a>,
141 ) where
142 <I as TuneInputs>::At<'a>: Clone + Send,
143 {
144 use alloc::vec::Vec;
145
146 let mut checks_outputs = Vec::new();
147 for i in 0..operations.len() {
148 let op = operations.fastest(i);
149 let result = op.execute(inputs.clone());
150 checks_outputs.push(result);
151 }
152 super::check_autotune_outputs(checks_outputs);
153 }
154
155 pub fn execute<'a, R: Runtime, I: TuneInputs, Out>(
158 &self,
159 id: &ID,
160 client: &ComputeClient<R>,
161 operations: Arc<TunableSet<AK, I, Out>>,
162 inputs: <I as TuneInputs>::At<'a>,
163 ) -> Out
164 where
165 <I as TuneInputs>::At<'a>: Clone + Send,
166 Out: AutotuneOutput,
167 {
168 #[cfg(std_io)]
169 if super::stack::stack_autotuner().is_some() && operations.stack_reference().is_some() {
170 return super::stack::try_execute_stack(
171 self.name, &id.to_string(), client, operations, inputs,
172 ).expect("Full-stack autotune failed; request was not replayed");
173 }
174 let key = operations.generate_key(&inputs);
175
176 let existing = self.state.read().as_ref().and_then(|state| state.get(id)).cloned();
177 let tuner = existing.unwrap_or_else(|| {
178 let mut state = self.state.write();
179 let state = state.get_or_insert_with(HashMap::new);
180 match state.get(id) {
181 Some(tuner) => tuner.clone(),
182 None => state.entry(id.clone())
183 .or_insert_with(|| {
184 let name = self.name.replace("::", "-");
185 Arc::new(Tuner::new(&name, &id.to_string()))
186 })
187 .clone(),
188 }
189 });
190
191 if let TuneCacheResult::Hit { fastest_index } = tuner.fastest(&key) {
193 #[cfg(feature = "runtime-autotune-checks")]
194 self.checks::<I, Out>(&operations, &inputs);
195 return operations
196 .fastest(fastest_index)
197 .execute(inputs)
198 .expect("Should run when selected by autotune.");
199 }
200
201 let fastest = tuner.check_tune::<R, I, Out>(
202 &key,
203 &inputs,
204 &operations,
205 || operations.compute_checksum(),
206 client,
207 );
208
209 match fastest {
211 TuneCacheResult::Hit { fastest_index } => {
212 #[cfg(feature = "runtime-autotune-checks")]
213 self.checks::<I, Out>(&operations, &inputs);
214
215 operations
216 .fastest(fastest_index)
217 .execute(inputs)
218 .expect("Should run when selected by autotune.")
219 }
220 TuneCacheResult::Unchecked | TuneCacheResult::Miss => {
221 panic!(
222 "Somehow we STILL didn't check a tuning checksum or start tuning, something has gone wrong."
223 )
224 }
225 TuneCacheResult::Pending => {
226 for i in 0..operations.len() {
228 if let Ok(output) = operations.fastest(i).execute(inputs.clone()) {
229 return output;
230 }
231 }
232 panic!("All autotune operations failed, no viable operation found.");
233 }
234 }
235 }
236}