use crate::npy;
use cortiq_core::format::{CMF_VERSION, CmfHeader, CmfModel, TensorSpec, TokenizerBundle};
use cortiq_core::quant::{bf16_to_f32, f16_to_f32, f32_to_f16};
use cortiq_core::types::{LayerType, LinearCoreConfig, ModelArch, MoeConfig, NormStyle, QuantType, TensorDtype, YarnConfig};
use std::collections::HashMap;
use std::fs;
use std::io::Read;
use std::path::Path;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::Duration;
const GROUP_SIZE: usize = 32;
const F16_TINY: f32 = 6.103_515_6e-5;
fn f16_scale(raw: f32) -> f32 {
f16_to_f32(f32_to_f16(raw)).max(F16_TINY)
}
pub(crate) fn canon_name(raw: &str) -> Option<String> {
if raw.contains(".visual.") || raw.starts_with("visual.") || raw.starts_with("mtp.") || raw.contains(".mtp.") {
return None;
}
for pfx in ["model.vision_embedder.", "model.embed_audio.", "model.embed_vision."] {
if raw.starts_with(pfx) {
return None;
}
}
for pfx in ["model.language_model.", "language_model.model.", "language_model."] {
if let Some(rest) = raw.strip_prefix(pfx) {
return Some(lfm2_canon(&format!("model.{rest}")));
}
}
if raw.ends_with(".mlp.experts.e_score_correction_bias") {
return Some(raw.replace(".mlp.experts.e_score_correction_bias", ".mlp.expert_bias"));
}
Some(lfm2_canon(raw))
}
fn lfm2_canon(name: &str) -> String {
let is_lfm2 =
name == "model.embedding_norm.weight" || name.contains(".operator_norm") || name.contains(".ffn_norm") || name.contains(".feed_forward.") || name.contains(".conv.") || name.contains(".self_attn.out_proj") || name.contains(".self_attn.q_layernorm") || name.contains(".self_attn.k_layernorm");
if !is_lfm2 {
return name.to_string();
}
if name == "model.embedding_norm.weight" {
return "model.norm.weight".to_string();
}
let mut n = name.to_string();
n = n.replace(".operator_norm.", ".input_layernorm.");
n = n.replace(".ffn_norm.", ".post_attention_layernorm.");
n = n.replace(".self_attn.out_proj.", ".self_attn.o_proj.");
n = n.replace(".self_attn.q_layernorm.", ".self_attn.q_norm.");
n = n.replace(".self_attn.k_layernorm.", ".self_attn.k_norm.");
n = n.replace(".conv.in_proj.", ".short_conv.in_proj.");
n = n.replace(".conv.out_proj.", ".short_conv.out_proj.");
n = n.replace(".conv.conv.", ".short_conv.conv.");
n = n.replace(".feed_forward.gate.weight", ".mlp.gate.weight");
n = n.replace(".feed_forward.expert_bias", ".mlp.expert_bias");
n = n.replace(".feed_forward.experts.", ".mlp.experts.");
n = n.replace(".feed_forward.", ".mlp.");
n = n.replace(".w1.weight", ".gate_proj.weight");
n = n.replace(".w3.weight", ".up_proj.weight");
n = n.replace(".w2.weight", ".down_proj.weight");
n
}
fn force_f16(name: &str) -> bool {
name.ends_with("linear_attn.in_proj_a.weight") || name.ends_with("linear_attn.in_proj_b.weight") || name.ends_with("mlp.gate.weight") || name.ends_with("shared_expert_gate.weight") || name.ends_with("self_attn.g_proj.weight")
}
#[derive(Clone, Copy, PartialEq)]
pub(crate) enum Quant {
Q8Row,
Q8_2f,
Q4Block,
F16,
Vbit,
Q4Tiled,
Q1,
Q1p,
Q1s,
Q1t,
}
pub(crate) fn quantize_2d(quant: Quant, vals: &[f32], out_dim: usize, in_dim: usize) -> (TensorDtype, Vec<u8>) {
match quant {
Quant::Q8Row => (TensorDtype::Q8Row, encode_q8_row(vals, out_dim, in_dim)),
Quant::Q8_2f => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
Quant::Q4Block => (TensorDtype::Q4Block, encode_q4_block(vals)),
Quant::F16 => (TensorDtype::F16, encode_f16(vals)),
Quant::Q4Tiled if in_dim % GROUP_SIZE == 0 => (TensorDtype::Q4Tiled, encode_q4_tiled(vals, out_dim, in_dim)),
Quant::Q4Tiled => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
Quant::Vbit if in_dim % GROUP_SIZE == 0 => (TensorDtype::VbitRo, encode_vbit_ro(vals, out_dim, in_dim)),
Quant::Vbit => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
Quant::Q1 if in_dim % GROUP_SIZE == 0 => (TensorDtype::Q1, encode_q1(vals, out_dim, in_dim)),
Quant::Q1 => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
Quant::Q1p if in_dim % GROUP_SIZE == 0 => (TensorDtype::Q1, encode_q1_ef(vals, out_dim, in_dim)),
Quant::Q1p => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
Quant::Q1s if in_dim % GROUP_SIZE == 0 => (TensorDtype::Q1S, encode_q1s(vals, out_dim, in_dim, q1s_keep_frac())),
Quant::Q1s => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
Quant::Q1t if in_dim % GROUP_SIZE == 0 => (TensorDtype::Q1T, crate::gptq::quantize_q1t(vals, out_dim, in_dim, &vec![1.0; in_dim], 0.0)),
Quant::Q1t => (TensorDtype::Q8_2f, encode_q8_2f(vals, out_dim, in_dim)),
}
}
pub(crate) fn parse_quant(s: &str) -> anyhow::Result<Quant> {
Ok(match s.to_ascii_lowercase().as_str() {
"q8" | "q8_row" | "q8row" => Quant::Q8Row,
"q8_2f" | "q82f" | "q8f" => Quant::Q8_2f,
"q4" | "q4_block" | "q4block" => Quant::Q4Block,
"f16" | "fp16" => Quant::F16,
"vbit" | "v_bit" => Quant::Vbit,
"q4t" | "q4_tiled" => Quant::Q4Tiled,
"q1" => Quant::Q1,
"q1p" | "q1_ptq" => Quant::Q1p,
"q1s" | "q1_mask" => Quant::Q1s,
"q1t" | "q1_ternary" => Quant::Q1t,
other => anyhow::bail!("unknown quant '{other}' (use q8, q8_2f, q4, q4t, f16, vbit, q1, q1p, q1s, or q1t)"),
})
}
fn encode_q8_row(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
let mut q = Vec::with_capacity(out_dim * in_dim);
let mut scales = Vec::with_capacity(out_dim * 2);
for o in 0..out_dim {
let row = &vals[o * in_dim..(o + 1) * in_dim];
let absmax = row.iter().fold(0f32, |m, v| m.max(v.abs()));
let scale = f16_scale(absmax / 127.0);
for &v in row {
q.push((v / scale).round_ties_even().clamp(-128.0, 127.0) as i8 as u8);
}
scales.extend_from_slice(&f32_to_f16(scale).to_le_bytes());
}
q.extend_from_slice(&scales);
q
}
fn encode_q4_block(vals: &[f32]) -> Vec<u8> {
let n_groups = vals.len().div_ceil(GROUP_SIZE);
let mut padded = vals.to_vec();
padded.resize(n_groups * GROUP_SIZE, 0.0);
let mut packed = Vec::with_capacity(n_groups * 16);
let mut scales = Vec::with_capacity(n_groups * 2);
for g in 0..n_groups {
let group = &padded[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
let absmax = group.iter().fold(0f32, |m, v| m.max(v.abs()));
let scale = f16_scale(absmax / 7.0);
for k in 0..16 {
let q0 = ((group[k * 2] / scale).round_ties_even().clamp(-8.0, 7.0) as i8 + 8) as u8;
let q1 = ((group[k * 2 + 1] / scale).round_ties_even().clamp(-8.0, 7.0) as i8 + 8) as u8;
packed.push((q0 & 0x0F) | (q1 << 4));
}
scales.extend_from_slice(&f32_to_f16(scale).to_le_bytes());
}
packed.extend_from_slice(&scales);
packed
}
fn encode_q4_tiled(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
debug_assert_eq!(vals.len(), out_dim * in_dim);
debug_assert_eq!(in_dim % GROUP_SIZE, 0);
let legacy = encode_q4_block(vals);
let n_groups = vals.len() / GROUP_SIZE;
let (packed, scales) = legacy.split_at(n_groups * 16);
let mut out = Vec::with_capacity(n_groups * 18);
for g in 0..n_groups {
out.extend_from_slice(&scales[g * 2..g * 2 + 2]);
out.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
}
out
}
fn encode_q1(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
debug_assert_eq!(vals.len(), out_dim * in_dim);
debug_assert_eq!(in_dim % GROUP_SIZE, 0);
let n_groups = vals.len() / GROUP_SIZE;
let mut out = Vec::with_capacity(n_groups * 6);
for g in 0..n_groups {
let grp = &vals[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
let mean_abs = grp.iter().map(|v| v.abs()).sum::<f32>() / GROUP_SIZE as f32;
let s = f16_scale(mean_abs);
out.extend_from_slice(&f32_to_f16(s).to_le_bytes());
for j in 0..GROUP_SIZE / 8 {
let mut byte = 0u8;
for k in 0..8 {
if grp[j * 8 + k] >= 0.0 {
byte |= 1 << k;
}
}
out.push(byte);
}
}
out
}
fn encode_q1_ef(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
debug_assert_eq!(vals.len(), out_dim * in_dim);
debug_assert_eq!(in_dim % GROUP_SIZE, 0);
let groups_per_row = in_dim / GROUP_SIZE;
let n_groups = vals.len() / GROUP_SIZE;
let mut out = Vec::with_capacity(n_groups * 6);
let mut carry = 0.0f32;
for g in 0..n_groups {
if g % groups_per_row == 0 {
carry = 0.0; }
let grp = &vals[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
let mean_abs = grp.iter().map(|v| v.abs()).sum::<f32>() / GROUP_SIZE as f32;
let s = f16_scale(mean_abs);
out.extend_from_slice(&f32_to_f16(s).to_le_bytes());
for j in 0..GROUP_SIZE / 8 {
let mut byte = 0u8;
for k in 0..8 {
let w = grp[j * 8 + k];
let v = w + carry;
let bit = v >= 0.0;
if bit {
byte |= 1 << k;
}
carry = v - if bit { s } else { -s };
}
out.push(byte);
}
}
out
}
fn q1s_keep_frac() -> f32 {
std::env::var("CMF_Q1S_KEEP").ok().and_then(|v| v.parse::<f32>().ok()).unwrap_or(0.01).clamp(0.0, 0.25)
}
fn encode_q1s(vals: &[f32], out_dim: usize, in_dim: usize, keep_frac: f32) -> Vec<u8> {
debug_assert_eq!(vals.len(), out_dim * in_dim);
debug_assert_eq!(in_dim % GROUP_SIZE, 0);
let n = vals.len();
let n_out = (((n as f32) * keep_frac).round() as usize).min(n);
let threshold = if n_out == 0 {
f32::INFINITY
} else {
let mut absv: Vec<f32> = vals.iter().map(|v| v.abs()).collect();
let k = n - n_out;
absv.select_nth_unstable_by(k, |a, b| a.partial_cmp(b).unwrap());
absv[k]
};
let is_out: Vec<bool> = (0..n).map(|i| n_out > 0 && vals[i].abs() >= threshold).collect();
let groups_per_row = in_dim / GROUP_SIZE;
let n_groups = n / GROUP_SIZE;
let n_out_actual = is_out.iter().filter(|&&o| o).count();
let mut out = Vec::with_capacity(n_groups * 6 + 4 + n_out_actual * 6);
let mut carry = 0.0f32;
for g in 0..n_groups {
if g % groups_per_row == 0 {
carry = 0.0;
}
let base = g * GROUP_SIZE;
let mut sum = 0.0f32;
let mut cnt = 0usize;
for j in 0..GROUP_SIZE {
if !is_out[base + j] {
sum += vals[base + j].abs();
cnt += 1;
}
}
let s = f16_scale(if cnt > 0 { sum / cnt as f32 } else { 0.0 });
out.extend_from_slice(&f32_to_f16(s).to_le_bytes());
for jb in 0..GROUP_SIZE / 8 {
let mut byte = 0u8;
for k in 0..8 {
let i = base + jb * 8 + k;
if is_out[i] {
if vals[i] >= 0.0 {
byte |= 1 << k;
}
} else {
let v = vals[i] + carry;
let bit = v >= 0.0;
if bit {
byte |= 1 << k;
}
carry = v - if bit { s } else { -s };
}
}
out.push(byte);
}
}
out.extend_from_slice(&(n_out_actual as u32).to_le_bytes());
for (i, &o) in is_out.iter().enumerate() {
if o {
out.extend_from_slice(&(i as u32).to_le_bytes());
out.extend_from_slice(&f32_to_f16(vals[i]).to_le_bytes());
}
}
out
}
fn encode_q8_2f(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
let mut col = vec![0f32; in_dim];
for (i, c) in col.iter_mut().enumerate() {
let mut acc = 0f64;
for o in 0..out_dim {
let v = vals[o * in_dim + i] as f64;
acc += v * v;
}
let rms = (acc / out_dim as f64).sqrt().max(1e-12) as f32;
*c = f16_to_f32(f32_to_f16(rms)).max(F16_TINY);
}
let mut q = Vec::with_capacity(out_dim * in_dim);
let mut scales = Vec::with_capacity(out_dim * 2);
for o in 0..out_dim {
let mut absmax = 0f32;
for i in 0..in_dim {
absmax = absmax.max((vals[o * in_dim + i] / col[i]).abs());
}
let scale = f16_scale(absmax.max(1e-12) / 127.0);
for i in 0..in_dim {
let wn = vals[o * in_dim + i] / col[i];
q.push((wn / scale).round_ties_even().clamp(-127.0, 127.0) as i8 as u8);
}
scales.extend_from_slice(&f32_to_f16(scale).to_le_bytes());
}
let mut out = q;
out.extend_from_slice(&scales);
for &c in &col {
out.extend_from_slice(&f32_to_f16(c).to_le_bytes());
}
out
}
const VBIT_LEVELS: [u8; 5] = [3, 4, 5, 6, 8];
static VBIT_MEAN_BITS_MILLI: AtomicU32 = AtomicU32::new(4250);
pub fn set_vbit_mean_bits(bits: f32) {
VBIT_MEAN_BITS_MILLI.store((bits.clamp(3.0, 8.0) * 1000.0) as u32, Ordering::Relaxed);
}
fn vbit_mean_bits() -> f32 {
VBIT_MEAN_BITS_MILLI.load(Ordering::Relaxed) as f32 / 1000.0
}
fn vbit_snap_level(x: f32) -> u8 {
let mut best = VBIT_LEVELS[0];
let mut bestd = (x - best as f32).abs();
for &lv in &VBIT_LEVELS[1..] {
let d = (x - lv as f32).abs();
if d < bestd {
bestd = d;
best = lv;
}
}
best
}
fn vbit_bits(vals: &[f32], out_dim: usize, in_dim: usize, mean_bits: f32) -> Vec<u8> {
let a: Vec<f32> = (0..out_dim)
.map(|o| {
let mx = vals[o * in_dim..(o + 1) * in_dim].iter().fold(0f32, |m, v| m.max(v.abs()));
mx.max(1e-12).log2()
})
.collect();
let amean = a.iter().sum::<f32>() / out_dim as f32;
a.iter().map(|&ar| vbit_snap_level(mean_bits + (ar - amean)).max(3)).collect()
}
struct BitWriter {
buf: Vec<u8>,
cur: u8,
nbits: u8,
}
impl BitWriter {
fn with_capacity(n: usize) -> Self {
Self { buf: Vec::with_capacity(n), cur: 0, nbits: 0 }
}
fn push(&mut self, v: u32, b: u32) {
for i in (0..b).rev() {
self.cur = (self.cur << 1) | ((v >> i) & 1) as u8;
self.nbits += 1;
if self.nbits == 8 {
self.buf.push(self.cur);
self.cur = 0;
self.nbits = 0;
}
}
}
fn flush_row(&mut self) {
if self.nbits > 0 {
self.buf.push(self.cur << (8 - self.nbits));
self.cur = 0;
self.nbits = 0;
}
}
}
fn encode_vbit(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
let ng = in_dim / GROUP_SIZE;
let bits = vbit_bits(vals, out_dim, in_dim, vbit_mean_bits());
let mut scale = vec![0f32; out_dim * ng];
let mut sc_bytes = Vec::with_capacity(out_dim * ng * 2);
for o in 0..out_dim {
let l = (2f32.powi(bits[o] as i32 - 1) - 1.0).max(1.0);
for g in 0..ng {
let base = o * in_dim + g * GROUP_SIZE;
let mx = vals[base..base + GROUP_SIZE].iter().fold(0f32, |m, v| m.max(v.abs()));
let s = f16_scale(mx / l);
scale[o * ng + g] = s;
sc_bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
}
}
let mut out = Vec::with_capacity(out_dim + sc_bytes.len() + out_dim * in_dim);
out.extend_from_slice(&bits);
out.extend_from_slice(&sc_bytes);
let mut bw = BitWriter::with_capacity(out_dim * in_dim);
for o in 0..out_dim {
let b = bits[o] as u32;
let l = 2f32.powi(bits[o] as i32 - 1) - 1.0;
let maxq = 2f32.powi(bits[o] as i32) - 1.0;
for c in 0..in_dim {
let s = scale[o * ng + c / GROUP_SIZE];
let q = ((vals[o * in_dim + c] / s).round_ties_even() + l).clamp(0.0, maxq) as u32;
bw.push(q, b);
}
bw.flush_row();
}
out.extend_from_slice(&bw.buf);
out
}
fn encode_vbit_ro(vals: &[f32], out_dim: usize, in_dim: usize) -> Vec<u8> {
let legacy = encode_vbit(vals, out_dim, in_dim);
let ng = in_dim / GROUP_SIZE;
let sc_len = out_dim * ng * 2;
let (head, packed) = legacy.split_at(out_dim + sc_len);
let bits = &head[..out_dim];
let mut out = Vec::with_capacity(legacy.len() + (out_dim + 1) * 4);
out.extend_from_slice(head);
let mut off = 0u32;
for &b in bits {
out.extend_from_slice(&off.to_le_bytes());
off += ((in_dim * b as usize).div_ceil(8)) as u32;
}
out.extend_from_slice(&off.to_le_bytes());
debug_assert_eq!(off as usize, packed.len());
out.extend_from_slice(packed);
out
}
pub(crate) fn encode_f16(vals: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(vals.len() * 2);
for &v in vals {
out.extend_from_slice(&f32_to_f16(v).to_le_bytes());
}
out
}
pub(crate) fn to_f32(dtype: &str, raw: &[u8]) -> anyhow::Result<Vec<f32>> {
Ok(match dtype {
"F32" => raw.chunks_exact(4).map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]])).collect(),
"F16" => raw.chunks_exact(2).map(|b| f16_to_f32(u16::from_le_bytes([b[0], b[1]]))).collect(),
"BF16" => raw.chunks_exact(2).map(|b| bf16_to_f32(u16::from_le_bytes([b[0], b[1]]))).collect(),
other => anyhow::bail!("unsupported safetensors dtype '{other}' (need F32/F16/BF16)"),
})
}
pub(crate) fn unpack_mlx(w_raw: &[u8], s_raw: &[u8], b_raw: Option<&[u8]>, out_dim: usize, in_dim: usize, bits: usize) -> anyhow::Result<Vec<f32>> {
let mut out = vec![0f32; out_dim * in_dim];
let num_groups = s_raw.len() / 2 / out_dim;
let group_size = in_dim / num_groups;
let w_u32: Vec<u32> = w_raw.chunks_exact(4).map(|b| u32::from_le_bytes([b[0], b[1], b[2], b[3]])).collect();
let s_f16: Vec<u16> = s_raw.chunks_exact(2).map(|b| u16::from_le_bytes([b[0], b[1]])).collect();
let b_f16: Option<Vec<u16>> = b_raw.map(|r| r.chunks_exact(2).map(|b| u16::from_le_bytes([b[0], b[1]])).collect());
let vals_per_u32 = 32 / bits;
let mask = (1 << bits) - 1;
for row in 0..out_dim {
for col in 0..in_dim {
let group = col / group_size;
let scale = f16_to_f32(s_f16[row * num_groups + group]);
let bias = b_f16.as_ref().map(|b| f16_to_f32(b[row * num_groups + group])).unwrap_or(0.0);
let u32_idx = (row * in_dim + col) / vals_per_u32;
let shift = (col % vals_per_u32) * bits;
let val = (w_u32[u32_idx] >> shift) & mask;
out[row * in_dim + col] = (val as f32) * scale + bias;
}
}
Ok(out)
}
pub(crate) fn exec_order_key(name: &str) -> (u32, u32, u32, u32, u32) {
let num_after = |marker: &str| name.split(marker).nth(1).and_then(|s| s.split('.').next()).and_then(|s| s.parse::<u32>().ok());
let expert = num_after(".experts.").unwrap_or(0);
let proj = if name.contains("q_proj") || name.contains("gate_proj") {
0
} else if name.contains("k_proj") || name.contains("up_proj") {
1
} else if name.contains("v_proj") || name.contains("down_proj") {
2
} else if name.contains("o_proj") {
3
} else {
4
};
if name.contains("embed_tokens") {
(0, 0, 0, 0, 0)
} else if let Some(l) = num_after(".layers.") {
let group = if name.contains("input_layernorm") {
0
} else if name.contains("self_attn") || name.contains("linear_attn") || name.contains("short_conv") {
1
} else if name.contains("post_attention_layernorm") {
2
} else if name.ends_with("mlp.gate.weight") || name.contains("shared_expert") || name.contains("expert_bias") {
3 } else if name.contains(".experts.") {
4
} else {
5 };
(1, l, group, expert, proj)
} else if name.contains("model.mtp") {
(4, 0, 0, 0, 0)
} else if name.contains("lm_head") {
(3, 0, 0, 0, 0)
} else if name.contains("model.norm") || name.ends_with("norm.weight") {
(2, 0, 0, 0, 0)
} else {
(5, 0, 0, 0, 0)
}
}
pub(crate) struct TensorMeta {
pub(crate) name: String,
pub(crate) dtype: String,
pub(crate) shape: Vec<usize>,
pub(crate) start: usize,
pub(crate) end: usize,
}
pub(crate) struct SafeTensors {
mmap: memmap2::Mmap,
data_start: usize,
pub(crate) tensors: Vec<TensorMeta>,
}
impl SafeTensors {
pub(crate) fn bytes(&self, m: &TensorMeta) -> &[u8] {
&self.mmap[self.data_start + m.start..self.data_start + m.end]
}
}
fn open_safetensors(path: &Path) -> anyhow::Result<SafeTensors> {
let file = fs::File::open(path).map_err(|e| anyhow::anyhow!("open {}: {e}", path.display()))?;
let mmap = unsafe { memmap2::Mmap::map(&file)? };
if mmap.len() < 8 {
anyhow::bail!("{}: too small to be safetensors", path.display());
}
let hlen = u64::from_le_bytes(mmap[0..8].try_into().unwrap()) as usize;
let header: serde_json::Value = serde_json::from_slice(&mmap[8..8 + hlen])?;
let data_start = 8 + hlen;
let obj = header.as_object().ok_or_else(|| anyhow::anyhow!("bad safetensors header"))?;
let mut tensors = Vec::new();
for (name, v) in obj {
if name == "__metadata__" {
continue;
}
let dtype = v["dtype"].as_str().unwrap_or("").to_string();
let shape: Vec<usize> = v["shape"].as_array().map(|a| a.iter().map(|x| x.as_u64().unwrap_or(0) as usize).collect()).unwrap_or_default();
let offs = v["data_offsets"].as_array().ok_or_else(|| anyhow::anyhow!("tensor '{name}': no data_offsets"))?;
let start = offs[0].as_u64().unwrap_or(0) as usize;
let end = offs[1].as_u64().unwrap_or(0) as usize;
tensors.push(TensorMeta { name: name.clone(), dtype, shape, start, end });
}
Ok(SafeTensors { mmap, data_start, tensors })
}
pub(crate) fn open_model(dir: &Path) -> anyhow::Result<Vec<SafeTensors>> {
let single = dir.join("model.safetensors");
if single.exists() {
return Ok(vec![open_safetensors(&single)?]);
}
let index = dir.join("model.safetensors.index.json");
if index.exists() {
let idx: serde_json::Value = serde_json::from_slice(&fs::read(&index)?)?;
let map = idx["weight_map"].as_object().ok_or_else(|| anyhow::anyhow!("bad index json"))?;
let mut files: Vec<String> = map.values().filter_map(|v| v.as_str().map(String::from)).collect();
files.sort();
files.dedup();
return files.iter().map(|f| open_safetensors(&dir.join(f))).collect();
}
anyhow::bail!("no model.safetensors or model.safetensors.index.json in {}", dir.display())
}
fn cfg_usize(c: &serde_json::Value, key: &str) -> Option<usize> {
c.get(key).and_then(|v| v.as_u64()).map(|x| x as usize)
}
fn build_arch(config: &serde_json::Value) -> anyhow::Result<ModelArch> {
let tc = config.get("text_config").unwrap_or(config);
let model_type = config.get("model_type").and_then(|v| v.as_str()).unwrap_or("unknown").to_string();
let hidden = cfg_usize(tc, "hidden_size").ok_or_else(|| anyhow::anyhow!("config: missing hidden_size"))?;
let n_heads = cfg_usize(tc, "num_attention_heads").ok_or_else(|| anyhow::anyhow!("config: missing num_attention_heads"))?;
let n_layers = cfg_usize(tc, "num_hidden_layers").ok_or_else(|| anyhow::anyhow!("config: missing num_hidden_layers"))?;
let layer_types: Vec<LayerType> = match tc.get("layer_types").and_then(|v| v.as_array()) {
Some(a) => a
.iter()
.map(|v| match v.as_str() {
Some("linear_attention") => LayerType::LinearAttention,
Some("conv") | Some("short_conv") => LayerType::ShortConv,
Some("sliding_attention") => LayerType::SlidingAttention,
_ => LayerType::FullAttention,
})
.collect(),
None => vec![LayerType::FullAttention; n_layers],
};
let has_linear = layer_types.iter().any(|t| matches!(t, LayerType::LinearAttention));
let lnv = cfg_usize(tc, "linear_num_value_heads");
let lvd = cfg_usize(tc, "linear_value_head_dim");
let linear_core = if has_linear {
Some(LinearCoreConfig {
kind: "gated_delta_net".into(),
num_heads: lnv.unwrap_or(0),
nphase: None,
value_head_dim: lvd.unwrap_or(0),
})
} else {
None
};
let rope_root = tc.get("rope_parameters");
let is_laguna_config = model_type.eq_ignore_ascii_case("laguna");
let rope = if is_laguna_config { rope_root.and_then(|r| r.get("full_attention")) } else { rope_root };
let local_rope = if is_laguna_config { rope_root.and_then(|r| r.get("sliding_attention")) } else { None };
let rope_theta = tc.get("rope_theta").and_then(|v| v.as_f64()).or_else(|| rope.and_then(|r| r.get("rope_theta")).and_then(|v| v.as_f64())).unwrap_or(10_000.0);
let prf = tc.get("partial_rotary_factor").and_then(|v| v.as_f64()).or_else(|| rope.and_then(|r| r.get("partial_rotary_factor")).and_then(|v| v.as_f64())).unwrap_or(1.0) as f32;
let local_prf = local_rope.and_then(|r| r.get("partial_rotary_factor")).and_then(|v| v.as_f64()).map(|v| v as f32);
let attention_heads_per_layer = tc
.get("num_attention_heads_per_layer")
.and_then(|v| v.as_array())
.map(|a| a.iter().map(|v| v.as_u64().map(|n| n as usize).ok_or_else(|| anyhow::anyhow!("num_attention_heads_per_layer must contain integers"))).collect::<anyhow::Result<Vec<_>>>())
.transpose()?;
if let Some(heads) = &attention_heads_per_layer {
anyhow::ensure!(heads.len() == n_layers, "num_attention_heads_per_layer has {} entries, expected {n_layers}", heads.len());
let nkv = cfg_usize(tc, "num_key_value_heads").unwrap_or(n_heads);
anyhow::ensure!(heads.iter().all(|&nh| nh > 0 && nh % nkv == 0), "every per-layer attention head count must be positive and divisible by num_key_value_heads={nkv}");
}
let moe = tc.get("num_experts").and_then(|v| v.as_u64()).filter(|&n| n > 0).map(|ne| {
let mt = model_type.to_lowercase();
let ntp_default = mt.starts_with("qwen3_5") || mt.contains("qwen3_next");
let is_lfm2 = mt.starts_with("lfm2");
let is_laguna = mt == "laguna";
MoeConfig {
num_experts: ne as usize,
top_k: cfg_usize(tc, "num_experts_per_tok").unwrap_or(2),
moe_intermediate_size: cfg_usize(tc, "moe_intermediate_size").unwrap_or(0),
norm_topk_prob: tc.get("norm_topk_prob").and_then(|v| v.as_bool()).unwrap_or(ntp_default),
shared_expert_intermediate_size: cfg_usize(tc, "shared_expert_intermediate_size"),
router_sigmoid: is_lfm2 || is_laguna,
routed_scaling_factor: tc.get("routed_scaling_factor").or_else(|| tc.get("moe_routed_scaling_factor")).and_then(|v| v.as_f64()).map(|v| v as f32).filter(|&v| (v - 1.0).abs() > 1e-9),
}
});
let head_dim = cfg_usize(tc, "head_dim").unwrap_or(hidden / n_heads.max(1));
let mt = model_type.to_lowercase();
let is_laguna = mt == "laguna";
if is_laguna {
anyhow::ensure!(!tc.get("swa_attention_sink_enabled").and_then(|v| v.as_bool()).unwrap_or(false), "laguna: learned SWA attention sinks are not supported yet");
anyhow::ensure!(tc.get("moe_router_logit_softcapping").and_then(|v| v.as_f64()).unwrap_or(0.0) == 0.0, "laguna: non-zero MoE router logit soft-capping is not supported");
anyhow::ensure!(!tc.get("moe_apply_router_weight_on_input").and_then(|v| v.as_bool()).unwrap_or(false), "laguna: moe_apply_router_weight_on_input=true is not supported");
}
let norm_style = if (mt.contains("gemma") && !mt.contains("gemma4")) || mt.starts_with("qwen3_5") || mt.contains("qwen3_next") {
NormStyle::Gemma
} else {
NormStyle::Qwen
};
let is_gemma = mt.contains("gemma");
let is_gemma4 = mt.contains("gemma4");
if tc.get("attn_logit_softcapping").and_then(|v| v.as_f64()).is_some() || (!is_gemma4 && tc.get("final_logit_softcapping").and_then(|v| v.as_f64()).is_some()) {
anyhow::bail!(
"{model_type}: attention logit soft-capping (Gemma-2) is not supported yet — \
Gemma-1/Gemma-3/Gemma-4 convert natively"
);
}
if is_gemma4 {
if tc.get("enable_moe_block").and_then(|v| v.as_bool()).unwrap_or(false) {
anyhow::bail!("{model_type}: gemma-4 MoE block (26B-A4B) is not supported yet");
}
if cfg_usize(tc, "hidden_size_per_layer_input").unwrap_or(0) > 0 {
anyhow::bail!(
"{model_type}: gemma-4 E-series per-layer inputs are not supported yet — \
the dense 12B/31B variants convert natively"
);
}
if cfg_usize(tc, "num_kv_shared_layers").unwrap_or(0) > 0 {
anyhow::bail!("{model_type}: gemma-4 KV-shared layers are not supported yet");
}
}
let (g4_rope_theta, g4_local_theta, g4_global_prf) = match rope {
Some(r) if is_gemma4 => {
let full = r.get("full_attention");
let slide = r.get("sliding_attention");
(
full.and_then(|f| f.get("rope_theta")).and_then(|v| v.as_f64()),
slide.and_then(|f| f.get("rope_theta")).and_then(|v| v.as_f64()),
full.and_then(|f| f.get("partial_rotary_factor")).and_then(|v| v.as_f64()).map(|v| v as f32),
)
}
_ => (None, None, None),
};
let rope_theta = g4_rope_theta.unwrap_or(rope_theta);
let yarn = rope
.filter(|r| r.get("rope_type").and_then(|v| v.as_str()) == Some("yarn"))
.map(|r| {
Ok::<YarnConfig, anyhow::Error>(YarnConfig {
factor: r.get("factor").and_then(|v| v.as_f64()).ok_or_else(|| anyhow::anyhow!("YaRN rope profile is missing factor"))? as f32,
original_max_position_embeddings: r.get("original_max_position_embeddings").and_then(|v| v.as_u64()).ok_or_else(|| anyhow::anyhow!("YaRN rope profile is missing original_max_position_embeddings"))? as usize,
beta_fast: r.get("beta_fast").and_then(|v| v.as_f64()).unwrap_or(32.0) as f32,
beta_slow: r.get("beta_slow").and_then(|v| v.as_f64()).unwrap_or(1.0) as f32,
attention_factor: r.get("attention_factor").and_then(|v| v.as_f64()).unwrap_or_else(|| {
let factor = r.get("factor").and_then(|v| v.as_f64()).unwrap_or(1.0);
0.1 * factor.ln() + 1.0
}) as f32,
})
})
.transpose()?;
let g4_pattern: Option<usize> = if is_gemma4 {
let fulls: Vec<usize> = tc.get("layer_types").and_then(|v| v.as_array()).map(|a| a.iter().enumerate().filter(|(_, v)| v.as_str() == Some("full_attention")).map(|(i, _)| i).collect()).unwrap_or_default();
let p = fulls.first().map(|f| f + 1).unwrap_or(0);
if p == 0 || fulls.iter().enumerate().any(|(k, &i)| i != p * (k + 1) - 1) || (n_layers / p) != fulls.len() {
anyhow::bail!("{model_type}: irregular full/sliding layer schedule not supported");
}
Some(p)
} else {
None
};
let hidden_act = match tc.get("hidden_activation").or_else(|| tc.get("hidden_act")).and_then(|v| v.as_str()).unwrap_or("silu") {
"gelu_pytorch_tanh" | "gelu_tanh" | "gelu_new" => "gelu_tanh".to_string(),
"silu" | "swish" => "silu".to_string(),
other => anyhow::bail!("unsupported hidden_act '{other}'"),
};
let embed_multiplier = if is_gemma { (hidden as f32).sqrt() } else { 1.0 };
let mut max_pos = cfg_usize(tc, "max_position_embeddings").unwrap_or(32768);
if let Some(rs) = tc.get("rope_scaling").filter(|v| !v.is_null()) {
let kind = rs.get("type").or_else(|| rs.get("rope_type")).and_then(|v| v.as_str());
match kind {
Some("longrope") | Some("su") | Some("yarn") | Some("linear") | Some("dynamic") | Some("mrope") => {
let orig = cfg_usize(tc, "original_max_position_embeddings").unwrap_or(4096);
eprintln!(" note: rope scaling '{:?}' — serving the exact {orig}-token native window", kind.unwrap());
max_pos = orig;
}
Some(other) => anyhow::bail!("rope_scaling '{other}' not supported yet"),
None => {}
}
}
Ok(ModelArch {
arch_name: model_type,
hidden_size: hidden,
intermediate_size: cfg_usize(tc, "intermediate_size").or_else(|| cfg_usize(tc, "moe_intermediate_size")).ok_or_else(|| anyhow::anyhow!("config: missing intermediate_size"))?,
num_layers: n_layers,
num_attention_heads: n_heads,
num_kv_heads: cfg_usize(tc, "num_key_value_heads").unwrap_or(n_heads),
head_dim,
vocab_size: cfg_usize(tc, "vocab_size").ok_or_else(|| anyhow::anyhow!("config: missing vocab_size"))?,
layer_types,
rms_norm_eps: tc.get("rms_norm_eps").or_else(|| tc.get("norm_eps")).and_then(|v| v.as_f64()).unwrap_or(1e-6),
norm_style,
rope_theta,
tie_word_embeddings: config.get("tie_word_embeddings").and_then(|v| v.as_bool()).unwrap_or(is_gemma),
partial_rotary_factor: prf,
yarn,
attention_heads_per_layer,
mtp: None,
moe,
linear_core,
max_position_embeddings: max_pos,
linear_conv_kernel_dim: cfg_usize(tc, "linear_conv_kernel_dim").or_else(|| cfg_usize(tc, "conv_L_cache")),
linear_num_key_heads: cfg_usize(tc, "linear_num_key_heads"),
linear_num_value_heads: lnv,
linear_key_head_dim: cfg_usize(tc, "linear_key_head_dim"),
linear_value_head_dim: lvd,
hidden_act,
embed_multiplier,
query_pre_attn_scalar: tc.get("query_pre_attn_scalar").and_then(|v| v.as_f64()).or(if is_gemma4 { Some(1.0) } else { None }),
sliding_window: cfg_usize(tc, "sliding_window").filter(|_| is_laguna || tc.get("sliding_window_pattern").is_some() || g4_pattern.is_some()),
sliding_window_pattern: cfg_usize(tc, "sliding_window_pattern").or(g4_pattern),
rope_local_base_freq: tc.get("rope_local_base_freq").and_then(|v| v.as_f64()).or(g4_local_theta).or_else(|| local_rope.and_then(|r| r.get("rope_theta")).and_then(|v| v.as_f64())),
local_partial_rotary_factor: local_prf,
global_head_dim: cfg_usize(tc, "global_head_dim").filter(|_| is_gemma4),
num_global_kv_heads: cfg_usize(tc, "num_global_key_value_heads").filter(|_| is_gemma4),
global_partial_rotary_factor: g4_global_prf,
final_logit_softcapping: if is_gemma4 { tc.get("final_logit_softcapping").and_then(|v| v.as_f64()) } else { None },
attn_v_norm: is_gemma4,
num_loops: cfg_usize(tc, "num_loops").unwrap_or(1),
loop_final_norm: !tc.get("skip_loop_final_norm").and_then(|v| v.as_bool()).unwrap_or(true),
})
}
fn eos_ids(gen_cfg: &serde_json::Value, config: &serde_json::Value) -> Vec<u32> {
for v in [gen_cfg.get("eos_token_id"), config.get("eos_token_id")].into_iter().flatten() {
if let Some(n) = v.as_u64() {
return vec![n as u32];
}
if let Some(a) = v.as_array() {
return a.iter().filter_map(|x| x.as_u64().map(|n| n as u32)).collect();
}
}
Vec::new()
}
pub(crate) fn looks_like_repo(s: &str) -> bool {
let s = s.trim_matches('/');
s.split('/').count() == 2 && !s.contains(char::is_whitespace) && !Path::new(s).exists()
}
fn hf_agent() -> ureq::Agent {
ureq::AgentBuilder::new().timeout_connect(Duration::from_secs(20)).timeout_read(Duration::from_secs(300)).build()
}
pub(crate) fn hf_repo_files(repo: &str, token: Option<&str>) -> Vec<String> {
repo_files(&hf_agent(), repo, token)
}
pub(crate) fn hf_fetch_file(repo: &str, filename: &str, token: Option<&str>) -> anyhow::Result<std::path::PathBuf> {
let dir = hf_cache_dir(repo)?;
let dest = dir.join(filename.replace('/', "__"));
let url = format!("https://huggingface.co/{repo}/resolve/main/{filename}");
fetch(&hf_agent(), &url, &dest, token, true, hf_threads())?;
Ok(dest)
}
fn hf_cache_dir(repo: &str) -> anyhow::Result<std::path::PathBuf> {
let base = std::env::var_os("HOME").map(|h| std::path::PathBuf::from(h).join(".cache/cortiq/hf")).unwrap_or_else(|| std::path::PathBuf::from(".cortiq-hf"));
let dir = base.join(repo.replace('/', "--"));
fs::create_dir_all(&dir)?;
Ok(dir)
}
const HF_CHUNK: u64 = 32 * 1024 * 1024;
fn hf_threads() -> usize {
std::env::var("CORTIQ_HF_THREADS").ok().and_then(|v| v.parse::<usize>().ok()).filter(|&n| n >= 1).unwrap_or(8).min(16)
}
fn cached(dest: &Path) -> bool {
dest.exists() && fs::metadata(dest).map(|m| m.len() > 0).unwrap_or(false)
}
fn auth(mut req: ureq::Request, token: Option<&str>) -> ureq::Request {
req = req.set("User-Agent", "cortiq-convert");
if let Some(t) = token {
req = req.set("Authorization", &format!("Bearer {t}"));
}
req
}
fn probe_size(agent: &ureq::Agent, url: &str, token: Option<&str>) -> Option<u64> {
let resp = auth(agent.get(url).set("Range", "bytes=0-0"), token).call().ok()?;
resp.header("Content-Range")?.rsplit('/').next()?.trim().parse::<u64>().ok()
}
fn get_range(agent: &ureq::Agent, url: &str, token: Option<&str>, start: u64, end: u64) -> anyhow::Result<Vec<u8>> {
let resp = auth(agent.get(url).set("Range", &format!("bytes={}-{}", start, end - 1)), token).call().map_err(|e| anyhow::anyhow!("{e}"))?;
let mut buf = Vec::with_capacity((end - start) as usize);
resp.into_reader().read_to_end(&mut buf)?;
Ok(buf)
}
fn write_at(path: &Path, offset: u64, data: &[u8]) -> std::io::Result<()> {
use std::io::{Seek, SeekFrom, Write};
let mut f = fs::OpenOptions::new().write(true).open(path)?;
f.seek(SeekFrom::Start(offset))?;
f.write_all(data)
}
fn with_retry<T>(attempts: u32, mut f: impl FnMut() -> anyhow::Result<T>) -> anyhow::Result<T> {
let mut delay = Duration::from_millis(400);
let mut last: Option<anyhow::Error> = None;
for a in 0..attempts {
match f() {
Ok(v) => return Ok(v),
Err(e) => {
last = Some(e);
if a + 1 < attempts {
std::thread::sleep(delay);
delay = (delay * 2).min(Duration::from_secs(8));
}
}
}
}
Err(last.unwrap())
}
fn fetch(agent: &ureq::Agent, url: &str, dest: &Path, token: Option<&str>, required: bool, threads: usize) -> anyhow::Result<bool> {
if cached(dest) {
return Ok(true);
}
let tmp = dest.with_extension("part");
let size = probe_size(agent, url, token);
if let Some(sz) = size {
if sz > HF_CHUNK && threads > 1 {
{
let f = fs::File::create(&tmp)?;
f.set_len(sz)?;
}
let chunks: Vec<(u64, u64)> = (0..sz).step_by(HF_CHUNK as usize).map(|s| (s, (s + HF_CHUNK).min(sz))).collect();
let total = chunks.len();
let queue = Mutex::new(chunks);
let err: Mutex<Option<String>> = Mutex::new(None);
let done = std::sync::atomic::AtomicUsize::new(0);
std::thread::scope(|scope| {
for _ in 0..threads {
scope.spawn(|| {
loop {
if err.lock().unwrap().is_some() {
break;
}
let Some((start, end)) = queue.lock().unwrap().pop() else {
break;
};
let r = with_retry(4, || get_range(agent, url, token, start, end)).and_then(|buf| write_at(&tmp, start, &buf).map_err(Into::into));
match r {
Ok(()) => {
let d = done.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1;
eprint!("\r downloading: {:>3}% ({d}/{total} chunks)", d * 100 / total);
}
Err(e) => {
*err.lock().unwrap() = Some(e.to_string());
break;
}
}
}
});
}
});
eprintln!();
if let Some(e) = err.into_inner().unwrap() {
anyhow::bail!("download {url}: {e}");
}
fs::rename(&tmp, dest)?;
return Ok(true);
}
}
let got = with_retry(4, || match auth(agent.get(url), token).call() {
Ok(resp) => {
let mut r = resp.into_reader();
let mut f = fs::File::create(&tmp)?;
std::io::copy(&mut r, &mut f)?;
Ok(Some(()))
}
Err(ureq::Error::Status(404, _)) if !required => Ok(None),
Err(e) => Err(anyhow::anyhow!("download {url}: {e}")),
})?;
match got {
Some(()) => {
fs::rename(&tmp, dest)?;
Ok(true)
}
None => Ok(false),
}
}
fn repo_files(agent: &ureq::Agent, repo: &str, token: Option<&str>) -> Vec<String> {
let url = format!("https://huggingface.co/api/models/{repo}");
match auth(agent.get(&url), token).call() {
Ok(resp) => resp
.into_json::<serde_json::Value>()
.ok()
.and_then(|j| j["siblings"].as_array().map(|a| a.iter().filter_map(|s| s["rfilename"].as_str().map(String::from)).collect()))
.unwrap_or_default(),
Err(_) => Vec::new(),
}
}
pub(crate) fn hf_download(repo: &str, token: Option<&str>) -> anyhow::Result<std::path::PathBuf> {
let dir = hf_cache_dir(repo)?;
let base = format!("https://huggingface.co/{repo}/resolve/main");
let threads = hf_threads();
let agent = ureq::AgentBuilder::new().timeout_connect(Duration::from_secs(20)).timeout_read(Duration::from_secs(300)).build();
if !fetch(&agent, &format!("{base}/config.json"), &dir.join("config.json"), token, false, threads)? {
let files = repo_files(&agent, repo, token);
let ggufs = files.iter().filter(|f| f.to_lowercase().ends_with(".gguf")).count();
if ggufs > 0 {
let src = repo.strip_suffix("-GGUF").or_else(|| repo.strip_suffix("-gguf")).filter(|s| !s.is_empty());
anyhow::bail!(
"'{repo}' is a GGUF repository ({ggufs} .gguf file(s), no config.json); \
`cortiq convert` needs a safetensors checkpoint. Either import a GGUF file \
directly with `cortiq import-gguf <file.gguf>` (dense llama/qwen2/qwen3, F32/F16/Q8_0), \
or convert the source safetensors repo instead{}.",
match src {
Some(s) => format!(" — try `--model {s}`"),
None => String::new(),
}
);
}
anyhow::bail!("'{repo}': no config.json — not a Hugging Face safetensors checkpoint");
}
for (f, required) in [
("tokenizer.json", true),
("tokenizer_config.json", false),
("generation_config.json", false),
("chat_template.jinja", false),
] {
fetch(&agent, &format!("{base}/{f}"), &dir.join(f), token, required, threads)?;
}
let idx = dir.join("model.safetensors.index.json");
if fetch(&agent, &format!("{base}/model.safetensors.index.json"), &idx, token, false, 1)? {
let j: serde_json::Value = serde_json::from_slice(&fs::read(&idx)?)?;
let map = j["weight_map"].as_object().ok_or_else(|| anyhow::anyhow!("bad safetensors index"))?;
let mut shards: Vec<String> = map.values().filter_map(|v| v.as_str().map(String::from)).collect();
shards.sort();
shards.dedup();
for (i, s) in shards.iter().enumerate() {
eprintln!(" shard {}/{} ({threads}× parallel): {s}", i + 1, shards.len());
fetch(&agent, &format!("{base}/{s}"), &dir.join(s), token, true, threads)?;
}
} else {
eprintln!(" model.safetensors ({threads}× parallel)");
fetch(&agent, &format!("{base}/model.safetensors"), &dir.join("model.safetensors"), token, true, threads)?;
}
Ok(dir)
}
fn split_fused_gdn(name: &str, w: &[f32], hid: usize, nk: usize, dk: usize, nv: usize, dv: usize) -> anyhow::Result<Vec<(String, Vec<f32>, usize)>> {
if nk == 0 || nv % nk != 0 {
anyhow::bail!("fused GDN: bad head config nk={nk} nv={nv}");
}
let r = nv / nk;
let row = |w: &[f32], gw: usize, g: usize, gr: usize| -> Vec<f32> {
let base = (g * gw + gr) * hid;
w[base..base + hid].to_vec()
};
if name.contains("in_proj_qkvz") {
let gw = 2 * dk + 2 * r * dv;
if w.len() != nk * gw * hid {
anyhow::bail!("fused GDN qkvz: {} values, expected {}", w.len(), nk * gw * hid);
}
let mut qkv = Vec::with_capacity((2 * nk * dk + nv * dv) * hid);
for g in 0..nk {
for rr in 0..dk {
qkv.extend_from_slice(&row(w, gw, g, rr));
}
}
for g in 0..nk {
for rr in 0..dk {
qkv.extend_from_slice(&row(w, gw, g, dk + rr));
}
}
for g in 0..nk {
for rr in 0..r * dv {
qkv.extend_from_slice(&row(w, gw, g, 2 * dk + rr));
}
}
let mut z = Vec::with_capacity(nv * dv * hid);
for g in 0..nk {
for rr in 0..r * dv {
z.extend_from_slice(&row(w, gw, g, 2 * dk + r * dv + rr));
}
}
let p = name.strip_suffix("in_proj_qkvz.weight").unwrap_or(name);
Ok(vec![(format!("{p}in_proj_qkv.weight"), qkv, 2 * nk * dk + nv * dv), (format!("{p}in_proj_z.weight"), z, nv * dv)])
} else {
let gw = 2 * r;
if w.len() != nk * gw * hid {
anyhow::bail!("fused GDN ba: {} values, expected {}", w.len(), nk * gw * hid);
}
let mut b = Vec::with_capacity(nv * hid);
let mut a = Vec::with_capacity(nv * hid);
for g in 0..nk {
for rr in 0..r {
b.extend_from_slice(&row(w, gw, g, rr));
}
}
for g in 0..nk {
for rr in 0..r {
a.extend_from_slice(&row(w, gw, g, r + rr));
}
}
let p = name.strip_suffix("in_proj_ba.weight").unwrap_or(name);
Ok(vec![(format!("{p}in_proj_b.weight"), b, nv), (format!("{p}in_proj_a.weight"), a, nv)])
}
}
enum FfnKind {
Gate,
Up,
Down,
}
fn ffn_kind(name: &str) -> Option<(usize, FfnKind)> {
let rest = name.strip_prefix("model.layers.")?;
let dot = rest.find('.')?;
let li: usize = rest[..dot].parse().ok()?;
let kind = match &rest[dot + 1..] {
"mlp.gate_proj.weight" => FfnKind::Gate,
"mlp.up_proj.weight" => FfnKind::Up,
"mlp.down_proj.weight" => FfnKind::Down,
_ => return None,
};
Some((li, kind))
}
fn slice_ffn(kind: &FfnKind, shape: &[usize], vals: &[f32], keep: &[bool]) -> anyhow::Result<(Vec<usize>, Vec<f32>)> {
let k = keep.iter().filter(|&&b| b).count();
match kind {
FfnKind::Gate | FfnKind::Up => {
let (inter, hidden) = (shape[0], shape[1]);
if keep.len() != inter {
anyhow::bail!("defrag: keep len {} != gate/up rows {inter}", keep.len());
}
let mut out = Vec::with_capacity(k * hidden);
for r in 0..inter {
if keep[r] {
out.extend_from_slice(&vals[r * hidden..(r + 1) * hidden]);
}
}
Ok((vec![k, hidden], out))
}
FfnKind::Down => {
let (hidden, inter) = (shape[0], shape[1]);
if keep.len() != inter {
anyhow::bail!("defrag: keep len {} != down cols {inter}", keep.len());
}
let mut out = Vec::with_capacity(hidden * k);
for r in 0..hidden {
for c in 0..inter {
if keep[c] {
out.push(vals[r * inter + c]);
}
}
}
Ok((vec![hidden, k], out))
}
}
}
fn effective_tensor(overlay: &HashMap<String, (Vec<usize>, Vec<f32>)>, files: &[SafeTensors], name: &str) -> anyhow::Result<(Vec<usize>, Vec<f32>)> {
if let Some((s, v)) = overlay.get(name) {
return Ok((s.clone(), v.clone()));
}
for f in files {
for m in &f.tensors {
if canon_name(&m.name).as_deref() == Some(name) {
return Ok((m.shape.clone(), to_f32(&m.dtype, f.bytes(m))?));
}
}
}
anyhow::bail!("defrag: tensor '{name}' not in overlay or base model")
}
struct DefragPlan {
overlay: HashMap<String, (Vec<usize>, Vec<f32>)>,
keep: HashMap<usize, Vec<bool>>,
}
fn build_defrag_plan(dir: &Path, arch: &ModelArch, files: &[SafeTensors]) -> anyhow::Result<DefragPlan> {
let mut overlay: HashMap<String, (Vec<usize>, Vec<f32>)> = HashMap::new();
let tdir = dir.join("tensors");
if tdir.is_dir() {
for entry in fs::read_dir(&tdir)? {
let p = entry?.path();
if p.extension().and_then(|e| e.to_str()) != Some("npy") {
continue;
}
let stem = p.file_stem().and_then(|s| s.to_str()).unwrap_or_default().to_string();
let a = npy::read(&p)?;
let vals = match a.data {
npy::NpyData::F32(v) => v,
npy::NpyData::Bool(_) => {
anyhow::bail!("defrag overlay {stem}: expected float, got bool")
}
};
overlay.insert(stem, (a.shape, vals));
}
}
println!(" Defrag overlay: {} baked tensors from {}", overlay.len(), dir.display());
let (nl, inter) = (arch.num_layers, arch.intermediate_size);
let mut keep: HashMap<usize, Vec<bool>> = HashMap::new();
let keep_path = dir.join("ffn_keep.npy");
if keep_path.exists() {
let a = npy::read(&keep_path)?;
if a.shape != [nl, inter] {
anyhow::bail!("ffn_keep.npy shape {:?} != model ({nl}, {inter})", a.shape);
}
let flags: Vec<bool> = match a.data {
npy::NpyData::Bool(v) => v,
npy::NpyData::F32(v) => v.iter().map(|&x| x != 0.0).collect(),
};
for li in 0..nl {
let row = flags[li * inter..(li + 1) * inter].to_vec();
if !row.iter().any(|&b| b) {
anyhow::bail!("defrag: layer {li} has 0 live neurons");
}
keep.insert(li, row);
}
} else {
println!(" Defrag: no ffn_keep.npy — autodetecting from zero down_proj columns");
for li in 0..nl {
let name = format!("model.layers.{li}.mlp.down_proj.weight");
let (shape, vals) = effective_tensor(&overlay, files, &name)?;
let (hidden, cols) = (shape[0], shape[1]);
let mut alive = vec![false; cols];
for r in 0..hidden {
for c in 0..cols {
if vals[r * cols + c] != 0.0 {
alive[c] = true;
}
}
}
if !alive.iter().any(|&b| b) {
anyhow::bail!("defrag: layer {li} autodetected 0 live neurons");
}
keep.insert(li, alive);
}
}
Ok(DefragPlan { overlay, keep })
}
pub fn run_convert(
model: &str,
quant: &str,
output: &str,
hf_token: Option<&str>,
defrag: Option<&str>,
o1_hint: Option<serde_json::Value>,
mut progress: impl FnMut(f32),
) -> anyhow::Result<()> {
let quant = parse_quant(quant)?;
let downloaded;
let dir: &Path = if Path::new(model).join("config.json").exists() {
Path::new(model)
} else if looks_like_repo(model) {
eprintln!("downloading {model} from Hugging Face…");
downloaded = hf_download(model, hf_token)?;
downloaded.as_path()
} else {
anyhow::bail!("'{model}': not a local model dir (no config.json) and not an HF repo id (owner/name)");
};
let config: serde_json::Value = serde_json::from_slice(&fs::read(dir.join("config.json")).map_err(|e| anyhow::anyhow!("read config.json: {e}"))?)?;
let mut arch = build_arch(&config)?;
let files = open_model(dir)?;
let orig_inter = arch.intermediate_size;
let defrag_plan = match defrag {
Some(d) => Some(build_defrag_plan(Path::new(d), &arch, &files)?),
None => None,
};
if let Some(plan) = &defrag_plan {
let max_kept = (0..arch.num_layers).filter_map(|li| plan.keep.get(&li).map(|k| k.iter().filter(|&&b| b).count())).max().unwrap_or(orig_inter);
arch.intermediate_size = max_kept;
}
let total: usize = files.iter().map(|f| f.tensors.len()).sum::<usize>().max(1);
let mut tensors: Vec<TensorSpec> = Vec::with_capacity(total);
let mut done = 0usize;
for file in &files {
for m in &file.tensors {
done += 1;
progress(done as f32 / total as f32);
let Some(name) = canon_name(&m.name) else {
continue;
};
if m.dtype == "F16" && (name.ends_with(".scales") || name.ends_with(".biases")) {
continue;
}
let (m_shape, m_vals) = if m.dtype == "U32" && m.name.ends_with(".weight") {
let scales_name = m.name.replace(".weight", ".scales");
let biases_name = m.name.replace(".weight", ".biases");
let mut scales_blob = None;
let mut biases_blob = None;
for f in &files {
if let Some(t) = f.tensors.iter().find(|t| t.name == scales_name) {
scales_blob = Some(f.bytes(t));
}
if let Some(t) = f.tensors.iter().find(|t| t.name == biases_name) {
biases_blob = Some(f.bytes(t));
}
}
let scales = scales_blob.ok_or_else(|| anyhow::anyhow!("missing {} for MLX unpacking", scales_name))?;
let out_dim = m.shape[0];
let w_cols = m.shape[1];
let num_groups = scales.len() / 2 / out_dim;
let mut bits = 0;
let mut in_dim = 0;
for b in [1, 2, 3, 4, 8] {
let possible_in_dim = w_cols * 32 / b;
if possible_in_dim % num_groups == 0 {
let gs = possible_in_dim / num_groups;
if gs == 32 || gs == 64 || gs == 128 {
bits = b;
in_dim = possible_in_dim;
break;
}
}
}
if bits == 0 {
anyhow::bail!("Could not deduce MLX bit width for shape {:?} and {} scale groups", m.shape, num_groups);
}
(vec![out_dim, in_dim], unpack_mlx(file.bytes(m), scales, biases_blob, out_dim, in_dim, bits)?)
} else {
(m.shape.clone(), to_f32(&m.dtype, file.bytes(m))?)
};
if name.contains(".linear_attn.in_proj_qkvz") || name.contains(".linear_attn.in_proj_ba") {
if m_shape.len() != 2 {
anyhow::bail!("fused GDN tensor '{name}': expected 2-D, got {:?}", m_shape);
}
let w = &m_vals;
let hid = m_shape[1];
let miss = |k: &str| anyhow::anyhow!("fused GDN needs {k} in config");
let nk = arch.linear_num_key_heads.ok_or_else(|| miss("linear_num_key_heads"))?;
let dk = arch.linear_key_head_dim.ok_or_else(|| miss("linear_key_head_dim"))?;
let nv = arch.linear_num_value_heads.ok_or_else(|| miss("linear_num_value_heads"))?;
let dv = arch.linear_value_head_dim.ok_or_else(|| miss("linear_value_head_dim"))?;
for (out_name, out_vals, out_rows) in split_fused_gdn(&name, w, hid, nk, dk, nv, dv)? {
let two_d = out_rows * hid >= GROUP_SIZE && !force_f16(&out_name);
let (dt, data) = if two_d { quantize_2d(quant, &out_vals, out_rows, hid) } else { (TensorDtype::F16, encode_f16(&out_vals)) };
tensors.push(TensorSpec {
name: out_name,
dtype: dt,
shape: vec![out_rows, hid],
data,
});
}
continue;
}
if name.ends_with(".self_attn.qkv_proj.weight") || name.ends_with(".mlp.gate_up_proj.weight") {
anyhow::ensure!(m_shape.len() == 2, "fused '{name}': expected 2-D");
let w = &m_vals;
let (rows, cols) = (m_shape[0], m_shape[1]);
let parts: Vec<(String, usize, usize)> = if name.contains("qkv_proj") {
let q = arch.num_attention_heads * arch.head_dim;
let kv = arch.num_kv_heads * arch.head_dim;
anyhow::ensure!(q + 2 * kv == rows, "'{name}': {rows} rows != q({q}) + 2·kv({kv})");
vec![(name.replace("qkv_proj", "q_proj"), 0, q), (name.replace("qkv_proj", "k_proj"), q, kv), (name.replace("qkv_proj", "v_proj"), q + kv, kv)]
} else {
anyhow::ensure!(rows % 2 == 0, "'{name}': odd row count {rows}");
vec![(name.replace("gate_up_proj", "gate_proj"), 0, rows / 2), (name.replace("gate_up_proj", "up_proj"), rows / 2, rows / 2)]
};
for (out_name, r0, nr) in parts {
let vals = &w[r0 * cols..(r0 + nr) * cols];
let (dt, data) = if nr * cols >= GROUP_SIZE && !force_f16(&out_name) { quantize_2d(quant, vals, nr, cols) } else { (TensorDtype::F16, encode_f16(vals)) };
tensors.push(TensorSpec { name: out_name, dtype: dt, shape: vec![nr, cols], data });
}
continue;
}
if arch.global_head_dim.is_some() && name.ends_with(".self_attn.k_proj.weight") {
let li: Option<usize> = name.split(".layers.").nth(1).and_then(|r| r.split('.').next()).and_then(|n| n.parse().ok());
let pat = arch.sliding_window_pattern.unwrap_or(usize::MAX);
if let Some(li) = li {
if (li + 1) % pat == 0 {
anyhow::ensure!(m_shape.len() == 2, "'{name}': expected 2-D");
let w = &m_vals;
let (rows, cols) = (m_shape[0], m_shape[1]);
for out_name in [name.clone(), name.replace("k_proj", "v_proj")] {
let (dt, data) = if rows * cols >= GROUP_SIZE && !force_f16(&out_name) { quantize_2d(quant, w, rows, cols) } else { (TensorDtype::F16, encode_f16(w)) };
tensors.push(TensorSpec {
name: out_name,
dtype: dt,
shape: vec![rows, cols],
data,
});
}
continue;
}
}
}
if let Some(plan) = defrag_plan.as_ref() {
if let Some((li, kind)) = ffn_kind(&name) {
if let Some(keep) = plan.keep.get(&li) {
let (shape, vals) = match plan.overlay.get(&name) {
Some((s, v)) => (s.clone(), v.clone()),
None => (m_shape.clone(), m_vals.clone()),
};
let (out_shape, out_vals) = slice_ffn(&kind, &shape, &vals, keep)?;
let numel = out_shape[0] * out_shape[1];
let two_d = numel >= GROUP_SIZE && !force_f16(&name);
let (dt, data) = if two_d { quantize_2d(quant, &out_vals, out_shape[0], out_shape[1]) } else { (TensorDtype::F16, encode_f16(&out_vals)) };
tensors.push(TensorSpec { name, dtype: dt, shape: out_shape, data });
continue;
}
}
}
let vals = m_vals;
let numel: usize = m_shape.iter().product();
if numel != vals.len() {
anyhow::bail!("tensor '{name}': {} values for shape {:?}", vals.len(), m_shape);
}
let two_d = m_shape.len() == 2 && numel >= GROUP_SIZE && !force_f16(&name);
let (dt, data) = if two_d { quantize_2d(quant, &vals, m_shape[0], m_shape[1]) } else { (TensorDtype::F16, encode_f16(&vals)) };
tensors.push(TensorSpec { name, dtype: dt, shape: m_shape.clone(), data });
}
}
let vocab = fs::read(dir.join("tokenizer.json")).ok();
let tok_cfg: serde_json::Value = fs::read(dir.join("tokenizer_config.json")).ok().and_then(|b| serde_json::from_slice(&b).ok()).unwrap_or(serde_json::Value::Null);
let gen_cfg: serde_json::Value = fs::read(dir.join("generation_config.json")).ok().and_then(|b| serde_json::from_slice(&b).ok()).unwrap_or(serde_json::Value::Null);
let chat_template = fs::read_to_string(dir.join("chat_template.jinja")).ok().filter(|s| !s.trim().is_empty()).or_else(|| tok_cfg.get("chat_template").and_then(|v| v.as_str().map(String::from)));
let bundle = TokenizerBundle {
chat_template,
eos_token_ids: eos_ids(&gen_cfg, &config),
bos_token_id: config.get("bos_token_id").and_then(|v| v.as_u64()).map(|n| n as u32),
pad_token_id: config.get("pad_token_id").and_then(|v| v.as_u64()).map(|n| n as u32),
};
let quant_type = match quant {
Quant::Q8Row => QuantType::Q8Row,
Quant::Q8_2f => QuantType::Q8_2f,
Quant::Q4Block => QuantType::Q4Block,
Quant::F16 => QuantType::F16,
Quant::Vbit => QuantType::Vbit,
Quant::Q4Tiled => QuantType::Q4Block,
Quant::Q1 | Quant::Q1p | Quant::Q1s | Quant::Q1t => QuantType::Vbit,
};
let provenance = match &defrag_plan {
Some(plan) => {
let kept: Vec<usize> = (0..arch.num_layers).map(|li| plan.keep.get(&li).map(|k| k.iter().filter(|&&b| b).count()).unwrap_or(orig_inter)).collect();
let live: usize = kept.iter().sum();
let ratio = 1.0 - live as f64 / (arch.num_layers as f64 * orig_inter as f64);
eprintln!(
"defrag: FFN pruned per-layer, {live}/{} live ({:.0}% pruned), inter {orig_inter} -> max {} (per-layer variable); masks dropped",
arch.num_layers * orig_inter,
ratio * 100.0,
arch.intermediate_size
);
serde_json::json!({
"tool": "cortiq convert",
"source_model": model,
"defrag": {
"source_skill": defrag,
"pre_intermediate": orig_inter,
"post_intermediate_max": arch.intermediate_size,
"kept_per_layer": kept,
"pruned_ratio": (ratio * 10000.0).round() / 10000.0,
}
})
}
None => serde_json::json!({ "tool": "cortiq convert", "source_model": model }),
};
let provenance = match o1_hint {
Some(h) => {
let mut p = provenance;
p["o1_attn"] = h;
p
}
None => provenance,
};
let header = CmfHeader {
format: "cmf".into(),
version: CMF_VERSION,
arch,
quant_type,
provenance: Some(provenance),
tokenizer_config: Some(bundle),
section_hashes: None,
skills: Vec::new(),
shard: None,
calibration: None,
};
tensors.sort_by(|a, b| exec_order_key(&a.name).cmp(&exec_order_key(&b.name)).then_with(|| a.name.cmp(&b.name)));
CmfModel::write(output, &header, &tensors, None, vocab.as_deref()).map_err(|e| anyhow::anyhow!("write {output}: {e}"))?;
progress(1.0);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use cortiq_core::quant::{dequant_q4_block, dequant_q8_2f, dequant_q8_row, dequant_vbit};
#[test]
fn laguna_config_maps_to_exact_cmf_contract() {
let config = serde_json::json!({
"model_type": "laguna",
"hidden_size": 32,
"intermediate_size": 64,
"num_hidden_layers": 4,
"num_attention_heads": 4,
"num_attention_heads_per_layer": [4, 6, 6, 6],
"num_key_value_heads": 2,
"head_dim": 8,
"vocab_size": 100,
"max_position_embeddings": 1048576,
"num_experts": 8,
"num_experts_per_tok": 2,
"moe_intermediate_size": 16,
"shared_expert_intermediate_size": 16,
"norm_topk_prob": true,
"moe_routed_scaling_factor": 2.5,
"sliding_window": 512,
"layer_types": ["full_attention", "sliding_attention", "sliding_attention", "sliding_attention"],
"rope_parameters": {
"full_attention": {
"rope_type": "yarn", "rope_theta": 500000.0,
"factor": 128.0, "original_max_position_embeddings": 8192,
"beta_fast": 32.0, "beta_slow": 1.0,
"attention_factor": 1.485203, "partial_rotary_factor": 0.5
},
"sliding_attention": {
"rope_type": "default", "rope_theta": 10000.0,
"partial_rotary_factor": 1.0
}
}
});
let arch = build_arch(&config).unwrap();
assert_eq!(arch.arch_name, "laguna");
assert_eq!(arch.attention_heads_per_layer, Some(vec![4, 6, 6, 6]));
assert_eq!(arch.layer_types, vec![LayerType::FullAttention, LayerType::SlidingAttention, LayerType::SlidingAttention, LayerType::SlidingAttention,]);
assert_eq!(arch.sliding_window, Some(512));
assert_eq!(arch.rope_theta, 500000.0);
assert_eq!(arch.rope_local_base_freq, Some(10000.0));
assert_eq!(arch.partial_rotary_factor, 0.5);
assert_eq!(arch.local_partial_rotary_factor, Some(1.0));
let yarn = arch.yarn.unwrap();
assert_eq!(yarn.factor, 128.0);
let moe = arch.moe.unwrap();
assert!(moe.router_sigmoid);
assert_eq!(moe.routed_scaling_factor, Some(2.5));
assert_eq!(canon_name("model.layers.1.mlp.experts.e_score_correction_bias").as_deref(), Some("model.layers.1.mlp.expert_bias"));
}
#[test]
fn exec_order_lays_out_by_layer_then_block() {
let mut names: Vec<&str> = vec![
"lm_head.weight",
"model.layers.10.mlp.experts.1.up_proj.weight",
"model.embed_tokens.weight",
"model.layers.2.self_attn.q_proj.weight",
"model.norm.weight",
"model.layers.2.mlp.gate.weight", "model.layers.2.mlp.experts.0.down_proj.weight",
"model.layers.2.input_layernorm.weight",
"model.layers.2.mlp.experts.0.gate_proj.weight",
"model.layers.10.self_attn.o_proj.weight",
];
names.sort_by(|a, b| exec_order_key(a).cmp(&exec_order_key(b)).then_with(|| a.cmp(b)));
assert_eq!(
names,
vec![
"model.embed_tokens.weight",
"model.layers.2.input_layernorm.weight",
"model.layers.2.self_attn.q_proj.weight",
"model.layers.2.mlp.gate.weight",
"model.layers.2.mlp.experts.0.gate_proj.weight",
"model.layers.2.mlp.experts.0.down_proj.weight",
"model.layers.10.self_attn.o_proj.weight",
"model.layers.10.mlp.experts.1.up_proj.weight",
"model.norm.weight",
"lm_head.weight",
]
);
}
#[test]
fn vbit_roundtrip_within_quant_error() {
let (rows, cols) = (5usize, 64usize);
let mut vals = vec![0f32; rows * cols];
for o in 0..rows {
for i in 0..cols {
vals[o * cols + i] = (o as f32 + 1.0) * 0.13 * (i as f32 * 0.27).sin();
}
}
let enc = encode_vbit(&vals, rows, cols);
let bits = &enc[..rows];
assert!(bits.iter().all(|&b| (3..=8).contains(&b)));
let mut dec = vec![0f32; rows * cols];
dequant_vbit(&enc, rows, cols, &mut dec).unwrap();
for o in 0..rows {
let amp = vals[o * cols..(o + 1) * cols].iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
for i in 0..cols {
let e = (dec[o * cols + i] - vals[o * cols + i]).abs();
assert!(e <= amp * 0.2, "row {o} col {i}: err {e} vs amp {amp} (bits {})", bits[o]);
}
}
}
#[test]
fn vbit_ro_roundtrip_and_validation() {
use cortiq_core::TensorDtype;
use cortiq_core::quant::{dequant_vbit_ro, validate_payload};
let (rows, cols) = (5usize, 64usize);
let mut vals = vec![0f32; rows * cols];
for o in 0..rows {
for i in 0..cols {
vals[o * cols + i] = (o as f32 + 1.0) * 0.13 * (i as f32 * 0.27).sin();
}
}
let enc = encode_vbit_ro(&vals, rows, cols);
validate_payload(TensorDtype::VbitRo, &[rows, cols], &enc).unwrap();
let mut dec = vec![0f32; rows * cols];
dequant_vbit_ro(&enc, rows, cols, &mut dec).unwrap();
let legacy = encode_vbit(&vals, rows, cols);
let mut dec_legacy = vec![0f32; rows * cols];
dequant_vbit(&legacy, rows, cols, &mut dec_legacy).unwrap();
assert_eq!(dec, dec_legacy, "vbit_ro must reconstruct exactly like vbit");
}
#[test]
fn fused_gdn_split_is_correct_permutation() {
let (nk, dk, nv, dv, hid) = (2usize, 3usize, 4usize, 2usize, 1usize);
let r = nv / nk; let gw = 2 * dk + 2 * r * dv; let w: Vec<f32> = (0..nk * gw * hid).map(|i| i as f32).collect();
let out = split_fused_gdn("m.linear_attn.in_proj_qkvz.weight", &w, hid, nk, dk, nv, dv).unwrap();
let qkv = &out[0];
assert_eq!(qkv.0, "m.linear_attn.in_proj_qkv.weight");
assert_eq!(qkv.2, 2 * nk * dk + nv * dv); assert_eq!(qkv.1[0..3], [0.0, 1.0, 2.0]);
assert_eq!(qkv.1[3..6], [14.0, 15.0, 16.0]);
assert_eq!(qkv.1[6..9], [3.0, 4.0, 5.0]);
assert_eq!(qkv.1[9..12], [17.0, 18.0, 19.0]);
assert_eq!(qkv.1[12..16], [6.0, 7.0, 8.0, 9.0]);
assert_eq!(qkv.1[16..20], [20.0, 21.0, 22.0, 23.0]);
let z = &out[1];
assert_eq!(z.0, "m.linear_attn.in_proj_z.weight");
assert_eq!(z.2, nv * dv); assert_eq!(z.1, [10.0, 11.0, 12.0, 13.0, 24.0, 25.0, 26.0, 27.0]);
let wb: Vec<f32> = (0..nk * 2 * r * hid).map(|i| i as f32).collect();
let outb = split_fused_gdn("m.linear_attn.in_proj_ba.weight", &wb, hid, nk, dk, nv, dv).unwrap();
assert_eq!(outb[0].0, "m.linear_attn.in_proj_b.weight");
assert_eq!(outb[0].1, [0.0, 1.0, 4.0, 5.0]);
assert_eq!(outb[1].0, "m.linear_attn.in_proj_a.weight");
assert_eq!(outb[1].1, [2.0, 3.0, 6.0, 7.0]);
}
#[test]
fn lfm2_names_map_to_canonical_layout() {
let c = |s: &str| canon_name(s).unwrap();
assert_eq!(c("model.embedding_norm.weight"), "model.norm.weight");
assert_eq!(c("model.layers.0.operator_norm.weight"), "model.layers.0.input_layernorm.weight");
assert_eq!(c("model.layers.0.ffn_norm.weight"), "model.layers.0.post_attention_layernorm.weight");
assert_eq!(c("model.layers.0.conv.in_proj.weight"), "model.layers.0.short_conv.in_proj.weight");
assert_eq!(c("model.layers.0.conv.conv.weight"), "model.layers.0.short_conv.conv.weight");
assert_eq!(c("model.layers.0.conv.out_proj.weight"), "model.layers.0.short_conv.out_proj.weight");
assert_eq!(c("model.layers.0.feed_forward.w1.weight"), "model.layers.0.mlp.gate_proj.weight");
assert_eq!(c("model.layers.0.feed_forward.w3.weight"), "model.layers.0.mlp.up_proj.weight");
assert_eq!(c("model.layers.0.feed_forward.w2.weight"), "model.layers.0.mlp.down_proj.weight");
assert_eq!(c("model.layers.2.self_attn.out_proj.weight"), "model.layers.2.self_attn.o_proj.weight");
assert_eq!(c("model.layers.2.self_attn.q_layernorm.weight"), "model.layers.2.self_attn.q_norm.weight");
assert_eq!(c("model.layers.2.self_attn.k_layernorm.weight"), "model.layers.2.self_attn.k_norm.weight");
assert_eq!(c("model.layers.2.feed_forward.gate.weight"), "model.layers.2.mlp.gate.weight");
assert_eq!(c("model.layers.2.feed_forward.expert_bias"), "model.layers.2.mlp.expert_bias");
assert_eq!(c("model.layers.2.feed_forward.experts.7.w1.weight"), "model.layers.2.mlp.experts.7.gate_proj.weight");
assert_eq!(c("model.layers.2.feed_forward.experts.7.w2.weight"), "model.layers.2.mlp.experts.7.down_proj.weight");
assert_eq!(c("model.layers.2.self_attn.q_proj.weight"), "model.layers.2.self_attn.q_proj.weight");
assert_eq!(c("model.layers.3.mlp.gate_proj.weight"), "model.layers.3.mlp.gate_proj.weight");
}
#[test]
fn lfm2_moe_arch_routing_and_layers() {
let cfg: serde_json::Value = serde_json::from_str(
r#"{"model_type":"lfm2_moe","hidden_size":2048,"num_hidden_layers":4,
"num_attention_heads":32,"num_key_value_heads":8,"intermediate_size":7168,
"moe_intermediate_size":1792,"vocab_size":128000,"norm_eps":1e-5,
"conv_L_cache":3,"num_experts":32,"num_experts_per_tok":4,
"norm_topk_prob":true,"use_expert_bias":true,"routed_scaling_factor":1.0,
"tie_word_embeddings":true,"rope_parameters":{"rope_theta":5000000},
"layer_types":["conv","conv","full_attention","conv"]}"#,
)
.unwrap();
let arch = build_arch(&cfg).unwrap();
assert_eq!(arch.layer_types[0], LayerType::ShortConv);
assert_eq!(arch.layer_types[2], LayerType::FullAttention);
assert_eq!(arch.head_dim, 64);
assert_eq!(arch.linear_conv_kernel_dim, Some(3));
assert!((arch.rms_norm_eps - 1e-5).abs() < 1e-12);
let moe = arch.moe.as_ref().unwrap();
assert!(moe.router_sigmoid, "lfm2_moe must route with a sigmoid gate");
assert_eq!(moe.top_k, 4);
assert!(moe.norm_topk_prob);
assert_eq!(moe.routed_scaling_factor, None);
}
#[test]
fn q1_ef_bit_identical_on_a_1bit_tensor() {
let (rows, cols) = (4usize, 96usize);
let onebit: Vec<f32> = (0..rows * cols).map(|i| if (i * 7 + 3) % 5 < 2 { 0.25 } else { -0.25 }).collect();
assert_eq!(encode_q1(&onebit, rows, cols), encode_q1_ef(&onebit, rows, cols), "error diffusion must be a no-op on a genuinely 1-bit tensor");
}
#[test]
fn q1s_roundtrip_restores_outliers_and_binarizes_the_rest() {
use cortiq_core::quant::dequant_q1s;
let (rows, cols) = (2usize, 64usize);
let mut vals: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin() * 0.1).collect();
let spikes = [5usize, 40, 70, 120];
for &i in &spikes {
vals[i] = if i % 2 == 0 { 3.0 } else { -3.0 };
}
let keep = spikes.len() as f32 / (rows * cols) as f32;
let bytes = encode_q1s(&vals, rows, cols, keep);
let mut dec = vec![0f32; rows * cols];
dequant_q1s(&bytes, &mut dec);
for &i in &spikes {
assert!((dec[i] - vals[i]).abs() < 0.02, "outlier {i}: {} vs {}", dec[i], vals[i]);
}
for i in 0..rows * cols {
if !spikes.contains(&i) {
assert!(dec[i].abs() < 2.0, "bulk {i} should be a small ±s, got {}", dec[i]);
}
}
}
#[test]
fn parse_quant_variants() {
for q in ["q8", "q8_row", "q8_2f", "q4", "q4_block", "f16"] {
assert!(parse_quant(q).is_ok(), "{q}");
}
assert!(parse_quant("nope").is_err());
}
#[test]
fn q8_row_roundtrips() {
let (o, i) = (4usize, 64usize);
let vals: Vec<f32> = (0..o * i).map(|k| (k as f32 * 0.017).sin() * 2.5).collect();
let bytes = encode_q8_row(&vals, o, i);
assert_eq!(bytes.len(), o * i + o * 2);
let mut back = vec![0f32; o * i];
dequant_q8_row(&bytes, o, i, &mut back);
for (a, b) in vals.iter().zip(&back) {
assert!((a - b).abs() < 0.05, "{a} vs {b}");
}
}
#[test]
fn q8_2f_roundtrips() {
let (o, i) = (8usize, 48usize);
let vals: Vec<f32> = (0..o * i).map(|k| (k as f32 * 0.023).cos() * 1.7).collect();
let bytes = encode_q8_2f(&vals, o, i);
assert_eq!(bytes.len(), o * i + o * 2 + i * 2);
let mut back = vec![0f32; o * i];
dequant_q8_2f(&bytes, o, i, &mut back);
for (a, b) in vals.iter().zip(&back) {
assert!((a - b).abs() < 0.1, "{a} vs {b}");
}
}
#[test]
fn q4_block_roundtrips() {
let vals: Vec<f32> = (0..128).map(|k| (k as f32 * 0.05).sin()).collect();
let bytes = encode_q4_block(&vals);
let mut back = vec![0f32; 128];
dequant_q4_block(&bytes, &mut back);
for (a, b) in vals.iter().zip(&back) {
assert!((a - b).abs() < 0.2, "{a} vs {b}");
}
}
fn tiny_safetensors(tensors: &[(&str, Vec<usize>, Vec<f32>)]) -> Vec<u8> {
let mut header = serde_json::Map::new();
let mut data = Vec::new();
for (name, shape, vals) in tensors {
let start = data.len();
for &v in vals {
data.extend_from_slice(&v.to_le_bytes());
}
header.insert(name.to_string(), serde_json::json!({"dtype":"F32","shape":shape,"data_offsets":[start, data.len()]}));
}
let hjson = serde_json::to_vec(&serde_json::Value::Object(header)).unwrap();
let mut out = (hjson.len() as u64).to_le_bytes().to_vec();
out.extend_from_slice(&hjson);
out.extend_from_slice(&data);
out
}
#[test]
fn convert_tiny_model_end_to_end() {
let dir = std::env::temp_dir().join(format!("cortiq-convtest-{}", std::process::id()));
let _ = fs::remove_dir_all(&dir);
fs::create_dir_all(&dir).unwrap();
fs::write(
dir.join("config.json"),
r#"{"model_type":"llama","hidden_size":64,"num_hidden_layers":1,"num_attention_heads":4,"num_key_value_heads":4,"intermediate_size":128,"vocab_size":32,"rms_norm_eps":0.000001,"tie_word_embeddings":true}"#,
)
.unwrap();
fs::write(dir.join("tokenizer.json"), b"{}").unwrap();
let st = tiny_safetensors(&[("model.embed_tokens.weight", vec![32, 64], (0..32 * 64).map(|k| (k as f32 * 0.01).sin()).collect()), ("model.norm.weight", vec![64], vec![1.0f32; 64])]);
fs::write(dir.join("model.safetensors"), &st).unwrap();
let out = dir.join("m.cmf");
run_convert(dir.to_str().unwrap(), "q8", out.to_str().unwrap(), None, None, None, |_| {}).unwrap();
let model = CmfModel::open(&out).unwrap();
assert_eq!(model.arch().vocab_size, 32);
assert_eq!(model.arch().num_layers, 1);
let _ = fs::remove_dir_all(&dir);
}
}