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::{Mutex, RwLock};
12
13type Sets = RwLock<Option<HashMap<TypeId, Arc<dyn Any + Send + Sync>>>>;
14
15pub struct LocalTuner<AK: AutotuneKey, ID> {
18 state: Mutex<Option<HashMap<ID, Arc<Tuner<AK>>>>>,
19 name: &'static str,
20 sets: Sets,
21 device_sets: Mutex<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: Mutex::new(None),
46 name,
47 sets: RwLock::new(None),
48 device_sets: Mutex::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 sets = {
73 let mut devices = self.device_sets.lock();
74 let devices = devices.get_or_insert_with(HashMap::new);
75 match devices.get(id) {
76 Some(sets) => sets.clone(),
77 None => devices.entry(id.clone())
78 .or_insert_with(|| Arc::new(RwLock::new(None)))
79 .clone(),
80 }
81 };
82 Self::init_set(&sets, init_set)
83 }
84
85 fn init_set<I, Out, F>(
86 sets: &Sets,
87 init_set: F,
88 ) -> Arc<TunableSet<AK, I, Out>>
89 where
90 F: Fn() -> TunableSet<AK, I, Out> + 'static + Send + Sync,
91 I: TuneInputs,
92 Out: AutotuneOutput,
93 {
94 let key = TypeId::of::<F>();
95 let read = sets.read();
96
97 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.";
98
99 if let Some(sets) = read.as_ref()
100 && let Some(set) = sets.get(&key)
101 {
102 return set.clone().downcast().expect(DOWNCAST_ERROR);
103 };
104
105 core::mem::drop(read);
106
107 let mut sets = sets.write();
108
109 if let Some(sets) = sets.as_ref()
110 && let Some(set) = sets.get(&key)
111 {
112 return set.clone().downcast().expect(DOWNCAST_ERROR);
113 };
114
115 let content = Arc::new(init_set());
116
117 if let Some(sets) = sets.as_mut() {
118 sets.insert(key, content.clone());
119 } else {
120 let mut map = HashMap::<TypeId, Arc<dyn Any + Send + Sync>>::new();
121 map.insert(key, content.clone());
122 *sets = Some(map);
123 };
124
125 content
126 }
127
128 pub fn clear(&self) {
130 if let Some(s) = self.state.lock().as_mut() {
131 s.clear()
132 }
133 }
134
135 #[cfg(feature = "runtime-autotune-checks")]
136 fn checks<'a, I: TuneInputs, Out: AutotuneOutput>(
137 &self,
138 operations: &TunableSet<AK, I, Out>,
139 inputs: &<I as TuneInputs>::At<'a>,
140 ) where
141 <I as TuneInputs>::At<'a>: Clone + Send,
142 {
143 use alloc::vec::Vec;
144
145 let mut checks_outputs = Vec::new();
146 for i in 0..operations.len() {
147 let op = operations.fastest(i);
148 let result = op.execute(inputs.clone());
149 checks_outputs.push(result);
150 }
151 super::check_autotune_outputs(checks_outputs);
152 }
153
154 pub fn execute<'a, R: Runtime, I: TuneInputs, Out>(
157 &self,
158 id: &ID,
159 client: &ComputeClient<R>,
160 operations: Arc<TunableSet<AK, I, Out>>,
161 inputs: <I as TuneInputs>::At<'a>,
162 ) -> Out
163 where
164 <I as TuneInputs>::At<'a>: Clone + Send,
165 Out: AutotuneOutput,
166 {
167 #[cfg(std_io)]
168 if super::stack::stack_autotuner().is_some() && operations.stack_reference().is_some() {
169 return super::stack::try_execute_stack(
170 self.name, &id.to_string(), client, operations, inputs,
171 ).expect("Full-stack autotune failed; request was not replayed");
172 }
173 let key = operations.generate_key(&inputs);
174
175 let tuner = {
176 let mut state = self.state.lock();
177 let state = state.get_or_insert_with(HashMap::new);
178 match state.get(id) {
179 Some(tuner) => tuner.clone(),
180 None => state.entry(id.clone())
181 .or_insert_with(|| {
182 let name = self.name.replace("::", "-");
183 Arc::new(Tuner::new(&name, &id.to_string()))
184 })
185 .clone(),
186 }
187 };
188
189 if let TuneCacheResult::Hit { fastest_index } = tuner.fastest(&key) {
191 #[cfg(feature = "runtime-autotune-checks")]
192 self.checks::<I, Out>(&operations, &inputs);
193 return operations
194 .fastest(fastest_index)
195 .execute(inputs)
196 .expect("Should run when selected by autotune.");
197 }
198
199 let fastest = tuner.check_tune::<R, I, Out>(
200 &key,
201 &inputs,
202 &operations,
203 || operations.compute_checksum(),
204 client,
205 );
206
207 match fastest {
209 TuneCacheResult::Hit { fastest_index } => {
210 #[cfg(feature = "runtime-autotune-checks")]
211 self.checks::<I, Out>(&operations, &inputs);
212
213 operations
214 .fastest(fastest_index)
215 .execute(inputs)
216 .expect("Should run when selected by autotune.")
217 }
218 TuneCacheResult::Unchecked | TuneCacheResult::Miss => {
219 panic!(
220 "Somehow we STILL didn't check a tuning checksum or start tuning, something has gone wrong."
221 )
222 }
223 TuneCacheResult::Pending => {
224 for i in 0..operations.len() {
226 if let Ok(output) = operations.fastest(i).execute(inputs.clone()) {
227 return output;
228 }
229 }
230 panic!("All autotune operations failed, no viable operation found.");
231 }
232 }
233 }
234}