1use crate::pipeline::Pipeline;
7use base64::Engine as _;
8use cortiq_core::CmfModel;
9use cortiq_core::quant::f16_to_f32;
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 {
37 continue;
38 };
39 if sel.metric != "mse" {
40 tracing::warn!("skill '{}': unknown metric '{}'", skill.id, sel.metric);
41 continue;
42 }
43 let phi = match phi_cache.iter().find(|(l, _)| *l == sel.phi_layer) {
44 Some((_, p)) => p.clone(),
45 None => {
46 let p = pipeline.probe_phi(ids, sel.phi_layer);
47 phi_cache.push((sel.phi_layer, p.clone()));
48 p
49 }
50 };
51 let (Some(mean), Some(basis)) = (decode_f16(&sel.mean), decode_f16(&sel.basis)) else {
52 tracing::error!("skill '{}': malformed selection payload", skill.id);
53 continue;
54 };
55 if mean.len() != hidden || basis.len() != sel.rank * hidden {
56 tracing::error!("skill '{}': selection dims mismatch", skill.id);
57 continue;
58 }
59 let r: Vec<f32> = phi.iter().zip(&mean).map(|(p, m)| p - m).collect();
63 let rr: f32 = r.iter().map(|v| v * v).sum();
64 let pp: f32 = phi.iter().map(|v| v * v).sum();
65 let mut proj = 0f32;
66 for k in 0..sel.rank {
67 let row = &basis[k * hidden..(k + 1) * hidden];
68 let c: f32 = row.iter().zip(&r).map(|(b, v)| b * v).sum();
69 proj += c * c;
70 }
71 let e = (rr - proj).max(0.0) / pp.max(1e-12);
72 out.push(SkillRoute {
73 id: skill.id.clone(),
74 error: e,
75 });
76 }
77 out.sort_by(|a, b| a.error.total_cmp(&b.error));
78 out
79}