use anyhow::{bail, Context};
use cortiq_core::{CmfModel, TensorDtype, TensorSpecRef};
use std::collections::HashMap;
use std::sync::Arc;
fn expert_parts(name: &str) -> Option<(usize, usize, &str)> {
let t = name.strip_prefix("model.layers.")?;
let (li, t) = t.split_once('.')?;
let t = t.strip_prefix("mlp.experts.")?;
let (e, rest) = t.split_once('.')?;
Some((li.parse().ok()?, e.parse().ok()?, rest))
}
fn router_layer(name: &str) -> Option<usize> {
let t = name.strip_prefix("model.layers.")?;
let (li, t) = t.split_once('.')?;
(t == "mlp.gate.weight").then(|| li.parse().ok())?
}
pub fn cmd_moe_defrag(
model_path: &str,
stats_path: &str,
cover: f64,
output: &str,
) -> anyhow::Result<()> {
if !(cover > 0.0 && cover <= 1.0) {
bail!("--cover must be in (0, 1]");
}
let model = Arc::new(CmfModel::open_sharded(model_path)?);
let stats: HashMap<String, Vec<u64>> = serde_json::from_str(
&std::fs::read_to_string(stats_path).with_context(|| format!("reading {stats_path}"))?,
)
.with_context(|| format!("parsing {stats_path}"))?;
let mut remap: HashMap<usize, HashMap<usize, usize>> = HashMap::new();
for (k, counts) in &stats {
let li: usize = match k.parse() {
Ok(v) => v,
Err(_) => continue,
};
let total: u64 = counts.iter().sum();
if total == 0 {
continue;
}
let mut order: Vec<usize> = (0..counts.len()).collect();
order.sort_unstable_by_key(|&e| std::cmp::Reverse(counts[e]));
let mut acc = 0u64;
let mut kept = Vec::new();
for &e in &order {
kept.push(e);
acc += counts[e];
if (acc as f64) >= cover * (total as f64) {
break;
}
}
kept.sort_unstable();
remap.insert(
li,
kept.iter().enumerate().map(|(new, &old)| (old, new)).collect(),
);
}
if remap.is_empty() {
bail!("no usable layer stats in {stats_path}");
}
let mut routers: HashMap<usize, (Vec<u8>, Vec<usize>)> = HashMap::new(); for (ti, entry) in model.tensors.iter().enumerate() {
let Some(li) = router_layer(&entry.name) else {
continue;
};
let Some(map) = remap.get(&li) else { continue };
if entry.dtype != TensorDtype::F32 {
bail!("{}: router dtype {:?} != F32", entry.name, entry.dtype);
}
let ne = entry.shape[0] as usize;
let hidden = entry.shape[1] as usize;
let src = model.entry_bytes(entry);
let row = hidden * 4;
let mut kept: Vec<usize> = map.keys().copied().collect();
kept.sort_unstable();
let mut data = Vec::with_capacity(kept.len() * row);
for &old in &kept {
if old >= ne {
bail!("{}: stats index {old} >= {ne} rows", entry.name);
}
data.extend_from_slice(&src[old * row..(old + 1) * row]);
}
routers.insert(ti, (data, vec![kept.len(), hidden]));
}
let mut specs: Vec<TensorSpecRef> = Vec::new();
let mut dropped = 0usize;
let mut kept_experts = 0usize;
for (ti, entry) in model.tensors.iter().enumerate() {
if let Some((data, shape)) = routers.get(&ti) {
specs.push(TensorSpecRef {
name: entry.name.clone(),
dtype: entry.dtype,
shape: shape.clone(),
data,
});
continue;
}
if let Some((li, e, rest)) = expert_parts(&entry.name) {
if let Some(map) = remap.get(&li) {
match map.get(&e) {
Some(&new) => {
kept_experts += 1;
specs.push(TensorSpecRef {
name: format!("model.layers.{li}.mlp.experts.{new}.{rest}"),
dtype: entry.dtype,
shape: entry.shape.iter().map(|&d| d as usize).collect(),
data: model.entry_bytes(entry),
});
}
None => dropped += 1,
}
continue;
}
}
specs.push(TensorSpecRef {
name: entry.name.clone(),
dtype: entry.dtype,
shape: entry.shape.iter().map(|&d| d as usize).collect(),
data: model.entry_bytes(entry),
});
}
if dropped == 0 {
bail!("nothing to drop — check the stats file / --cover");
}
CmfModel::write_ref(
output,
&model.header,
&specs,
Some(&model.masks),
model.vocab.as_deref(),
)?;
let in_sz = std::fs::metadata(model_path)?.len() as f64 / 1e9;
let out_sz = std::fs::metadata(output)?.len() as f64 / 1e9;
println!(
"moe-defrag: kept {kept_experts} expert tensors, dropped {dropped} ({} MoE layers, cover {:.0}%)\n{model_path} {in_sz:.1} GB -> {output} {out_sz:.1} GB ({:+.0}%)",
remap.len(),
cover * 100.0,
(out_sz / in_sz - 1.0) * 100.0
);
Ok(())
}