1use base64::Engine as _;
16use cortiq_core::SelectionDescriptor;
17use cortiq_core::quant::f16_to_f32;
18
19pub struct RoutableSkill {
21 pub idx: usize,
23 pub id: String,
24 pub phi_layer: usize,
25 mean: Vec<f32>,
26 basis: Vec<f32>,
27 rank: usize,
28 unit: bool,
31}
32
33fn decode_f16(b64: &str) -> Option<Vec<f32>> {
34 let bytes = base64::engine::general_purpose::STANDARD.decode(b64).ok()?;
35 Some(
36 bytes
37 .chunks_exact(2)
38 .map(|c| f16_to_f32(u16::from_le_bytes([c[0], c[1]])))
39 .collect(),
40 )
41}
42
43impl RoutableSkill {
44 pub fn from_descriptor(
45 idx: usize,
46 id: String,
47 sel: &SelectionDescriptor,
48 hidden: usize,
49 ) -> Option<Self> {
50 let unit = match sel.metric.as_str() {
51 "mse" => false,
52 "mse_unit" => true,
53 _ => return None,
54 };
55 let mean = decode_f16(&sel.mean)?;
56 let basis = decode_f16(&sel.basis)?;
57 if mean.len() != hidden || basis.len() != sel.rank * hidden {
58 return None;
59 }
60 Some(Self {
61 idx,
62 id,
63 phi_layer: sel.phi_layer,
64 mean,
65 basis,
66 rank: sel.rank,
67 unit,
68 })
69 }
70
71 pub fn error(&self, phi: &[f32]) -> f32 {
75 let hidden = self.mean.len();
76 if phi.len() != hidden {
77 return f32::INFINITY;
78 }
79 let norm2: f32 = phi.iter().map(|v| v * v).sum();
80 let scale = if self.unit && norm2 > 0.0 {
81 1.0 / norm2.sqrt()
82 } else {
83 1.0
84 };
85 let r: Vec<f32> = phi
86 .iter()
87 .zip(&self.mean)
88 .map(|(p, m)| p * scale - m)
89 .collect();
90 let rr: f32 = r.iter().map(|v| v * v).sum();
91 let pp: f32 = norm2 * scale * scale;
92 let mut proj = 0f32;
93 for k in 0..self.rank {
94 let row = &self.basis[k * hidden..(k + 1) * hidden];
95 let c: f32 = row.iter().zip(&r).map(|(b, v)| b * v).sum();
96 proj += c * c;
97 }
98 (rr - proj).max(0.0) / pp.max(1e-12)
99 }
100}
101
102pub struct DynRouter {
104 pub skills: Vec<RoutableSkill>,
105 pub e_on: f32,
107 pub e_off: f32,
109 pub margin: f32,
111 pub period: usize,
113 active: Option<usize>,
115 tick: usize,
116 pub switches: Vec<(usize, Option<String>, Option<String>)>,
118 last_best_e: f32,
121}
122
123impl DynRouter {
124 pub fn new(skills: Vec<RoutableSkill>) -> Self {
125 let e_on = std::env::var("CMF_ROUTE_EON")
126 .ok()
127 .and_then(|v| v.parse().ok())
128 .unwrap_or(0.62);
129 let e_off = std::env::var("CMF_ROUTE_EOFF")
130 .ok()
131 .and_then(|v| v.parse().ok())
132 .unwrap_or(0.74);
133 let margin = std::env::var("CMF_ROUTE_MARGIN")
134 .ok()
135 .and_then(|v| v.parse().ok())
136 .unwrap_or(0.03);
137 let period = std::env::var("CMF_ROUTE_PERIOD")
138 .ok()
139 .and_then(|v| v.parse().ok())
140 .unwrap_or(8usize)
141 .max(1);
142 Self {
143 skills,
144 e_on,
145 e_off,
146 margin,
147 period,
148 active: None,
149 tick: 0,
150 switches: Vec::new(),
151 last_best_e: f32::INFINITY,
152 }
153 }
154
155 pub fn phi_layer(&self) -> Option<usize> {
158 self.skills.first().map(|s| s.phi_layer)
159 }
160
161 pub fn step(&mut self, phi: &[f32], token_no: usize) -> Option<Option<usize>> {
166 self.tick += 1;
167 if self.tick % self.period != 0 || phi.is_empty() || self.skills.is_empty() {
168 return None;
169 }
170 let mut best_idx = None;
172 let mut best_e = f32::INFINITY;
173 let mut active_e = f32::INFINITY;
174 for s in &self.skills {
175 let e = s.error(phi);
176 if Some(s.idx) == self.active {
177 active_e = e;
178 }
179 if e < best_e {
180 best_e = e;
181 best_idx = Some(s.idx);
182 }
183 }
184
185 self.last_best_e = best_e; let next = decide(
188 self.active,
189 active_e,
190 best_idx,
191 best_e,
192 self.e_on,
193 self.e_off,
194 self.margin,
195 );
196
197 if next != self.active {
198 let from = self
199 .active
200 .and_then(|i| self.skills.iter().find(|s| s.idx == i))
201 .map(|s| s.id.clone());
202 let to = next
203 .and_then(|i| self.skills.iter().find(|s| s.idx == i))
204 .map(|s| s.id.clone());
205 self.switches.push((token_no, from, to));
206 self.active = next;
207 return Some(next);
208 }
209 None
210 }
211
212 pub fn active(&self) -> Option<usize> {
213 self.active
214 }
215
216 pub fn active_id(&self) -> Option<String> {
218 self.active
219 .and_then(|i| self.skills.iter().find(|s| s.idx == i))
220 .map(|s| s.id.clone())
221 }
222
223 pub fn last_best_e(&self) -> f32 {
225 self.last_best_e
226 }
227
228 pub fn reset(&mut self) {
231 self.active = None;
232 self.tick = 0;
233 self.switches.clear();
234 self.last_best_e = f32::INFINITY;
235 }
236}
237
238#[allow(clippy::too_many_arguments)]
242pub fn decide(
243 active: Option<usize>,
244 active_e: f32,
245 best_idx: Option<usize>,
246 best_e: f32,
247 e_on: f32,
248 e_off: f32,
249 margin: f32,
250) -> Option<usize> {
251 match active {
252 None => {
254 if best_e < e_on {
255 best_idx
256 } else {
257 None
258 }
259 }
260 Some(cur) => {
261 if active_e > e_off {
262 if best_e < e_on { best_idx } else { None }
264 } else if best_idx != Some(cur) && best_e + margin < active_e {
265 best_idx
267 } else {
268 Some(cur)
269 }
270 }
271 }
272}
273
274#[cfg(test)]
275mod tests {
276 use super::{RoutableSkill, decide};
277 use base64::Engine as _;
278
279 fn f16_b64(v: &[f32]) -> String {
280 let bytes: Vec<u8> = v
281 .iter()
282 .flat_map(|x| cortiq_core::quant::f32_to_f16(*x).to_le_bytes())
283 .collect();
284 base64::engine::general_purpose::STANDARD.encode(bytes)
285 }
286
287 #[test]
291 fn mse_unit_descriptors_are_routable_and_scale_invariant() {
292 let hidden = 4;
293 let mut mean = vec![0.0f32; hidden];
294 mean[0] = 1.0;
295 let mut basis = vec![0.0f32; hidden];
296 basis[1] = 1.0;
297 let sel = |metric: &str| cortiq_core::SelectionDescriptor {
298 metric: metric.into(),
299 phi_layer: 2,
300 mean: f16_b64(&mean),
301 basis: f16_b64(&basis),
302 rank: 1,
303 err_mean: None,
304 err_std: None,
305 holdout: None,
306 holdout_n: None,
307 };
308 let unit = RoutableSkill::from_descriptor(0, "u".into(), &sel("mse_unit"), hidden)
309 .expect("mse_unit accepted");
310 let raw = RoutableSkill::from_descriptor(1, "r".into(), &sel("mse"), hidden)
311 .expect("mse accepted");
312 assert!(RoutableSkill::from_descriptor(2, "x".into(), &sel("cosine"), hidden).is_none());
313 for s in [0.5f32, 1.0, 7.0] {
315 let phi = vec![s, 0.0, 0.0, 0.0];
316 assert!(unit.error(&phi) < 1e-6, "scale {s}: {}", unit.error(&phi));
317 }
318 assert!(raw.error(&[7.0, 0.0, 0.0, 0.0]) > 0.5);
320 assert!(unit.error(&[3.0, 1.0, 0.0, 0.0]) < 0.01);
322 assert!(unit.error(&[0.0, 0.0, 1.0, 0.0]) > 0.9);
323 }
324
325 #[test]
326 fn hysteresis_barrier_suppresses_thrashing() {
327 let (e_on, e_off, m) = (0.60, 0.75, 0.03);
328
329 assert_eq!(
331 decide(None, f32::INFINITY, Some(0), 0.70, e_on, e_off, m),
332 None
333 );
334 assert_eq!(
336 decide(None, f32::INFINITY, Some(0), 0.55, e_on, e_off, m),
337 Some(0)
338 );
339
340 assert_eq!(
343 decide(Some(0), 0.70, Some(1), 0.68, e_on, e_off, m),
344 Some(0)
345 );
346 assert_eq!(
348 decide(Some(0), 0.80, Some(1), 0.55, e_on, e_off, m),
349 Some(1)
350 );
351 assert_eq!(decide(Some(0), 0.80, Some(1), 0.70, e_on, e_off, m), None);
353 assert_eq!(
355 decide(Some(0), 0.70, Some(1), 0.69, e_on, e_off, m),
356 Some(0)
357 );
358 assert_eq!(
359 decide(Some(0), 0.70, Some(1), 0.66, e_on, e_off, m),
360 Some(1)
361 );
362 }
363}