use base64::Engine as _;
use cortiq_core::SelectionDescriptor;
use cortiq_core::quant::f16_to_f32;
pub struct RoutableSkill {
pub idx: usize,
pub id: String,
pub phi_layer: usize,
mean: Vec<f32>,
basis: Vec<f32>,
rank: usize,
}
fn decode_f16(b64: &str) -> Option<Vec<f32>> {
let bytes = base64::engine::general_purpose::STANDARD.decode(b64).ok()?;
Some(
bytes
.chunks_exact(2)
.map(|c| f16_to_f32(u16::from_le_bytes([c[0], c[1]])))
.collect(),
)
}
impl RoutableSkill {
pub fn from_descriptor(
idx: usize,
id: String,
sel: &SelectionDescriptor,
hidden: usize,
) -> Option<Self> {
if sel.metric != "mse" {
return None;
}
let mean = decode_f16(&sel.mean)?;
let basis = decode_f16(&sel.basis)?;
if mean.len() != hidden || basis.len() != sel.rank * hidden {
return None;
}
Some(Self {
idx,
id,
phi_layer: sel.phi_layer,
mean,
basis,
rank: sel.rank,
})
}
pub fn error(&self, phi: &[f32]) -> f32 {
let hidden = self.mean.len();
if phi.len() != hidden {
return f32::INFINITY;
}
let r: Vec<f32> = phi.iter().zip(&self.mean).map(|(p, m)| p - m).collect();
let rr: f32 = r.iter().map(|v| v * v).sum();
let pp: f32 = phi.iter().map(|v| v * v).sum();
let mut proj = 0f32;
for k in 0..self.rank {
let row = &self.basis[k * hidden..(k + 1) * hidden];
let c: f32 = row.iter().zip(&r).map(|(b, v)| b * v).sum();
proj += c * c;
}
(rr - proj).max(0.0) / pp.max(1e-12)
}
}
pub struct DynRouter {
pub skills: Vec<RoutableSkill>,
pub e_on: f32,
pub e_off: f32,
pub margin: f32,
pub period: usize,
active: Option<usize>,
tick: usize,
pub switches: Vec<(usize, Option<String>, Option<String>)>,
last_best_e: f32,
}
impl DynRouter {
pub fn new(skills: Vec<RoutableSkill>) -> Self {
let e_on = std::env::var("CMF_ROUTE_EON")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.62);
let e_off = std::env::var("CMF_ROUTE_EOFF")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.74);
let margin = std::env::var("CMF_ROUTE_MARGIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.03);
let period = std::env::var("CMF_ROUTE_PERIOD")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(8usize)
.max(1);
Self {
skills,
e_on,
e_off,
margin,
period,
active: None,
tick: 0,
switches: Vec::new(),
last_best_e: f32::INFINITY,
}
}
pub fn phi_layer(&self) -> Option<usize> {
self.skills.first().map(|s| s.phi_layer)
}
pub fn step(&mut self, phi: &[f32], token_no: usize) -> Option<Option<usize>> {
self.tick += 1;
if self.tick % self.period != 0 || phi.is_empty() || self.skills.is_empty() {
return None;
}
let mut best_idx = None;
let mut best_e = f32::INFINITY;
let mut active_e = f32::INFINITY;
for s in &self.skills {
let e = s.error(phi);
if Some(s.idx) == self.active {
active_e = e;
}
if e < best_e {
best_e = e;
best_idx = Some(s.idx);
}
}
self.last_best_e = best_e;
let next = decide(
self.active,
active_e,
best_idx,
best_e,
self.e_on,
self.e_off,
self.margin,
);
if next != self.active {
let from = self
.active
.and_then(|i| self.skills.iter().find(|s| s.idx == i))
.map(|s| s.id.clone());
let to = next
.and_then(|i| self.skills.iter().find(|s| s.idx == i))
.map(|s| s.id.clone());
self.switches.push((token_no, from, to));
self.active = next;
return Some(next);
}
None
}
pub fn active(&self) -> Option<usize> {
self.active
}
pub fn active_id(&self) -> Option<String> {
self.active
.and_then(|i| self.skills.iter().find(|s| s.idx == i))
.map(|s| s.id.clone())
}
pub fn last_best_e(&self) -> f32 {
self.last_best_e
}
pub fn reset(&mut self) {
self.active = None;
self.tick = 0;
self.switches.clear();
self.last_best_e = f32::INFINITY;
}
}
#[allow(clippy::too_many_arguments)]
pub fn decide(
active: Option<usize>,
active_e: f32,
best_idx: Option<usize>,
best_e: f32,
e_on: f32,
e_off: f32,
margin: f32,
) -> Option<usize> {
match active {
None => {
if best_e < e_on {
best_idx
} else {
None
}
}
Some(cur) => {
if active_e > e_off {
if best_e < e_on { best_idx } else { None }
} else if best_idx != Some(cur) && best_e + margin < active_e {
best_idx
} else {
Some(cur)
}
}
}
}
#[cfg(test)]
mod tests {
use super::decide;
#[test]
fn hysteresis_barrier_suppresses_thrashing() {
let (e_on, e_off, m) = (0.60, 0.75, 0.03);
assert_eq!(
decide(None, f32::INFINITY, Some(0), 0.70, e_on, e_off, m),
None
);
assert_eq!(
decide(None, f32::INFINITY, Some(0), 0.55, e_on, e_off, m),
Some(0)
);
assert_eq!(
decide(Some(0), 0.70, Some(1), 0.68, e_on, e_off, m),
Some(0)
);
assert_eq!(
decide(Some(0), 0.80, Some(1), 0.55, e_on, e_off, m),
Some(1)
);
assert_eq!(decide(Some(0), 0.80, Some(1), 0.70, e_on, e_off, m), None);
assert_eq!(
decide(Some(0), 0.70, Some(1), 0.69, e_on, e_off, m),
Some(0)
);
assert_eq!(
decide(Some(0), 0.70, Some(1), 0.66, e_on, e_off, m),
Some(1)
);
}
}