Skip to main content

cortiq_engine/
router.rs

1//! Recon-argmin skill routing (spec §9, P1 signal-consistency): the
2//! container's selection descriptors define per-skill affine subspaces
3//! over φ(x); the winner is the skill that reconstructs φ best. No
4//! trained gate — routing is a property of the skills themselves.
5
6use 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    /// Normalized reconstruction error E ∈ [0, 1]; lower = closer.
15    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
28/// Score every routable skill; sorted best-first. Empty when the file
29/// carries no selection descriptors.
30pub 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        // r = φ − mean;  E = ‖r − B·Bᵀr‖² / ‖φ‖²  (B rows orthonormal).
58        // Normalizing by ‖φ‖ (not ‖r‖!) keeps the distance-to-mean
59        // signal — the whole point of the affine subspace.
60        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}