1use crate::pipeline::Pipeline;
24use base64::Engine as _;
25use cortiq_core::format::{CmfHeader, RoutingCalibration, SelectionDescriptor, SkillRecord};
26use cortiq_core::knowledge::{BACKBONE_CLASS_ID, METRIC_MSE_UNIT, parse_hex64};
27use cortiq_core::quant::{f16_to_f32, f32_to_f16};
28use cortiq_core::{CmfModel, PhiSpec, RouterPolicy, hash64};
29
30pub const NOVELTY_W_ENERGY: f32 = 0.5;
32pub const NOVELTY_W_MARGIN: f32 = 0.25;
33pub const NOVELTY_W_CONF: f32 = 0.25;
34pub const NOVELTY_MARGIN_K: f32 = 8.0;
35
36#[derive(Debug, Clone)]
37pub struct SkillRoute {
38 pub id: String,
39 pub error: f32,
41 pub raw_error: f32,
43 pub probability: f32,
46}
47
48#[derive(Debug, Clone)]
50pub struct Routing {
51 pub scores: Vec<SkillRoute>,
53 pub confidence: f32,
55 pub margin: f32,
57 pub novelty: f32,
59 pub is_novel: bool,
61 pub calibrated: bool,
62}
63
64impl Routing {
65 pub fn winner(&self) -> Option<&SkillRoute> {
66 self.scores.first()
67 }
68}
69
70pub fn decode_f16(b64: &str) -> Option<Vec<f32>> {
71 let bytes = base64::engine::general_purpose::STANDARD.decode(b64).ok()?;
72 Some(
73 bytes
74 .chunks_exact(2)
75 .map(|c| f16_to_f32(u16::from_le_bytes([c[0], c[1]])))
76 .collect(),
77 )
78}
79
80fn sigmoid(x: f32) -> f32 {
81 1.0 / (1.0 + (-x).exp())
82}
83
84fn softmax(v: &mut [f32]) {
85 if v.is_empty() {
86 return;
87 }
88 let mx = v.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
89 let mut s = 0.0f32;
90 for x in v.iter_mut() {
91 *x = (*x - mx).exp();
92 s += *x;
93 }
94 for x in v.iter_mut() {
95 *x /= s.max(1e-30);
96 }
97}
98
99pub fn recon_error(phi: &[f32], mean: &[f32], basis: &[f32], rank: usize) -> f32 {
101 let hidden = phi.len();
102 let r: Vec<f32> = phi.iter().zip(mean).map(|(p, m)| p - m).collect();
103 let rr: f32 = r.iter().map(|v| v * v).sum();
104 let mut proj = 0f32;
105 for k in 0..rank {
106 let row = &basis[k * hidden..(k + 1) * hidden];
107 let c: f32 = row.iter().zip(&r).map(|(b, v)| b * v).sum();
108 proj += c * c;
109 }
110 (rr - proj).max(0.0)
111}
112
113pub fn decide(
116 rows: &[ErrorRow],
117 calib: Option<&cortiq_core::format::RoutingCalibration>,
118 fallback_tau: f32,
119) -> Routing {
120 let mut idx: Vec<usize> = (0..rows.len()).collect();
121 idx.sort_by(|&a, &b| rows[a].1.total_cmp(&rows[b].1));
122 let mut scores: Vec<SkillRoute> = idx
123 .iter()
124 .map(|&i| SkillRoute {
125 id: rows[i].0.clone(),
126 error: rows[i].1 / rows[i].4.max(1e-12),
127 raw_error: rows[i].1,
128 probability: 0.0,
129 })
130 .collect();
131 if scores.is_empty() {
132 return Routing {
133 scores,
134 confidence: 0.0,
135 margin: 0.0,
136 novelty: 1.0,
137 is_novel: true,
138 calibrated: calib.is_some(),
139 };
140 }
141 let (Some(c), Some(&top)) = (calib, idx.first()) else {
142 let e_min = scores[0].error;
143 return Routing {
144 scores,
145 confidence: 0.0,
146 margin: 0.0,
147 novelty: f32::NAN,
148 is_novel: e_min > fallback_tau,
149 calibrated: false,
150 };
151 };
152 let mut logits: Vec<f32> = idx
154 .iter()
155 .map(|&i| -rows[i].1 / c.temperature.max(1e-3))
156 .collect();
157 softmax(&mut logits);
158 for (s, p) in scores.iter_mut().zip(&logits) {
159 s.probability = *p;
160 }
161 let confidence = logits[0];
162 let inv = |e: f32| 1.0 / (1.0 + e);
164 let margin = if idx.len() > 1 {
165 inv(rows[idx[0]].1) - inv(rows[idx[1]].1)
166 } else {
167 inv(rows[idx[0]].1)
168 };
169 let (em, es) = (
171 rows[top].2.unwrap_or(0.0),
172 rows[top].3.unwrap_or(1.0).max(1e-4),
173 );
174 let z = (rows[top].1 - em) / es;
175 let novelty = NOVELTY_W_ENERGY * sigmoid(z)
176 + NOVELTY_W_MARGIN / (1.0 + margin * NOVELTY_MARGIN_K)
177 + NOVELTY_W_CONF * (1.0 - confidence);
178 Routing {
179 scores,
180 confidence,
181 margin,
182 novelty,
183 is_novel: novelty > c.novelty_theta,
184 calibrated: true,
185 }
186}
187
188pub fn error_rows(
190 model: &CmfModel,
191 phi_of_layer: &mut dyn FnMut(usize) -> Vec<f32>,
192) -> Vec<ErrorRow> {
193 let hidden = model.arch().hidden_size;
194 let mut rows = Vec::new();
195 for skill in &model.header.skills {
196 if skill.is_v2() {
200 continue;
201 }
202 let Some(sel) = &skill.selection else {
203 continue;
204 };
205 let unit = match sel.metric.as_str() {
206 "mse" => false,
207 "mse_unit" => true,
208 m => {
209 tracing::warn!("skill '{}': unknown metric '{}'", skill.id, m);
210 continue;
211 }
212 };
213 let mut phi = phi_of_layer(sel.phi_layer);
214 if unit {
215 let n = phi.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-12);
216 for x in phi.iter_mut() {
217 *x /= n;
218 }
219 }
220 let (Some(mean), Some(basis)) = (decode_f16(&sel.mean), decode_f16(&sel.basis)) else {
221 tracing::error!("skill '{}': malformed selection payload", skill.id);
222 continue;
223 };
224 if mean.len() != hidden || basis.len() != sel.rank * hidden || phi.len() != hidden {
225 tracing::error!("skill '{}': selection dims mismatch", skill.id);
226 continue;
227 }
228 let e = recon_error(&phi, &mean, &basis, sel.rank);
229 let pp: f32 = phi.iter().map(|v| v * v).sum();
230 rows.push((skill.id.clone(), e, sel.err_mean, sel.err_std, pp));
231 }
232 rows
233}
234
235pub fn route_full(
237 model: &CmfModel,
238 pipeline: &mut Pipeline,
239 ids: &[u32],
240 fallback_tau: f32,
241) -> Routing {
242 let mut phi_cache: Vec<(usize, Vec<f32>)> = Vec::new();
243 let mut phi_of = |layer: usize| -> Vec<f32> {
244 if let Some((_, p)) = phi_cache.iter().find(|(l, _)| *l == layer) {
245 return p.clone();
246 }
247 let p = pipeline.probe_phi(ids, layer);
248 phi_cache.push((layer, p.clone()));
249 p
250 };
251 let rows = error_rows(model, &mut phi_of);
252 decide(&rows, model.header.routing.as_ref(), fallback_tau)
253}
254
255pub fn route(model: &CmfModel, pipeline: &mut Pipeline, ids: &[u32]) -> Vec<SkillRoute> {
257 route_full(model, pipeline, ids, 0.30).scores
258}
259
260pub fn holdout_phis(model: &CmfModel) -> Vec<(usize, Vec<f32>)> {
263 let hidden = model.arch().hidden_size;
264 let mut out = Vec::new();
265 for (si, skill) in model.header.skills.iter().enumerate() {
266 let Some(sel) = &skill.selection else {
267 continue;
268 };
269 let (Some(h), Some(n)) = (sel.holdout.as_ref(), sel.holdout_n) else {
270 continue;
271 };
272 let Some(v) = decode_f16(h) else { continue };
273 if v.len() != n * hidden {
274 continue;
275 }
276 for i in 0..n {
277 out.push((si, v[i * hidden..(i + 1) * hidden].to_vec()));
278 }
279 }
280 out
281}
282
283pub fn calibrate(
290 model: &CmfModel,
291 target_fpr: f32,
292) -> Option<cortiq_core::format::RoutingCalibration> {
293 let samples = holdout_phis(model);
294 if samples.is_empty() {
295 return None;
296 }
297 let skills = &model.header.skills;
300 let mut per_sample: Vec<(Vec<ErrorRow>, usize)> = Vec::new();
301 for (si, phi) in &samples {
302 let layer = skills[*si]
303 .selection
304 .as_ref()
305 .map(|s| s.phi_layer)
306 .unwrap_or(0);
307 let mut phi_of =
308 |l: usize| -> Vec<f32> { if l == layer { phi.clone() } else { Vec::new() } };
309 let rows = error_rows(model, &mut phi_of);
310 let Some(pos) = rows.iter().position(|r| r.0 == skills[*si].id) else {
311 continue;
312 };
313 per_sample.push((rows, pos));
314 }
315 if per_sample.is_empty() {
316 return None;
317 }
318 let mut best_t = 1.0f32;
320 let mut best_nll = f32::INFINITY;
321 let mut t = 1e-3f32;
322 while t <= 1e6 {
325 let mut nll = 0.0f32;
326 for (rows, pos) in &per_sample {
327 let mut logits: Vec<f32> = rows.iter().map(|r| -r.1 / t).collect();
328 softmax(&mut logits);
329 nll -= logits[*pos].max(1e-9).ln();
330 }
331 if nll < best_nll {
332 best_nll = nll;
333 best_t = t;
334 }
335 t *= 1.15;
336 }
337 let mut cal = cortiq_core::format::RoutingCalibration {
338 temperature: best_t,
339 novelty_theta: 0.5,
340 samples: per_sample.len(),
341 target_fpr,
342 };
343 let mut nov: Vec<f32> = per_sample
345 .iter()
346 .map(|(rows, _)| decide(rows, Some(&cal), 1.0).novelty)
347 .filter(|v| v.is_finite())
348 .collect();
349 nov.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
350 if !nov.is_empty() {
351 let q = (1.0 - target_fpr).clamp(0.0, 1.0);
352 let idx = (((nov.len() - 1) as f32) * q).round() as usize;
353 cal.novelty_theta = (nov[idx.min(nov.len() - 1)] + 1e-4).min(0.999);
354 }
355 Some(cal)
356}
357
358pub type ErrorRow = (String, f32, Option<f32>, Option<f32>, f32);
363
364#[derive(Debug, Clone, PartialEq, Eq)]
366pub enum RouteTarget {
367 Backbone,
368 Skill(String),
369}
370
371#[derive(Debug, Clone)]
374pub struct RouteDecision {
375 pub target: RouteTarget,
376 pub routing: Routing,
377 pub reason: String,
378}
379
380impl RouteDecision {
381 pub fn skill(&self) -> Option<&str> {
383 match &self.target {
384 RouteTarget::Skill(s) => Some(s),
385 RouteTarget::Backbone => None,
386 }
387 }
388
389 pub fn forced(target: RouteTarget, reason: impl Into<String>) -> Self {
393 Self {
394 target,
395 routing: Routing {
396 scores: Vec::new(),
397 confidence: 0.0,
398 margin: 0.0,
399 novelty: f32::NAN,
400 is_novel: false,
401 calibrated: false,
402 },
403 reason: reason.into(),
404 }
405 }
406
407 pub fn target_label(&self) -> &str {
409 self.skill().unwrap_or("backbone")
410 }
411
412 pub fn e_base(&self) -> Option<f32> {
414 self.routing
415 .scores
416 .iter()
417 .find(|s| s.id == BACKBONE_CLASS_ID)
418 .map(|s| s.error)
419 }
420
421 pub fn e_skill(&self) -> Option<f32> {
424 self.routing
425 .scores
426 .iter()
427 .find(|s| s.id != BACKBONE_CLASS_ID)
428 .map(|s| s.error)
429 }
430
431 pub fn nearest_skill(&self) -> Option<&str> {
433 self.routing
434 .scores
435 .iter()
436 .find(|s| s.id != BACKBONE_CLASS_ID)
437 .map(|s| s.id.as_str())
438 }
439
440 pub fn describe(&self) -> String {
443 let f = |v: Option<f32>| match v {
444 Some(x) if x.is_finite() => format!("{x:.4}"),
445 _ => "—".to_string(),
446 };
447 let nov = if self.routing.novelty.is_finite() {
448 format!("{:.3}", self.routing.novelty)
449 } else {
450 "—".to_string()
451 };
452 format!(
453 "route: {} | novelty {nov} | E_base {} | E_skill {} | {}",
454 self.target_label(),
455 f(self.e_base()),
456 f(self.e_skill()),
457 self.reason
458 )
459 }
460
461 pub fn summary_json(&self) -> serde_json::Value {
464 let num = |v: Option<f32>| -> serde_json::Value {
465 match v {
466 Some(x) if x.is_finite() => serde_json::json!(x),
467 _ => serde_json::Value::Null,
468 }
469 };
470 serde_json::json!({
471 "target": self.target_label(),
472 "novelty": num(Some(self.routing.novelty)),
473 "e_base": num(self.e_base()),
474 "e_skill": num(self.e_skill()),
475 "reason": self.reason,
476 })
477 }
478}
479
480#[derive(Debug, Clone, Copy, Default)]
482pub struct RouteOptions {
483 pub include_quarantine: bool,
489}
490
491pub fn is_routable(s: &SkillRecord, opts: RouteOptions) -> bool {
497 s.is_auto_routable()
498 || (opts.include_quarantine
499 && s.is_v2()
500 && s.selection.is_some()
501 && s.status.as_deref() != Some("retired"))
502}
503
504pub fn routable_skills<'a>(header: &'a CmfHeader, opts: RouteOptions) -> Vec<&'a SkillRecord> {
507 calibration_classes(header)
508 .into_iter()
509 .filter(|s| is_routable(s, opts))
510 .collect()
511}
512
513pub fn no_candidate_reason(header: &CmfHeader, opts: RouteOptions) -> Option<String> {
518 let Some(policy) = &header.router else {
519 return Some("the file declares no router policy (ROUTER_V2)".into());
520 };
521 if header.routing.is_none() {
522 return Some("the router is not calibrated (header.routing absent)".into());
523 }
524 let h = skills_hash(header);
525 if parse_hex64(&policy.skills_hash) != Some(h) {
526 return Some(format!(
527 "skills_hash mismatch (router {}, file {h:016x}): the calibration is stale",
528 policy.skills_hash
529 ));
530 }
531 if routable_skills(header, opts).is_empty() {
532 let classes: Vec<String> = calibration_classes(header)
533 .iter()
534 .map(|s| format!("{}={}", s.id, s.status.as_deref().unwrap_or("?")))
535 .collect();
536 return Some(format!(
537 "no routable skill class (classes {classes:?}; include_quarantine {})",
538 opts.include_quarantine
539 ));
540 }
541 None
542}
543
544pub const PROMPT_CONTRACT_CMF_IM_V1: &str = "cmf-im-v1";
549const IM_USER: &str = "<|im_start|>user\n";
550const IM_ASSISTANT: &str = "<|im_end|>\n<|im_start|>assistant\n";
551
552pub fn render_cmf_im_v1(user_text: &str) -> String {
555 format!("{IM_USER}{user_text}{IM_ASSISTANT}")
556}
557
558pub fn cmf_im_v1_single_user_turn(prompt: &str) -> Option<&str> {
562 let q = prompt.strip_prefix(IM_USER)?.strip_suffix(IM_ASSISTANT)?;
563 (!q.contains("<|im_start|>") && !q.contains("<|im_end|>")).then_some(q)
564}
565
566#[derive(Debug, Clone, Copy, PartialEq, Eq)]
568pub enum PromptFrame {
569 Raw,
571 CmfImV1,
574 Other,
576}
577
578impl PromptFrame {
579 fn label(self) -> &'static str {
580 match self {
581 Self::Raw => "a raw completion prompt",
582 Self::CmfImV1 => "cmf-im-v1",
583 Self::Other => "another chat template",
584 }
585 }
586}
587
588pub fn chat_frame(tok: &crate::tokenizer::Tokenizer) -> PromptFrame {
592 match &tok.chat_template {
593 Some(t) if t.contains("<|im_start|>") && t.contains("<|im_end|>") => PromptFrame::CmfImV1,
594 Some(_) => PromptFrame::Other,
595 None if tok.im_start_id.is_some() && tok.im_end_id.is_some() => PromptFrame::CmfImV1,
596 None => PromptFrame::Other,
597 }
598}
599
600pub fn enforce_prompt_contract(
612 header: &CmfHeader,
613 decision: RouteDecision,
614 frame: PromptFrame,
615 can_render: bool,
616) -> (RouteDecision, bool) {
617 let Some(id) = decision.skill() else {
618 return (decision, false);
619 };
620 let Some(contract) = header
621 .skills
622 .iter()
623 .find(|s| s.id == id)
624 .and_then(|s| s.prompt_contract.clone())
625 else {
626 return (decision, false);
627 };
628 let known = contract == PROMPT_CONTRACT_CMF_IM_V1;
629 if known && frame == PromptFrame::CmfImV1 {
630 return (decision, false);
631 }
632 if known && frame == PromptFrame::Raw && can_render {
633 let mut d = decision;
634 d.reason = format!(
635 "{}; prompt rendered as {PROMPT_CONTRACT_CMF_IM_V1} (the skill's contract)",
636 d.reason
637 );
638 return (d, true);
639 }
640 let reason = format!(
641 "prompt contract '{contract}' of skill '{id}' not satisfied ({}): the backbone runs \
642 [decision was: {}]",
643 frame.label(),
644 decision.reason
645 );
646 (
647 RouteDecision {
648 target: RouteTarget::Backbone,
649 routing: decision.routing,
650 reason,
651 },
652 false,
653 )
654}
655
656pub fn is_router_v2(model: &CmfModel) -> bool {
659 model.header.router.is_some()
660}
661
662pub fn route_request(model: &CmfModel, backbone: &mut Pipeline, user_text: &str) -> RouteDecision {
676 route_request_with(model, backbone, user_text, RouteOptions::default())
677}
678
679pub fn route_request_with(
681 model: &CmfModel,
682 backbone: &mut Pipeline,
683 user_text: &str,
684 opts: RouteOptions,
685) -> RouteDecision {
686 let q_ids = if model.header.router.is_some() {
692 backbone.tokenizer.encode_plain(user_text)
693 } else {
694 backbone.tokenizer.encode(user_text)
695 };
696 route_request_ids(model, backbone, &q_ids, opts)
697}
698
699pub fn route_request_ids(
701 model: &CmfModel,
702 backbone: &mut Pipeline,
703 q_ids: &[u32],
704 opts: RouteOptions,
705) -> RouteDecision {
706 match &model.header.router {
707 Some(policy) => {
708 let (ids, span) = phi_span_ids(&policy.phi, q_ids);
709 let phi = backbone.probe_phi_span(&ids, policy.phi.layer, span);
710 route_policy_with(&model.header, &phi, opts)
711 }
712 None => route_legacy(model, backbone, q_ids),
713 }
714}
715
716fn route_legacy(model: &CmfModel, backbone: &mut Pipeline, ids: &[u32]) -> RouteDecision {
720 let tau = std::env::var("CMF_OOD_TAU")
721 .ok()
722 .and_then(|v| v.parse().ok())
723 .unwrap_or(0.30);
724 let routing = route_full(model, backbone, ids, tau);
725 let Some(w) = routing.winner().cloned() else {
726 return RouteDecision {
727 target: RouteTarget::Backbone,
728 routing,
729 reason: "no routable skill in this container".into(),
730 };
731 };
732 if routing.is_novel {
733 let reason = if routing.calibrated {
734 format!(
735 "novel input (novelty {:.3} over θ): nearest '{}' not taken",
736 routing.novelty, w.id
737 )
738 } else {
739 format!(
740 "novel input (E_min {:.4} > τ {tau}, uncalibrated file): nearest '{}' not taken",
741 w.error, w.id
742 )
743 };
744 return RouteDecision {
745 target: RouteTarget::Backbone,
746 routing,
747 reason,
748 };
749 }
750 let reason = format!(
751 "legacy recon-argmin: '{}' (E {:.4}), in scope",
752 w.id, w.error
753 );
754 RouteDecision {
755 target: RouteTarget::Skill(w.id),
756 routing,
757 reason,
758 }
759}
760
761pub fn calibration_classes(header: &CmfHeader) -> Vec<&SkillRecord> {
766 let mut v: Vec<&SkillRecord> = header
767 .skills
768 .iter()
769 .filter(|s| s.is_v2() && s.selection.is_some() && s.status.as_deref() != Some("retired"))
770 .collect();
771 v.sort_by(|a, b| a.id.as_bytes().cmp(b.id.as_bytes()));
772 v
773}
774
775pub fn skills_hash(header: &CmfHeader) -> u64 {
784 fn push(buf: &mut Vec<u8>, id: &str, d: &SelectionDescriptor) {
785 buf.extend_from_slice(id.as_bytes());
786 buf.push(0);
787 buf.extend_from_slice(d.metric.as_bytes());
788 buf.push(0);
789 buf.extend_from_slice(&(d.phi_layer as u64).to_le_bytes());
790 buf.extend_from_slice(&(d.rank as u64).to_le_bytes());
791 buf.extend_from_slice(d.mean.as_bytes());
792 buf.push(0);
793 buf.extend_from_slice(d.basis.as_bytes());
794 buf.push(0);
795 for x in [d.err_mean, d.err_std] {
796 match x {
797 Some(v) => {
798 buf.push(1);
799 buf.extend_from_slice(&v.to_bits().to_le_bytes());
800 }
801 None => buf.push(0),
802 }
803 }
804 buf.push(0x1e);
805 }
806 let mut buf = b"cmf-skills-v2\0".to_vec();
807 match &header.router {
808 Some(r) => push(&mut buf, BACKBONE_CLASS_ID, &r.base),
809 None => buf.push(0),
810 }
811 for s in calibration_classes(header) {
812 push(&mut buf, &s.id, s.selection.as_ref().expect("filtered"));
813 }
814 hash64(&buf)
815}
816
817pub fn phi_span_ids(policy: &PhiSpec, q_ids: &[u32]) -> (Vec<u32>, std::ops::Range<usize>) {
822 let mut ids =
823 Vec::with_capacity(policy.prefix_ids.len() + q_ids.len() + policy.suffix_ids.len());
824 ids.extend_from_slice(&policy.prefix_ids);
825 let start = ids.len();
826 ids.extend_from_slice(q_ids);
827 let end = ids.len();
828 ids.extend_from_slice(&policy.suffix_ids);
829 (ids, start..end)
830}
831
832pub fn pool_span_unit(hiddens: &[f32], hidden: usize, span: std::ops::Range<usize>) -> Vec<f32> {
837 let rows = hiddens.len() / hidden.max(1);
838 let (a, b) = (span.start.min(rows), span.end.min(rows));
839 let mut phi = vec![0.0f32; hidden];
840 if b <= a {
841 return phi;
842 }
843 for r in a..b {
844 for (p, x) in phi.iter_mut().zip(&hiddens[r * hidden..(r + 1) * hidden]) {
845 *p += x;
846 }
847 }
848 let n = phi.iter().map(|x| x * x).sum::<f32>().sqrt();
849 if n > 0.0 {
850 phi.iter_mut().for_each(|x| *x /= n);
851 }
852 phi
853}
854
855pub fn holdout_count(n: usize) -> usize {
862 (n / 5).clamp(1, n.saturating_sub(2).max(1))
863}
864
865pub fn encode_f16(v: &[f32]) -> String {
869 let bytes: Vec<u8> = v
870 .iter()
871 .flat_map(|x| f32_to_f16(*x).to_le_bytes())
872 .collect();
873 base64::engine::general_purpose::STANDARD.encode(bytes)
874}
875
876fn f16_round(v: &[f32]) -> Vec<f32> {
877 v.iter().map(|x| f16_to_f32(f32_to_f16(*x))).collect()
878}
879
880fn gauss_vec(seed: u64, n: usize) -> Vec<f32> {
883 let mut s = seed
884 .wrapping_mul(6_364_136_223_846_793_005)
885 .wrapping_add(1_442_695_040_888_963_407);
886 let mut next = || {
887 s = s
888 .wrapping_mul(6_364_136_223_846_793_005)
889 .wrapping_add(1_442_695_040_888_963_407);
890 (s >> 11) as f64 / (1u64 << 53) as f64
891 };
892 let mut out = Vec::with_capacity(n + 1);
893 while out.len() < n {
894 let (u1, u2) = (next().max(1e-12), next());
895 let r = (-2.0 * u1.ln()).sqrt();
896 let t = 2.0 * std::f64::consts::PI * u2;
897 out.push((r * t.cos()) as f32);
898 out.push((r * t.sin()) as f32);
899 }
900 out.truncate(n);
901 out
902}
903
904pub fn top_principal_rows(
912 centered: &[Vec<f32>],
913 h: usize,
914 k: usize,
915 iters: usize,
916 seed: u64,
917) -> Vec<f32> {
918 let mut q = gauss_vec(seed.wrapping_add(777), k * h);
919 let mut spare = seed.wrapping_add(1);
920 let mut orth = |q: &mut [f32]| {
921 for i in 0..k {
922 loop {
923 for j in 0..i {
924 let dot: f32 = (0..h).map(|t| q[i * h + t] * q[j * h + t]).sum();
925 for t in 0..h {
926 q[i * h + t] -= dot * q[j * h + t];
927 }
928 }
929 let nrm: f32 = (0..h)
930 .map(|t| q[i * h + t] * q[i * h + t])
931 .sum::<f32>()
932 .sqrt();
933 if nrm > 1e-12 {
934 for t in 0..h {
935 q[i * h + t] /= nrm;
936 }
937 break;
938 }
939 let fresh = gauss_vec(spare, h);
940 spare = spare.wrapping_add(1);
941 q[i * h..(i + 1) * h].copy_from_slice(&fresh);
942 }
943 }
944 };
945 orth(&mut q);
946 let mut tmp = vec![0f32; k * h];
947 let mut coef = vec![0f32; k];
948 for _ in 0..iters {
949 tmp.fill(0.0);
950 for c in centered {
951 for (i, ci) in coef.iter_mut().enumerate() {
952 *ci = q[i * h..(i + 1) * h]
953 .iter()
954 .zip(c)
955 .map(|(a, b)| a * b)
956 .sum();
957 }
958 for (i, &ci) in coef.iter().enumerate() {
959 if ci != 0.0 {
960 for (t, x) in tmp[i * h..(i + 1) * h].iter_mut().zip(c) {
961 *t += ci * x;
962 }
963 }
964 }
965 }
966 q.copy_from_slice(&tmp);
967 orth(&mut q);
968 }
969 q
970}
971
972pub fn fit_descriptor(
983 phis_raw: &[Vec<f32>],
984 phi_layer: usize,
985 rank: usize,
986) -> Option<SelectionDescriptor> {
987 let n = phis_raw.len();
988 if n < 2 {
989 return None;
990 }
991 let h = phis_raw[0].len();
992 if h == 0 || phis_raw.iter().any(|p| p.len() != h) {
993 return None;
994 }
995 let phis: Vec<Vec<f32>> = phis_raw.iter().map(|p| unit_copy(p)).collect();
996 let n_hold = holdout_count(n);
997 let (train, hold) = phis.split_at(n - n_hold);
998 let nt = train.len();
999 let mut mean = vec![0f32; h];
1000 for p in train {
1001 for (m, v) in mean.iter_mut().zip(p) {
1002 *m += v / nt as f32;
1003 }
1004 }
1005 let centered: Vec<Vec<f32>> = train
1006 .iter()
1007 .map(|p| p.iter().zip(&mean).map(|(v, m)| v - m).collect())
1008 .collect();
1009 let rank = rank.min(nt.saturating_sub(1)).max(1);
1010 let basis = top_principal_rows(¢ered, h, rank, 120, 99);
1011 let (mq, bq) = (f16_round(&mean), f16_round(&basis));
1012 let errs: Vec<f32> = train.iter().map(|p| recon_error(p, &mq, &bq, rank)).collect();
1013 let em = errs.iter().sum::<f32>() / nt as f32;
1014 let es = (errs.iter().map(|e| (e - em).powi(2)).sum::<f32>() / nt as f32)
1015 .sqrt()
1016 .max(1e-4);
1017 let hold_flat: Vec<f32> = hold.concat();
1018 Some(SelectionDescriptor {
1019 metric: METRIC_MSE_UNIT.into(),
1020 phi_layer,
1021 mean: encode_f16(&mean),
1022 basis: encode_f16(&basis),
1023 rank,
1024 err_mean: Some(em),
1025 err_std: Some(es),
1026 holdout: Some(encode_f16(&hold_flat)),
1027 holdout_n: Some(hold.len()),
1028 })
1029}
1030
1031pub fn cmf_im_v1_phi_spec(tok: &crate::tokenizer::Tokenizer, layer: usize) -> Result<PhiSpec, String> {
1039 for s in ["<|im_start|>", "<|im_end|>"] {
1040 if tok.encode(s).len() != 1 {
1041 return Err(format!(
1042 "the tokenizer has no added token {s}: the cmf-im-v1 frame cannot be rendered"
1043 ));
1044 }
1045 }
1046 Ok(PhiSpec {
1047 layer,
1048 pool: "span_mean".into(),
1049 norm: "unit".into(),
1050 prefix_ids: tok.encode(IM_USER),
1051 suffix_ids: tok.encode(IM_ASSISTANT),
1052 })
1053}
1054
1055fn unit_copy(phi: &[f32]) -> Vec<f32> {
1056 let n = phi.iter().map(|x| x * x).sum::<f32>().sqrt();
1057 if n > 0.0 {
1058 phi.iter().map(|x| x / n).collect()
1059 } else {
1060 phi.to_vec()
1061 }
1062}
1063
1064pub fn descriptor_row(
1066 id: &str,
1067 sel: &SelectionDescriptor,
1068 phi_unit: &[f32],
1069 hidden: usize,
1070) -> Option<ErrorRow> {
1071 if sel.metric != METRIC_MSE_UNIT {
1072 return None;
1073 }
1074 let (mean, basis) = (decode_f16(&sel.mean)?, decode_f16(&sel.basis)?);
1075 if mean.len() != hidden || basis.len() != sel.rank * hidden || phi_unit.len() != hidden {
1076 return None;
1077 }
1078 let e = recon_error(phi_unit, &mean, &basis, sel.rank);
1079 let pp: f32 = phi_unit.iter().map(|v| v * v).sum();
1080 Some((id.to_string(), e, sel.err_mean, sel.err_std, pp))
1081}
1082
1083pub fn policy_rows(header: &CmfHeader, phi: &[f32]) -> (Option<ErrorRow>, Vec<ErrorRow>) {
1092 let Some(r) = &header.router else {
1093 return (None, Vec::new());
1094 };
1095 let hidden = header.arch.hidden_size;
1096 let phi = unit_copy(phi);
1097 let base = descriptor_row(BACKBONE_CLASS_ID, &r.base, &phi, hidden);
1098 let skills = calibration_classes(header)
1099 .into_iter()
1100 .filter_map(|s| {
1101 let sel = s.selection.as_ref()?;
1102 if sel.phi_layer != r.phi.layer {
1103 return None;
1104 }
1105 descriptor_row(&s.id, sel, &phi, hidden)
1106 })
1107 .collect();
1108 (base, skills)
1109}
1110
1111pub fn decide_backbone_gated(
1115 base: Option<&ErrorRow>,
1116 skills: &[ErrorRow],
1117 calibration: Option<&RoutingCalibration>,
1118 policy: Option<&RouterPolicy>,
1119 file_skills_hash: u64,
1120) -> RouteDecision {
1121 decide_backbone_gated_with(
1122 base,
1123 skills,
1124 &|_| true,
1125 calibration,
1126 policy,
1127 file_skills_hash,
1128 )
1129}
1130
1131#[allow(clippy::neg_cmp_op_on_partial_ord)]
1144pub fn decide_backbone_gated_with(
1145 base: Option<&ErrorRow>,
1146 skills: &[ErrorRow],
1147 routable: &dyn Fn(&str) -> bool,
1148 calibration: Option<&RoutingCalibration>,
1149 policy: Option<&RouterPolicy>,
1150 file_skills_hash: u64,
1151) -> RouteDecision {
1152 let mut rows: Vec<ErrorRow> = Vec::with_capacity(skills.len() + 1);
1153 rows.extend(base.cloned());
1154 rows.extend(skills.iter().cloned());
1155 let backbone = |routing: Routing, reason: String| RouteDecision {
1156 target: RouteTarget::Backbone,
1157 routing,
1158 reason,
1159 };
1160 let Some(policy) = policy else {
1161 return backbone(
1162 decide(&rows, None, 1.0),
1163 "no router policy (ROUTER_V2 absent): the backbone runs".into(),
1164 );
1165 };
1166 if parse_hex64(&policy.skills_hash) != Some(file_skills_hash) {
1167 return backbone(
1168 decide(&rows, None, 1.0),
1169 format!(
1170 "skills_hash mismatch (router {}, file {file_skills_hash:016x}): the calibration \
1171 is stale — recalibrate",
1172 policy.skills_hash
1173 ),
1174 );
1175 }
1176 let Some(cal) = calibration else {
1177 return backbone(decide(&rows, None, 1.0), "router not calibrated".into());
1178 };
1179 let Some(base) = base else {
1180 return backbone(
1181 decide(&rows, Some(cal), 1.0),
1182 "no backbone descriptor row".into(),
1183 );
1184 };
1185 if !(base.4 > 1e-12) {
1186 return backbone(
1187 decide(&rows, Some(cal), 1.0),
1188 "degenerate φ (empty span)".into(),
1189 );
1190 }
1191 let routing = decide(&rows, Some(cal), 1.0);
1192 if !skills.iter().any(|r| routable(&r.0)) {
1193 return backbone(
1194 routing,
1195 "no routable skill (v2, active, measured gate)".into(),
1196 );
1197 }
1198 let e_base = base.1 / base.4.max(1e-12);
1199 let Some(winner) = routing.winner().cloned() else {
1200 return backbone(routing, "no scored class".into());
1201 };
1202 if winner.id == BACKBONE_CLASS_ID {
1203 return backbone(
1204 routing,
1205 format!("the backbone is the nearest class (E_base {e_base:.4})"),
1206 );
1207 }
1208 if !routable(&winner.id) {
1209 return backbone(
1210 routing,
1211 format!(
1212 "nearest class '{}' is not routable (not active with a measured gate): the \
1213 backbone runs",
1214 winner.id
1215 ),
1216 );
1217 }
1218 if routing.is_novel {
1219 let reason = format!(
1220 "novel input (novelty {:.3} > θ {:.3})",
1221 routing.novelty, cal.novelty_theta
1222 );
1223 return backbone(routing, reason);
1224 }
1225 let e_skill = winner.error;
1226 if !(e_skill + policy.margin < e_base) {
1227 return backbone(
1228 routing,
1229 format!(
1230 "margin not beaten: E_skill {e_skill:.4} + margin {:.4} ≥ E_base {e_base:.4}",
1231 policy.margin
1232 ),
1233 );
1234 }
1235 let reason = format!(
1236 "skill '{}': E {e_skill:.4} + margin {:.4} < E_base {e_base:.4}, confidence {:.3}, \
1237 novelty {:.3} ≤ θ {:.3}",
1238 winner.id, policy.margin, routing.confidence, routing.novelty, cal.novelty_theta
1239 );
1240 RouteDecision {
1241 target: RouteTarget::Skill(winner.id),
1242 routing,
1243 reason,
1244 }
1245}
1246
1247pub fn route_policy(header: &CmfHeader, phi: &[f32]) -> RouteDecision {
1250 route_policy_with(header, phi, RouteOptions::default())
1251}
1252
1253pub fn route_policy_with(header: &CmfHeader, phi: &[f32], opts: RouteOptions) -> RouteDecision {
1256 let (base, skills) = policy_rows(header, phi);
1257 let routable = |id: &str| {
1258 header
1259 .skills
1260 .iter()
1261 .find(|s| s.id == id)
1262 .is_some_and(|s| is_routable(s, opts))
1263 };
1264 decide_backbone_gated_with(
1265 base.as_ref(),
1266 &skills,
1267 &routable,
1268 header.routing.as_ref(),
1269 header.router.as_ref(),
1270 skills_hash(header),
1271 )
1272}
1273
1274pub fn clopper_pearson_upper(k: usize, n: usize, confidence: f64) -> f64 {
1278 if n == 0 || k >= n {
1279 return 1.0;
1280 }
1281 let q = 1.0 - (1.0 - confidence) / 2.0;
1282 let (a, b) = ((k + 1) as f64, (n - k) as f64);
1283 if k == 0 {
1284 return 1.0 - (1.0 - q).powf(1.0 / n as f64);
1286 }
1287 let (mut lo, mut hi) = (0.0f64, 1.0f64);
1289 for _ in 0..200 {
1290 let mid = 0.5 * (lo + hi);
1291 if reg_inc_beta(a, b, mid) < q {
1292 lo = mid;
1293 } else {
1294 hi = mid;
1295 }
1296 }
1297 0.5 * (lo + hi)
1298}
1299
1300fn ln_gamma(x: f64) -> f64 {
1301 const C: [f64; 9] = [
1303 0.999_999_999_999_809_9,
1304 676.520_368_121_885_1,
1305 -1_259.139_216_722_402_8,
1306 771.323_428_777_653_1,
1307 -176.615_029_162_140_6,
1308 12.507_343_278_686_905,
1309 -0.138_571_095_265_720_12,
1310 9.984_369_578_019_572e-6,
1311 1.505_632_735_149_311_6e-7,
1312 ];
1313 if x < 0.5 {
1314 return (std::f64::consts::PI / (std::f64::consts::PI * x).sin()).ln() - ln_gamma(1.0 - x);
1315 }
1316 let x = x - 1.0;
1317 let mut acc = C[0];
1318 for (i, c) in C.iter().enumerate().skip(1) {
1319 acc += c / (x + i as f64);
1320 }
1321 let t = x + 7.5;
1322 0.5 * (2.0 * std::f64::consts::PI).ln() + (x + 0.5) * t.ln() - t + acc.ln()
1323}
1324
1325fn reg_inc_beta(a: f64, b: f64, x: f64) -> f64 {
1327 if x <= 0.0 {
1328 return 0.0;
1329 }
1330 if x >= 1.0 {
1331 return 1.0;
1332 }
1333 let ln_front = ln_gamma(a + b) - ln_gamma(a) - ln_gamma(b) + a * x.ln() + b * (1.0 - x).ln();
1334 let cf = |a: f64, b: f64, x: f64| -> f64 {
1335 let tiny = 1e-300;
1336 let (qab, qap, qam) = (a + b, a + 1.0, a - 1.0);
1337 let mut c = 1.0;
1338 let mut d = 1.0 - qab * x / qap;
1339 if d.abs() < tiny {
1340 d = tiny;
1341 }
1342 d = 1.0 / d;
1343 let mut h = d;
1344 for m in 1..400 {
1345 let m = m as f64;
1346 let m2 = 2.0 * m;
1347 let aa = m * (b - m) * x / ((qam + m2) * (a + m2));
1348 d = 1.0 + aa * d;
1349 if d.abs() < tiny {
1350 d = tiny;
1351 }
1352 c = 1.0 + aa / c;
1353 if c.abs() < tiny {
1354 c = tiny;
1355 }
1356 d = 1.0 / d;
1357 h *= d * c;
1358 let aa = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2));
1359 d = 1.0 + aa * d;
1360 if d.abs() < tiny {
1361 d = tiny;
1362 }
1363 c = 1.0 + aa / c;
1364 if c.abs() < tiny {
1365 c = tiny;
1366 }
1367 d = 1.0 / d;
1368 let del = d * c;
1369 h *= del;
1370 if (del - 1.0).abs() < 1e-15 {
1371 break;
1372 }
1373 }
1374 h
1375 };
1376 if x < (a + 1.0) / (a + b + 2.0) {
1377 ln_front.exp() * cf(a, b, x) / a
1378 } else {
1379 1.0 - ln_front.exp() * cf(b, a, 1.0 - x) / b
1380 }
1381}
1382
1383pub fn calibrate_v2(
1396 header: &CmfHeader,
1397 target_fpr: f32,
1398) -> Result<(RoutingCalibration, serde_json::Value), String> {
1399 let policy = header
1400 .router
1401 .as_ref()
1402 .ok_or("calibrate_v2: the file declares no router policy (header.router)")?;
1403 let hidden = header.arch.hidden_size;
1404 let classes = calibration_classes(header);
1405 if classes.is_empty() {
1406 return Err("calibrate_v2: no v2 skill with a selection descriptor".into());
1407 }
1408 let mut descs: Vec<(&str, &SelectionDescriptor)> = vec![(BACKBONE_CLASS_ID, &policy.base)];
1410 for s in &classes {
1411 let sel = s.selection.as_ref().expect("filtered");
1412 if sel.metric != METRIC_MSE_UNIT || sel.phi_layer != policy.phi.layer {
1413 return Err(format!(
1414 "calibrate_v2: skill '{}' descriptor is {}@{}, the policy needs {METRIC_MSE_UNIT}@{}",
1415 s.id, sel.metric, sel.phi_layer, policy.phi.layer
1416 ));
1417 }
1418 descs.push((&s.id, sel));
1419 }
1420 let holdout = |d: &SelectionDescriptor| -> Vec<Vec<f32>> {
1421 let (Some(h), Some(n)) = (d.holdout.as_ref(), d.holdout_n) else {
1422 return Vec::new();
1423 };
1424 match decode_f16(h) {
1425 Some(v) if v.len() == n * hidden => (0..n)
1426 .map(|i| unit_copy(&v[i * hidden..(i + 1) * hidden]))
1427 .collect(),
1428 _ => Vec::new(),
1429 }
1430 };
1431 let mut samples: Vec<(usize, Vec<ErrorRow>)> = Vec::new();
1433 for (ci, (_, d)) in descs.iter().enumerate() {
1434 for phi in holdout(d) {
1435 let rows: Option<Vec<ErrorRow>> = descs
1436 .iter()
1437 .map(|(id, dd)| descriptor_row(id, dd, &phi, hidden))
1438 .collect();
1439 let rows = rows.ok_or("calibrate_v2: a descriptor is malformed (dims / base64)")?;
1440 samples.push((ci, rows));
1441 }
1442 }
1443 let n_general = samples.iter().filter(|(c, _)| *c == 0).count();
1444 let n_in = samples.len() - n_general;
1445 if n_general == 0 {
1446 return Err("calibrate_v2: router.base carries no holdout (general prompts)".into());
1447 }
1448 if n_in == 0 {
1449 return Err("calibrate_v2: no skill carries an in-scope holdout".into());
1450 }
1451
1452 let mut best_t = 1.0f32;
1454 let mut best_nll = f32::INFINITY;
1455 let mut t = 1e-3f32;
1456 while t <= 1e6 {
1457 let mut nll = 0.0f32;
1458 for (ci, rows) in &samples {
1459 let mut logits: Vec<f32> = rows.iter().map(|r| -r.1 / t).collect();
1460 softmax(&mut logits);
1461 nll -= logits[*ci].max(1e-9).ln();
1462 }
1463 if nll < best_nll {
1464 best_nll = nll;
1465 best_t = t;
1466 }
1467 t *= 1.15;
1468 }
1469 let mut cal = RoutingCalibration {
1470 temperature: best_t,
1471 novelty_theta: 0.5,
1472 samples: samples.len(),
1473 target_fpr,
1474 };
1475 let mut nov: Vec<f32> = samples
1477 .iter()
1478 .filter(|(c, _)| *c > 0)
1479 .map(|(_, rows)| decide(rows, Some(&cal), 1.0).novelty)
1480 .filter(|v| v.is_finite())
1481 .collect();
1482 nov.sort_by(|a, b| a.total_cmp(b));
1483 if !nov.is_empty() {
1484 let q = (1.0 - target_fpr).clamp(0.0, 1.0);
1485 let idx = (((nov.len() - 1) as f32) * q).round() as usize;
1486 cal.novelty_theta = (nov[idx.min(nov.len() - 1)] + 1e-4).min(0.999);
1487 }
1488
1489 let hash = skills_hash(header);
1491 let mut pol = policy.clone();
1492 pol.skills_hash = format!("{hash:016x}");
1493 let (mut fa, mut hit) = (0usize, 0usize);
1494 let mut per: Vec<(usize, usize)> = vec![(0, 0); descs.len()];
1495 for (ci, rows) in &samples {
1496 let d = decide_backbone_gated(Some(&rows[0]), &rows[1..], Some(&cal), Some(&pol), hash);
1497 if *ci == 0 {
1498 fa += usize::from(d.skill().is_some());
1499 } else {
1500 per[*ci].0 += 1;
1501 if d.skill() == Some(descs[*ci].0) {
1502 hit += 1;
1503 per[*ci].1 += 1;
1504 }
1505 }
1506 }
1507 let per_skill: serde_json::Map<String, serde_json::Value> = descs
1508 .iter()
1509 .enumerate()
1510 .skip(1)
1511 .map(|(i, (id, _))| {
1512 let (n, h) = per[i];
1513 (
1514 id.to_string(),
1515 serde_json::json!({"n": n, "recall": if n > 0 { h as f64 / n as f64 } else { f64::NAN }}),
1516 )
1517 })
1518 .collect();
1519 let measured = serde_json::json!({
1520 "set": "calibration",
1521 "n_in": n_in,
1522 "n_general": n_general,
1523 "in_scope_recall": hit as f64 / n_in as f64,
1524 "false_accept": fa as f64 / n_general as f64,
1525 "false_accept_upper95": clopper_pearson_upper(fa, n_general, 0.95),
1526 "temperature": cal.temperature,
1527 "novelty_theta": cal.novelty_theta,
1528 "target_fpr": target_fpr,
1529 "margin": policy.margin,
1530 "skills_hash": format!("{hash:016x}"),
1531 "per_skill": per_skill,
1532 });
1533 Ok((cal, measured))
1534}
1535
1536#[cfg(test)]
1537mod tests {
1538 use super::*;
1539 use cortiq_core::knowledge::skill_kind;
1540
1541 const HID: usize = 16;
1542
1543 #[test]
1544 fn holdout_count_is_a_fifth_at_least_one_and_leaves_train_samples() {
1545 for (n, want) in [(1, 1), (2, 1), (3, 1), (5, 1), (6, 1), (10, 2), (50, 10), (4000, 800)] {
1546 assert_eq!(holdout_count(n), want, "n = {n}");
1547 }
1548 }
1549
1550 #[test]
1551 fn encode_f16_round_trips_through_decode() {
1552 let v: Vec<f32> = (0..HID).map(|i| (i as f32 - 7.5) * 0.125).collect();
1553 assert_eq!(decode_f16(&encode_f16(&v)).unwrap(), v);
1554 }
1555
1556 #[test]
1561 fn fit_descriptor_recovers_the_dominant_direction() {
1562 let h = 8;
1563 let base: Vec<f32> = (0..h).map(|i| 1.0 + 0.05 * i as f32).collect();
1564 let dir: Vec<f32> = (0..h).map(|i| if i % 2 == 0 { 1.0 } else { -1.0 }).collect();
1565 let noise = gauss_vec(5, 40 * h);
1566 let samples: Vec<Vec<f32>> = (0..40)
1567 .map(|s| {
1568 let a = (s as f32 / 39.0 - 0.5) * 2.0;
1569 (0..h)
1570 .map(|i| base[i] + a * dir[i] + 0.02 * noise[s * h + i])
1571 .collect()
1572 })
1573 .collect();
1574 let d = fit_descriptor(&samples, 3, 1).expect("fit");
1575 assert_eq!(d.metric, METRIC_MSE_UNIT);
1576 assert_eq!(d.phi_layer, 3);
1577 assert_eq!(d.rank, 1);
1578 assert_eq!(d.holdout_n, Some(8));
1579 let basis = decode_f16(&d.basis).unwrap();
1580 assert_eq!(basis.len(), h);
1581 let mean = decode_f16(&d.mean).unwrap();
1585 let unit = |v: &[f32]| unit_copy(v);
1586 let mut proj = vec![0f32; h];
1587 let (hi, lo) = (unit(&samples[39]), unit(&samples[0]));
1588 for i in 0..h {
1589 proj[i] = hi[i] - lo[i];
1590 }
1591 let proj = unit(&proj);
1592 let cos: f32 = basis.iter().zip(&proj).map(|(a, b)| a * b).sum();
1593 assert!(cos.abs() > 0.95, "cos {cos}");
1594 assert!(d.err_mean.unwrap().is_finite() && d.err_std.unwrap() >= 1e-4);
1595 assert_eq!(mean.len(), h);
1596 let hold = decode_f16(d.holdout.as_ref().unwrap()).unwrap();
1597 assert_eq!(hold.len(), 8 * h);
1598 let last = unit(&samples[39]);
1600 for (a, b) in hold[7 * h..].iter().zip(&last) {
1601 assert!((a - b).abs() < 2e-3, "{a} vs {b}");
1602 }
1603 let d2 = fit_descriptor(&samples, 3, 2).unwrap();
1605 let b2 = decode_f16(&d2.basis).unwrap();
1606 let dot: f32 = b2[..h].iter().zip(&b2[h..]).map(|(a, b)| a * b).sum();
1607 let n1: f32 = b2[..h].iter().map(|a| a * a).sum::<f32>().sqrt();
1608 assert!(dot.abs() < 1e-2 && (n1 - 1.0).abs() < 1e-2, "dot {dot} n1 {n1}");
1609 assert_eq!(fit_descriptor(&samples[..3], 0, 16).unwrap().rank, 1);
1611 assert!(fit_descriptor(&samples[..1], 0, 1).is_none());
1612 assert!(fit_descriptor(&[vec![1.0], vec![1.0, 2.0]], 0, 1).is_none());
1613 assert_eq!(fit_descriptor(&samples, 3, 2).unwrap().basis, d2.basis);
1615 }
1616
1617 fn f16_b64(v: &[f32]) -> String {
1618 let bytes: Vec<u8> = v
1619 .iter()
1620 .flat_map(|x| cortiq_core::quant::f32_to_f16(*x).to_le_bytes())
1621 .collect();
1622 base64::engine::general_purpose::STANDARD.encode(bytes)
1623 }
1624
1625 fn header() -> CmfHeader {
1626 serde_json::from_value(serde_json::json!({
1627 "version": 2,
1628 "quant_type": "F32",
1629 "arch": {
1630 "arch_name": "synthetic", "hidden_size": HID, "intermediate_size": 32,
1631 "num_layers": 8, "num_attention_heads": 2, "num_kv_heads": 1, "head_dim": 8,
1632 "vocab_size": 16, "layer_types": vec!["FullAttention"; 8],
1633 "rms_norm_eps": 1e-6, "max_position_embeddings": 64
1634 }
1635 }))
1636 .unwrap()
1637 }
1638
1639 struct Rng(u64);
1641 impl Rng {
1642 fn next(&mut self) -> f32 {
1643 self.0 ^= self.0 << 13;
1644 self.0 ^= self.0 >> 7;
1645 self.0 ^= self.0 << 17;
1646 ((self.0 >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
1647 }
1648 }
1649
1650 fn sample(rng: &mut Rng, axis: usize, sigma: f32) -> Vec<f32> {
1652 let mut v: Vec<f32> = (0..HID).map(|_| sigma * rng.next()).collect();
1653 v[axis] += 1.0;
1654 unit_copy(&v)
1655 }
1656
1657 fn descriptor(
1660 train: &[Vec<f32>],
1661 holdout: &[Vec<f32>],
1662 noise_axis: usize,
1663 ) -> SelectionDescriptor {
1664 let mut mean = vec![0.0f32; HID];
1665 for s in train {
1666 for (m, x) in mean.iter_mut().zip(s) {
1667 *m += x / train.len() as f32;
1668 }
1669 }
1670 let mut basis = vec![0.0f32; HID];
1671 basis[noise_axis] = 1.0;
1672 let q = |v: &[f32]| -> Vec<f32> {
1673 v.iter()
1674 .map(|x| f16_to_f32(cortiq_core::quant::f32_to_f16(*x)))
1675 .collect()
1676 };
1677 let (mq, bq) = (q(&mean), q(&basis));
1678 let errs: Vec<f32> = train.iter().map(|s| recon_error(s, &mq, &bq, 1)).collect();
1679 let em = errs.iter().sum::<f32>() / errs.len() as f32;
1680 let es = (errs.iter().map(|e| (e - em).powi(2)).sum::<f32>() / errs.len() as f32).sqrt();
1681 SelectionDescriptor {
1682 metric: "mse_unit".into(),
1683 phi_layer: 2,
1684 mean: f16_b64(&mean),
1685 basis: f16_b64(&basis),
1686 rank: 1,
1687 err_mean: Some(em),
1688 err_std: Some(es),
1689 holdout: Some(f16_b64(&holdout.concat())),
1690 holdout_n: Some(holdout.len()),
1691 }
1692 }
1693
1694 fn policy(base: SelectionDescriptor) -> RouterPolicy {
1695 RouterPolicy {
1696 version: 2,
1697 policy: "backbone_gated".into(),
1698 granularity: "request".into(),
1699 phi: PhiSpec {
1700 layer: 2,
1701 pool: "span_mean".into(),
1702 norm: "unit".into(),
1703 prefix_ids: vec![1, 7, 8],
1704 suffix_ids: vec![2, 1, 9],
1705 },
1706 base,
1707 margin: 0.05,
1708 skills_hash: "0".into(),
1709 measured: None,
1710 }
1711 }
1712
1713 fn skill(id: &str, sel: SelectionDescriptor, status: &str) -> SkillRecord {
1714 SkillRecord {
1715 id: id.into(),
1716 layers: vec![5],
1717 selection: Some(sel),
1718 kind: Some(skill_kind::FFN_REPLACE.into()),
1719 status: Some(status.into()),
1720 gate: Some(serde_json::json!({"status": "measured"})),
1721 ..Default::default()
1722 }
1723 }
1724
1725 fn row(id: &str, e: f32) -> ErrorRow {
1726 (id.into(), e, Some(0.05), Some(0.05), 1.0)
1727 }
1728
1729 fn truth_policy() -> RouterPolicy {
1730 let mut p = policy(SelectionDescriptor {
1731 metric: "mse_unit".into(),
1732 phi_layer: 2,
1733 mean: String::new(),
1734 basis: String::new(),
1735 rank: 0,
1736 err_mean: None,
1737 err_std: None,
1738 holdout: None,
1739 holdout_n: None,
1740 });
1741 p.skills_hash = "00000000000000ab".into();
1742 p
1743 }
1744
1745 #[test]
1746 fn backbone_gated_truth_table() {
1747 let pol = truth_policy();
1748 let cal = RoutingCalibration {
1749 temperature: 0.1,
1750 novelty_theta: 0.99,
1751 samples: 10,
1752 target_fpr: 0.05,
1753 };
1754 let (base, near) = (row(BACKBONE_CLASS_ID, 0.9), row("herbs", 0.05));
1755 let go = |b: Option<&ErrorRow>,
1756 s: &[ErrorRow],
1757 c: Option<&RoutingCalibration>,
1758 p: Option<&RouterPolicy>,
1759 h: u64| { decide_backbone_gated(b, s, c, p, h) };
1760 let std_skills = std::slice::from_ref(&near);
1761
1762 let d = go(Some(&base), std_skills, Some(&cal), Some(&pol), 0xab);
1764 assert_eq!(d.target, RouteTarget::Skill("herbs".into()), "{}", d.reason);
1765
1766 let cases: Vec<(&str, RouteDecision, &str)> = vec![
1767 (
1768 "no policy",
1769 go(Some(&base), std_skills, Some(&cal), None, 0xab),
1770 "no router policy",
1771 ),
1772 (
1773 "hash mismatch",
1774 go(Some(&base), std_skills, Some(&cal), Some(&pol), 0xac),
1775 "skills_hash mismatch",
1776 ),
1777 (
1778 "no calibration",
1779 go(Some(&base), std_skills, None, Some(&pol), 0xab),
1780 "not calibrated",
1781 ),
1782 (
1783 "no base row",
1784 go(None, std_skills, Some(&cal), Some(&pol), 0xab),
1785 "no backbone descriptor",
1786 ),
1787 (
1788 "no skills",
1789 go(Some(&base), &[], Some(&cal), Some(&pol), 0xab),
1790 "no routable skill",
1791 ),
1792 (
1793 "winner is the backbone",
1794 go(
1795 Some(&row(BACKBONE_CLASS_ID, 0.01)),
1796 std_skills,
1797 Some(&cal),
1798 Some(&pol),
1799 0xab,
1800 ),
1801 "backbone is the nearest",
1802 ),
1803 (
1804 "novel",
1805 go(
1806 Some(&base),
1807 std_skills,
1808 Some(&RoutingCalibration {
1809 novelty_theta: 0.0,
1810 ..cal.clone()
1811 }),
1812 Some(&pol),
1813 0xab,
1814 ),
1815 "novel input",
1816 ),
1817 (
1818 "margin not beaten",
1819 go(
1820 Some(&row(BACKBONE_CLASS_ID, 0.32)),
1821 &[row("herbs", 0.30)],
1822 Some(&cal),
1823 Some(&pol),
1824 0xab,
1825 ),
1826 "margin not beaten",
1827 ),
1828 (
1829 "degenerate φ",
1830 go(
1831 Some(&(BACKBONE_CLASS_ID.into(), 0.9, None, None, 0.0)),
1832 std_skills,
1833 Some(&cal),
1834 Some(&pol),
1835 0xab,
1836 ),
1837 "degenerate",
1838 ),
1839 ];
1840 for (name, d, needle) in cases {
1841 assert_eq!(
1842 d.target,
1843 RouteTarget::Backbone,
1844 "case '{name}' routed to a skill"
1845 );
1846 assert!(
1847 d.reason.contains(needle),
1848 "case '{name}': reason '{}'",
1849 d.reason
1850 );
1851 }
1852 }
1853
1854 fn sample_mix(rng: &mut Rng, axes: &[(usize, f32)], sigma: f32) -> Vec<f32> {
1856 let mut v: Vec<f32> = (0..HID).map(|_| sigma * rng.next()).collect();
1857 for &(a, w) in axes {
1858 v[a] += w;
1859 }
1860 unit_copy(&v)
1861 }
1862
1863 #[test]
1870 fn a_prompt_nearest_to_a_quarantined_class_runs_the_backbone() {
1871 let mut rng = Rng(0x5eed_1234_abcd_0001);
1872 let draw = |rng: &mut Rng, axes: &[(usize, f32)], n: usize| -> Vec<Vec<f32>> {
1873 (0..n).map(|_| sample_mix(rng, axes, 0.02)).collect()
1874 };
1875 let (gen_ax, herb_ax, mush_ax) = (vec![(0, 1.0)], vec![(1, 1.0)], vec![(1, 1.0), (4, 0.8)]);
1876 let (gt, gh) = (draw(&mut rng, &gen_ax, 200), draw(&mut rng, &gen_ax, 100));
1877 let (ht, hh) = (draw(&mut rng, &herb_ax, 200), draw(&mut rng, &herb_ax, 100));
1878 let (mt, mh) = (draw(&mut rng, &mush_ax, 200), draw(&mut rng, &mush_ax, 100));
1879 let mut h = header();
1880 h.router = Some(policy(descriptor(>, &gh, 3)));
1881 h.skills.push(skill("herbs", descriptor(&ht, &hh, 4), "active"));
1884 h.skills
1885 .push(skill("mushrooms", descriptor(&mt, &mh, 5), "quarantine"));
1886 let (mut cal, _) = calibrate_v2(&h, 0.05).unwrap();
1887 cal.novelty_theta = 0.999;
1889 h.routing = Some(cal.clone());
1890 let hash = skills_hash(&h);
1891 h.router.as_mut().unwrap().skills_hash = format!("{hash:016x}");
1892
1893 let (mut to_backbone, mut old_to_herbs) = (0usize, 0usize);
1894 let n = 100;
1895 for _ in 0..n {
1896 let q = sample_mix(&mut rng, &mush_ax, 0.02);
1897 let d = route_policy(&h, &q);
1898 assert_eq!(
1899 d.nearest_skill(),
1900 Some("mushrooms"),
1901 "the quarantined class is scored: {}",
1902 d.reason
1903 );
1904 if d.target == RouteTarget::Backbone {
1905 to_backbone += 1;
1906 assert!(d.reason.contains("not routable"), "{}", d.reason);
1907 }
1908 let (base, rows) = policy_rows(&h, &q);
1910 let herbs_only: Vec<ErrorRow> = rows.into_iter().filter(|r| r.0 == "herbs").collect();
1911 let old = decide_backbone_gated(base.as_ref(), &herbs_only, Some(&cal), h.router.as_ref(), hash);
1912 old_to_herbs += usize::from(old.skill() == Some("herbs"));
1913 let dq = route_policy_with(&h, &q, RouteOptions { include_quarantine: true });
1915 assert_eq!(dq.skill(), Some("mushrooms"), "{}", dq.reason);
1916 }
1917 assert_eq!(to_backbone, n, "a mushroom prompt ran a skill");
1918 assert!(
1919 old_to_herbs * 10 >= n * 9,
1920 "the pre-fix decision did not misroute ({old_to_herbs}/{n}) — the regression is untested"
1921 );
1922 for _ in 0..50 {
1924 let d = route_policy(&h, &sample_mix(&mut rng, &herb_ax, 0.02));
1925 assert_eq!(d.skill(), Some("herbs"), "{}", d.reason);
1926 }
1927 h.skills[1].status = Some("stale_regate".into());
1929 let q = sample_mix(&mut rng, &mush_ax, 0.02);
1930 assert_eq!(route_policy(&h, &q).target, RouteTarget::Backbone);
1931 assert_eq!(
1932 route_policy_with(&h, &q, RouteOptions { include_quarantine: true }).skill(),
1933 Some("mushrooms")
1934 );
1935 assert!(no_candidate_reason(&h, RouteOptions::default()).is_none());
1936 h.skills[0].status = Some("quarantine".into());
1937 assert!(
1938 no_candidate_reason(&h, RouteOptions::default())
1939 .is_some_and(|r| r.contains("no routable skill class"))
1940 );
1941 assert!(no_candidate_reason(&h, RouteOptions { include_quarantine: true }).is_none());
1942 }
1943
1944 #[test]
1945 fn prompt_contract_is_enforced_in_one_place() {
1946 let mut h = header();
1947 h.skills.push(SkillRecord {
1948 prompt_contract: Some(PROMPT_CONTRACT_CMF_IM_V1.into()),
1949 ..skill("herbs", truth_policy().base, "active")
1950 });
1951 h.skills.push(SkillRecord {
1952 prompt_contract: Some("other-v9".into()),
1953 ..skill("odd", truth_policy().base, "active")
1954 });
1955 h.skills.push(skill("free", truth_policy().base, "active"));
1956 let pick = |id: &str| RouteDecision::forced(RouteTarget::Skill(id.into()), "t");
1957 let (d, r) = enforce_prompt_contract(&h, pick("herbs"), PromptFrame::CmfImV1, false);
1959 assert_eq!((d.skill(), r), (Some("herbs"), false));
1960 let (d, r) = enforce_prompt_contract(&h, pick("herbs"), PromptFrame::Raw, true);
1961 assert_eq!((d.skill(), r), (Some("herbs"), true));
1962 for (id, frame, can) in [
1964 ("herbs", PromptFrame::Raw, false),
1965 ("herbs", PromptFrame::Other, true),
1966 ("odd", PromptFrame::CmfImV1, true),
1967 ] {
1968 let (d, r) = enforce_prompt_contract(&h, pick(id), frame, can);
1969 assert_eq!(d.target, RouteTarget::Backbone, "{id} {frame:?}");
1970 assert!(!r && d.reason.contains("not satisfied"), "{}", d.reason);
1971 }
1972 let (d, r) = enforce_prompt_contract(&h, pick("free"), PromptFrame::Raw, false);
1974 assert_eq!((d.skill(), r), (Some("free"), false));
1975 let bb = RouteDecision::forced(RouteTarget::Backbone, "t");
1976 assert_eq!(enforce_prompt_contract(&h, bb, PromptFrame::Other, false).0.skill(), None);
1977 let p = render_cmf_im_v1("Что лечит зверобой?");
1979 assert_eq!(p, "<|im_start|>user\nЧто лечит зверобой?<|im_end|>\n<|im_start|>assistant\n");
1980 assert_eq!(cmf_im_v1_single_user_turn(&p), Some("Что лечит зверобой?"));
1981 assert_eq!(cmf_im_v1_single_user_turn("Что лечит зверобой?"), None);
1982 let two = format!("{}{}", render_cmf_im_v1("a"), "b<|im_end|>\n<|im_start|>user\nc<|im_end|>\n<|im_start|>assistant\n");
1983 assert_eq!(cmf_im_v1_single_user_turn(&two), None, "history is not one turn");
1984 assert_eq!(
1985 cmf_im_v1_single_user_turn(
1986 "<|im_start|>system\ns<|im_end|>\n<|im_start|>user\nq<|im_end|>\n<|im_start|>assistant\n"
1987 ),
1988 None
1989 );
1990 }
1991
1992 #[test]
1993 fn skills_hash_tracks_descriptors_not_status() {
1994 let mut rng = Rng(7);
1995 let tr: Vec<Vec<f32>> = (0..20).map(|_| sample(&mut rng, 1, 0.1)).collect();
1996 let mut h = header();
1997 h.router = Some(policy(descriptor(&tr, &[], 3)));
1998 h.skills
1999 .push(skill("herbs", descriptor(&tr, &[], 4), "quarantine"));
2000 let h0 = skills_hash(&h);
2001 h.skills[0].status = Some("active".into());
2002 h.skills[0].gate = Some(serde_json::json!({"status": "measured", "x": 1}));
2003 h.skills[0].selection.as_mut().unwrap().holdout = Some("AAAA".into());
2004 assert_eq!(
2005 skills_hash(&h),
2006 h0,
2007 "status/gate/holdout do not invalidate the calibration"
2008 );
2009 h.skills.push(SkillRecord {
2011 id: "legacy".into(),
2012 selection: Some(descriptor(&tr, &[], 5)),
2013 ..Default::default()
2014 });
2015 assert_eq!(skills_hash(&h), h0);
2016 let mut h2 = h.clone();
2018 h2.skills[0].selection.as_mut().unwrap().err_mean = Some(0.123);
2019 assert_ne!(skills_hash(&h2), h0);
2020 let mut h3 = h.clone();
2021 h3.skills
2022 .push(skill("more", descriptor(&tr, &[], 6), "quarantine"));
2023 assert_ne!(skills_hash(&h3), h0);
2024 let mut h4 = h.clone();
2025 h4.skills[0].status = Some("retired".into());
2026 assert_ne!(skills_hash(&h4), h0);
2027 let mut h5 = h.clone();
2028 h5.router.as_mut().unwrap().base.rank = 0;
2029 assert_ne!(skills_hash(&h5), h0);
2030 }
2031
2032 #[test]
2033 fn phi_span_is_the_user_text_only() {
2034 let p = policy(descriptor(&[vec![1.0; HID]], &[], 0)).phi;
2035 let (ids, span) = phi_span_ids(&p, &[40, 41, 42, 43]);
2036 assert_eq!(ids, vec![1, 7, 8, 40, 41, 42, 43, 2, 1, 9]);
2037 assert_eq!(span, 3..7);
2038 assert_eq!(&ids[span.clone()], &[40, 41, 42, 43]);
2039 let hiddens: Vec<f32> = (0..ids.len())
2041 .flat_map(|r| {
2042 (0..HID).map(move |c| {
2043 if (3..7).contains(&r) {
2044 (c == 0) as u8 as f32
2045 } else {
2046 100.0
2047 }
2048 })
2049 })
2050 .collect();
2051 let phi = pool_span_unit(&hiddens, HID, span);
2052 assert!((phi[0] - 1.0).abs() < 1e-6 && phi[1..].iter().all(|x| *x == 0.0));
2053 assert!(
2054 pool_span_unit(&hiddens, HID, 5..5)
2055 .iter()
2056 .all(|x| *x == 0.0)
2057 );
2058 }
2059
2060 #[test]
2061 fn clopper_pearson_matches_reference_values() {
2062 assert!((clopper_pearson_upper(0, 10, 0.95) - 0.308_5).abs() < 1e-3);
2064 assert!((clopper_pearson_upper(5, 10, 0.95) - 0.812_9).abs() < 1e-3);
2065 assert!((clopper_pearson_upper(1, 100, 0.95) - 0.054_5).abs() < 1e-3);
2066 let u = clopper_pearson_upper(0, 500, 0.95);
2067 assert!((u - 0.007_351).abs() < 1e-5, "{u}");
2068 assert_eq!(clopper_pearson_upper(3, 3, 0.95), 1.0);
2069 }
2070
2071 #[test]
2075 fn calibrate_v2_separates_backbone_and_skill() {
2076 let mut rng = Rng(0x9e37_79b9_7f4a_7c15);
2077 let draw = |rng: &mut Rng, axis: usize, n: usize| -> Vec<Vec<f32>> {
2078 (0..n).map(|_| sample(rng, axis, 0.1)).collect()
2079 };
2080 let (gt, gh) = (draw(&mut rng, 0, 200), draw(&mut rng, 0, 120));
2081 let (st, sh) = (draw(&mut rng, 1, 200), draw(&mut rng, 1, 120));
2082 let mut h = header();
2083 h.router = Some(policy(descriptor(>, &gh, 3)));
2084 h.skills
2085 .push(skill("herbs", descriptor(&st, &sh, 4), "active"));
2086
2087 let (cal, measured) = calibrate_v2(&h, 0.01).unwrap();
2088 assert_eq!(measured["false_accept"], 0.0, "{measured}");
2089 assert!(
2090 measured["in_scope_recall"].as_f64().unwrap() >= 0.98,
2091 "{measured}"
2092 );
2093 assert!(
2094 measured["false_accept_upper95"].as_f64().unwrap() < 0.031,
2095 "{measured}"
2096 );
2097 assert_eq!(measured["n_general"], 120);
2098 assert_eq!(measured["n_in"], 120);
2099
2100 h.routing = Some(cal);
2102 let hash = skills_hash(&h);
2103 h.router.as_mut().unwrap().skills_hash = format!("{hash:016x}");
2104 let (mut fa, mut hit) = (0, 0);
2105 for _ in 0..200 {
2106 fa += usize::from(
2107 route_policy(&h, &sample(&mut rng, 0, 0.1))
2108 .skill()
2109 .is_some(),
2110 );
2111 hit +=
2112 usize::from(route_policy(&h, &sample(&mut rng, 1, 0.1)).skill() == Some("herbs"));
2113 }
2114 assert_eq!(fa, 0);
2115 assert!(hit >= 190, "recall {hit}/200");
2116 let d = route_policy(&h, &sample(&mut rng, 9, 0.1));
2118 assert_eq!(d.target, RouteTarget::Backbone, "{}", d.reason);
2119
2120 let mut q = h.clone();
2122 q.skills[0].status = Some("quarantine".into());
2123 assert_eq!(
2124 route_policy(&q, &sample(&mut rng, 1, 0.1)).target,
2125 RouteTarget::Backbone
2126 );
2127 let mut stale = h.clone();
2128 stale
2129 .skills
2130 .push(skill("new", descriptor(&st, &[], 5), "active"));
2131 let d = route_policy(&stale, &sample(&mut rng, 1, 0.1));
2132 assert_eq!(d.target, RouteTarget::Backbone);
2133 assert!(d.reason.contains("stale"), "{}", d.reason);
2134 }
2135}