1use crate::pipeline::Pipeline;
7use base64::Engine as _;
8use cortiq_core::quant::f16_to_f32;
9use cortiq_core::CmfModel;
10
11#[derive(Debug, Clone)]
12pub struct SkillRoute {
13 pub id: String,
14 pub error: f32,
16}
17
18fn decode_f16(b64: &str) -> Option<Vec<f32>> {
19 let bytes = base64::engine::general_purpose::STANDARD.decode(b64).ok()?;
20 Some(
21 bytes
22 .chunks_exact(2)
23 .map(|c| f16_to_f32(u16::from_le_bytes([c[0], c[1]])))
24 .collect(),
25 )
26}
27
28pub fn route(model: &CmfModel, pipeline: &mut Pipeline, ids: &[u32]) -> Vec<SkillRoute> {
31 let hidden = model.arch().hidden_size;
32 let mut phi_cache: Vec<(usize, Vec<f32>)> = Vec::new();
33 let mut out = Vec::new();
34
35 for skill in &model.header.skills {
36 let Some(sel) = &skill.selection else { continue };
37 if sel.metric != "mse" {
38 tracing::warn!("skill '{}': unknown metric '{}'", skill.id, sel.metric);
39 continue;
40 }
41 let phi = match phi_cache.iter().find(|(l, _)| *l == sel.phi_layer) {
42 Some((_, p)) => p.clone(),
43 None => {
44 let p = pipeline.probe_phi(ids, sel.phi_layer);
45 phi_cache.push((sel.phi_layer, p.clone()));
46 p
47 }
48 };
49 let (Some(mean), Some(basis)) = (decode_f16(&sel.mean), decode_f16(&sel.basis)) else {
50 tracing::error!("skill '{}': malformed selection payload", skill.id);
51 continue;
52 };
53 if mean.len() != hidden || basis.len() != sel.rank * hidden {
54 tracing::error!("skill '{}': selection dims mismatch", skill.id);
55 continue;
56 }
57 let r: Vec<f32> = phi.iter().zip(&mean).map(|(p, m)| p - m).collect();
61 let rr: f32 = r.iter().map(|v| v * v).sum();
62 let pp: f32 = phi.iter().map(|v| v * v).sum();
63 let mut proj = 0f32;
64 for k in 0..sel.rank {
65 let row = &basis[k * hidden..(k + 1) * hidden];
66 let c: f32 = row.iter().zip(&r).map(|(b, v)| b * v).sum();
67 proj += c * c;
68 }
69 let e = (rr - proj).max(0.0) / pp.max(1e-12);
70 out.push(SkillRoute {
71 id: skill.id.clone(),
72 error: e,
73 });
74 }
75 out.sort_by(|a, b| a.error.total_cmp(&b.error));
76 out
77}