mod common;
use common::{
assert_all_three_paths_match, assert_decoder_matches_on_all_three_paths, graph_caches,
graph_fixture_path, kl_vs_golden, load_graph_fixture, worst_vs, GRAPH_PROMPT, GRAPH_TOL,
};
use frink_models::capability::{resolve_architecture, ArchPath, WEIGHTED_LAYER_NORM};
use frink_models::config::{ModelConfig, RopeLayout};
use frink_models::loader::LoadError;
use frink_models::norm::NormOp;
use frink_models::parallel_residual::ParallelNorm;
use frink_models::Decoder;
const COMMAND_R: &str = "command_r";
const NOSCALE: &str = "command_r_noscale";
const PLUS: &str = "command_r_plus";
const COMMAND_R_GOLDEN: [f32; 48] = [
0.07120129,
-0.0658564,
0.075167604,
-0.03013523,
-0.005875893,
-0.098735176,
0.090072975,
-0.050787985,
-0.0146826655,
0.016819555,
-0.014875419,
-0.073848434,
-0.12244835,
-0.0037116595,
-0.062389098,
-0.05820402,
-0.04317017,
0.20395382,
-0.06363838,
-0.09944917,
0.035045836,
-0.15262796,
-0.004360944,
-0.10448907,
-0.0140059665,
-0.11419049,
-0.14803372,
0.16377491,
-0.06749434,
0.03222097,
-0.057085104,
0.08956219,
0.0523116,
-0.28689998,
-0.026316267,
-0.0629535,
-0.05267317,
-0.09888828,
-0.034774955,
-0.020797234,
-0.16666314,
-0.1091035,
0.115461305,
-0.105705485,
-0.061054923,
-0.12041834,
0.015898105,
-0.030510504,
];
const COMMAND_R_NOSCALE_GOLDEN: [f32; 48] = [
1.1392206,
-1.0537024,
1.2026817,
-0.48216367,
-0.09401429,
-1.5797628,
1.4411676,
-0.81260777,
-0.23492265,
0.26911288,
-0.23800671,
-1.181575,
-1.9591736,
-0.05938655,
-0.99822557,
-0.93126434,
-0.6907227,
3.263261,
-1.0182141,
-1.5911868,
0.5607334,
-2.4420474,
-0.069775105,
-1.6718252,
-0.22409546,
-1.8270478,
-2.3685396,
2.6203985,
-1.0799094,
0.51553553,
-0.91336167,
1.4329951,
0.8369856,
-4.5903997,
-0.42106026,
-1.007256,
-0.8427707,
-1.5822124,
-0.5563993,
-0.33275574,
-2.6666102,
-1.745656,
1.8473809,
-1.6912878,
-0.97687876,
-1.9266934,
0.25436968,
-0.48816806,
];
fn decode(decoder: &Decoder) -> Vec<f32> {
let mut kv = graph_caches(decoder);
let mut out = Vec::new();
for (pos, &tok) in GRAPH_PROMPT.iter().enumerate() {
out = decoder.forward_token(tok, pos, &mut kv);
}
out
}
#[test]
fn command_r_matches_llama_cpp_on_all_three_paths() {
assert_all_three_paths_match(COMMAND_R, &COMMAND_R_GOLDEN);
}
#[test]
fn command_r_without_the_logit_scale_key_matches_llama_cpp() {
assert_all_three_paths_match(NOSCALE, &COMMAND_R_NOSCALE_GOLDEN);
}
#[test]
fn report_kl_against_llama_cpp() {
for (name, golden) in [
(COMMAND_R, &COMMAND_R_GOLDEN),
(NOSCALE, &COMMAND_R_NOSCALE_GOLDEN),
] {
let out = decode(&load_graph_fixture(name));
println!(
"{name}: KL(llama.cpp || frink) = {:.3e}, max |delta| = {:.3e}",
kl_vs_golden(&out, golden),
worst_vs(&out, golden)
);
}
}
#[test]
fn the_loaded_layers_are_the_graph() {
assert!(WEIGHTED_LAYER_NORM.contains(&"command-r"));
assert!(matches!(
resolve_architecture("command-r"),
Some(ArchPath::GenericGqa {
rope: RopeLayout::Norm
})
));
let d = load_graph_fixture(COMMAND_R);
assert!(d.config.parallel_residual);
assert_eq!(d.config.logit_multiplier, Some(0.0625));
assert_eq!(d.config.rope_theta, 8_000_000.0);
assert_eq!(d.config.rms_norm_eps, 1e-5, "attention.layer_norm_epsilon");
let file = frink_gguf::GgufFile::open(graph_fixture_path(COMMAND_R)).unwrap();
assert!(file.find_tensor("output.weight").is_none());
assert_eq!(
(d.output_head.rows(), d.output_head.cols()),
(d.embedding.rows(), d.embedding.cols())
);
for layer in &d.layers {
assert_eq!(layer.moe.parallel, Some(ParallelNorm::SharedNorm));
assert!(matches!(layer.attn.norm_weight, NormOp::LayerNorm(_)));
assert!(matches!(layer.moe.norm_weight, NormOp::None));
assert!(layer.attn.q_norm.is_none() && layer.attn.k_norm.is_none());
}
assert!(matches!(d.final_norm, NormOp::LayerNorm(_)));
let n = load_graph_fixture(NOSCALE);
assert_eq!(n.config.logit_multiplier, None, "absent means no scale");
let ratio = COMMAND_R_GOLDEN[3] / COMMAND_R_NOSCALE_GOLDEN[3];
assert!((ratio - 0.0625).abs() < 1e-6, "{ratio}");
}
#[test]
fn command_r_plus_is_refused_by_name() {
let file = frink_gguf::GgufFile::open(graph_fixture_path(PLUS)).expect("fixture opens");
let q = file
.find_tensor("blk.63.attn_q_norm.weight")
.expect("REQUIRED at 64 layers");
assert_eq!(
q.shape.iter().product::<u64>(),
16,
"{{8, 2}}: two heads of eight"
);
let config = ModelConfig::from_gguf(&file).expect("the header is the served shape's");
assert_eq!(config.n_layers, 64);
match Decoder::from_gguf(graph_fixture_path(PLUS), config) {
Err(LoadError::UnsupportedFeature(_, msg)) => {
assert!(msg.contains("per-head LayerNorm"), "{msg}");
assert!(msg.contains("command-r.cpp:28-31,80,87"), "{msg}");
}
Err(other) => panic!("expected the QK LayerNorm refused by name, got {other:?}"),
Ok(_) => panic!("Command-R+ loaded"),
}
}
#[test]
fn each_seam_is_visible_in_the_logits() {
let mut d = load_graph_fixture(COMMAND_R);
assert_decoder_matches_on_all_three_paths(&d, &COMMAND_R_GOLDEN, GRAPH_TOL, "baseline");
let weight_of = |op: &NormOp| -> Vec<f32> {
let NormOp::LayerNorm(w) = op else {
unreachable!()
};
w.clone()
};
let saved: Vec<NormOp> = d
.layers
.iter_mut()
.map(|l| {
let w = weight_of(&l.attn.norm_weight);
std::mem::replace(&mut l.attn.norm_weight, NormOp::Rms(w))
})
.collect();
let worst = worst_vs(&decode(&d), &COMMAND_R_GOLDEN);
assert!(worst > 1e-2, "LayerNorm vs RMSNorm not seen: {worst}");
for (l, s) in d.layers.iter_mut().zip(saved) {
l.attn.norm_weight = s;
}
for l in d.layers.iter_mut() {
l.moe.parallel = None;
}
let worst = worst_vs(&decode(&d), &COMMAND_R_GOLDEN);
assert!(worst > 1e-2, "the parallel residual not seen: {worst}");
for l in d.layers.iter_mut() {
l.moe.parallel = Some(ParallelNorm::SharedNorm);
}
let saved = d.config.logit_multiplier.take();
let worst = worst_vs(&decode(&d), &COMMAND_R_GOLDEN);
assert!(worst > 1.0, "the logit multiply not seen: {worst}");
d.config.logit_multiplier = saved;
assert_decoder_matches_on_all_three_paths(&d, &COMMAND_R_GOLDEN, GRAPH_TOL, "restored");
}