1use base64::Engine as _;
18use cortiq_core::SelectionDescriptor;
19use cortiq_core::quant::f16_to_f32;
20
21pub struct RoutableSkill {
23 pub idx: usize,
25 pub id: String,
26 pub phi_layer: usize,
27 mean: Vec<f32>,
28 basis: Vec<f32>,
29 rank: usize,
30}
31
32fn decode_f16(b64: &str) -> Option<Vec<f32>> {
33 let bytes = base64::engine::general_purpose::STANDARD.decode(b64).ok()?;
34 Some(
35 bytes
36 .chunks_exact(2)
37 .map(|c| f16_to_f32(u16::from_le_bytes([c[0], c[1]])))
38 .collect(),
39 )
40}
41
42impl RoutableSkill {
43 pub fn from_descriptor(
44 idx: usize,
45 id: String,
46 sel: &SelectionDescriptor,
47 hidden: usize,
48 ) -> Option<Self> {
49 if sel.metric != "mse" {
50 return None;
51 }
52 let mean = decode_f16(&sel.mean)?;
53 let basis = decode_f16(&sel.basis)?;
54 if mean.len() != hidden || basis.len() != sel.rank * hidden {
55 return None;
56 }
57 Some(Self {
58 idx,
59 id,
60 phi_layer: sel.phi_layer,
61 mean,
62 basis,
63 rank: sel.rank,
64 })
65 }
66
67 pub fn error(&self, phi: &[f32]) -> f32 {
70 let hidden = self.mean.len();
71 if phi.len() != hidden {
72 return f32::INFINITY;
73 }
74 let r: Vec<f32> = phi.iter().zip(&self.mean).map(|(p, m)| p - m).collect();
75 let rr: f32 = r.iter().map(|v| v * v).sum();
76 let pp: f32 = phi.iter().map(|v| v * v).sum();
77 let mut proj = 0f32;
78 for k in 0..self.rank {
79 let row = &self.basis[k * hidden..(k + 1) * hidden];
80 let c: f32 = row.iter().zip(&r).map(|(b, v)| b * v).sum();
81 proj += c * c;
82 }
83 (rr - proj).max(0.0) / pp.max(1e-12)
84 }
85}
86
87pub struct DynRouter {
89 pub skills: Vec<RoutableSkill>,
90 pub e_on: f32,
92 pub e_off: f32,
94 pub margin: f32,
96 pub period: usize,
98 active: Option<usize>,
100 tick: usize,
101 pub switches: Vec<(usize, Option<String>, Option<String>)>,
103 last_best_e: f32,
106}
107
108impl DynRouter {
109 pub fn new(skills: Vec<RoutableSkill>) -> Self {
110 let e_on = std::env::var("CMF_ROUTE_EON")
111 .ok()
112 .and_then(|v| v.parse().ok())
113 .unwrap_or(0.62);
114 let e_off = std::env::var("CMF_ROUTE_EOFF")
115 .ok()
116 .and_then(|v| v.parse().ok())
117 .unwrap_or(0.74);
118 let margin = std::env::var("CMF_ROUTE_MARGIN")
119 .ok()
120 .and_then(|v| v.parse().ok())
121 .unwrap_or(0.03);
122 let period = std::env::var("CMF_ROUTE_PERIOD")
123 .ok()
124 .and_then(|v| v.parse().ok())
125 .unwrap_or(8usize)
126 .max(1);
127 Self {
128 skills,
129 e_on,
130 e_off,
131 margin,
132 period,
133 active: None,
134 tick: 0,
135 switches: Vec::new(),
136 last_best_e: f32::INFINITY,
137 }
138 }
139
140 pub fn phi_layer(&self) -> Option<usize> {
143 self.skills.first().map(|s| s.phi_layer)
144 }
145
146 pub fn step(&mut self, phi: &[f32], token_no: usize) -> Option<Option<usize>> {
151 self.tick += 1;
152 if self.tick % self.period != 0 || phi.is_empty() || self.skills.is_empty() {
153 return None;
154 }
155 let mut best_idx = None;
157 let mut best_e = f32::INFINITY;
158 let mut active_e = f32::INFINITY;
159 for s in &self.skills {
160 let e = s.error(phi);
161 if Some(s.idx) == self.active {
162 active_e = e;
163 }
164 if e < best_e {
165 best_e = e;
166 best_idx = Some(s.idx);
167 }
168 }
169
170 self.last_best_e = best_e; let next = decide(
173 self.active,
174 active_e,
175 best_idx,
176 best_e,
177 self.e_on,
178 self.e_off,
179 self.margin,
180 );
181
182 if next != self.active {
183 let from = self
184 .active
185 .and_then(|i| self.skills.iter().find(|s| s.idx == i))
186 .map(|s| s.id.clone());
187 let to = next
188 .and_then(|i| self.skills.iter().find(|s| s.idx == i))
189 .map(|s| s.id.clone());
190 self.switches.push((token_no, from, to));
191 self.active = next;
192 return Some(next);
193 }
194 None
195 }
196
197 pub fn active(&self) -> Option<usize> {
198 self.active
199 }
200
201 pub fn active_id(&self) -> Option<String> {
203 self.active
204 .and_then(|i| self.skills.iter().find(|s| s.idx == i))
205 .map(|s| s.id.clone())
206 }
207
208 pub fn last_best_e(&self) -> f32 {
210 self.last_best_e
211 }
212
213 pub fn reset(&mut self) {
216 self.active = None;
217 self.tick = 0;
218 self.switches.clear();
219 self.last_best_e = f32::INFINITY;
220 }
221}
222
223#[allow(clippy::too_many_arguments)]
227pub fn decide(
228 active: Option<usize>,
229 active_e: f32,
230 best_idx: Option<usize>,
231 best_e: f32,
232 e_on: f32,
233 e_off: f32,
234 margin: f32,
235) -> Option<usize> {
236 match active {
237 None => {
239 if best_e < e_on {
240 best_idx
241 } else {
242 None
243 }
244 }
245 Some(cur) => {
246 if active_e > e_off {
247 if best_e < e_on { best_idx } else { None }
249 } else if best_idx != Some(cur) && best_e + margin < active_e {
250 best_idx
252 } else {
253 Some(cur)
254 }
255 }
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::decide;
262
263 #[test]
264 fn hysteresis_barrier_suppresses_thrashing() {
265 let (e_on, e_off, m) = (0.60, 0.75, 0.03);
266
267 assert_eq!(
269 decide(None, f32::INFINITY, Some(0), 0.70, e_on, e_off, m),
270 None
271 );
272 assert_eq!(
274 decide(None, f32::INFINITY, Some(0), 0.55, e_on, e_off, m),
275 Some(0)
276 );
277
278 assert_eq!(
281 decide(Some(0), 0.70, Some(1), 0.68, e_on, e_off, m),
282 Some(0)
283 );
284 assert_eq!(
286 decide(Some(0), 0.80, Some(1), 0.55, e_on, e_off, m),
287 Some(1)
288 );
289 assert_eq!(decide(Some(0), 0.80, Some(1), 0.70, e_on, e_off, m), None);
291 assert_eq!(
293 decide(Some(0), 0.70, Some(1), 0.69, e_on, e_off, m),
294 Some(0)
295 );
296 assert_eq!(
297 decide(Some(0), 0.70, Some(1), 0.66, e_on, e_off, m),
298 Some(1)
299 );
300 }
301}