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_compact(model_path: &str, output: &str) -> anyhow::Result<()> {
let model = Arc::new(CmfModel::open_sharded(model_path)?);
let specs: Vec<TensorSpecRef> = model
.tensors
.iter()
.map(|entry| TensorSpecRef {
name: entry.name.clone(),
dtype: entry.dtype,
shape: entry.shape.iter().map(|&d| d as usize).collect(),
data: model.entry_bytes(entry),
})
.collect();
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!(
"compact: {} tensors\n{model_path} {in_sz:.2} GB -> {output} {out_sz:.2} GB ({:+.1}%)",
specs.len(),
(out_sz / in_sz - 1.0) * 100.0
);
Ok(())
}
pub fn cmd_moe_mask(
model_path: &str,
stats_path: &str,
cover: f64,
name: &str,
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 arch = model.arch().clone();
let Some(moe) = arch.moe.as_ref() else {
bail!("{model_path}: not a MoE model (no arch.moe block)");
};
let ne = moe.num_experts;
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 expert_b = ne.div_ceil(8);
let mut expert_masks: Vec<Vec<u8>> = vec![Vec::new(); arch.num_layers];
let mut kept_total = 0usize;
let mut masked_layers = 0usize;
for (k, counts) in &stats {
let Ok(li) = k.parse::<usize>() else { continue };
if li >= arch.num_layers || counts.len() != ne {
continue;
}
let total: u64 = counts.iter().sum();
if total == 0 {
continue;
}
let mut order: Vec<usize> = (0..ne).collect();
order.sort_unstable_by_key(|&e| std::cmp::Reverse(counts[e]));
let mut bits = vec![0u8; expert_b];
let mut acc = 0u64;
for &e in &order {
bits[e / 8] |= 1 << (e % 8);
kept_total += 1;
acc += counts[e];
if (acc as f64) >= cover * (total as f64) {
break;
}
}
expert_masks[li] = bits;
masked_layers += 1;
}
if masked_layers == 0 {
bail!("no usable layer stats in {stats_path}");
}
let mut catalog = model.masks.clone();
if catalog.masks.iter().any(|m| m.name == name) {
bail!("mask '{name}' already exists in {model_path}");
}
let task_id = catalog.masks.iter().map(|m| m.task_id + 1).max().unwrap_or(1);
let sparsity = 1.0 - kept_total as f32 / (masked_layers * ne) as f32;
catalog.masks.push(cortiq_core::TaskMask {
task_id,
name: name.to_string(),
description: Some(format!(
"MoE expert mask (cover {:.0}%, {} layers)",
cover * 100.0,
masked_layers
)),
sparsity,
quality: None, ffn_masks: vec![vec![0xFF; arch.ffn_mask_bytes()]; arch.num_layers],
head_masks: vec![vec![0xFF; arch.head_mask_bytes()]; arch.num_layers],
layer_gates: vec![true; arch.num_layers],
expert_masks,
parent: None,
has_hot_pack: false,
priority: cortiq_core::MaskPriority::Normal,
});
let specs: Vec<TensorSpecRef> = model
.tensors
.iter()
.map(|entry| TensorSpecRef {
name: entry.name.clone(),
dtype: entry.dtype,
shape: entry.shape.iter().map(|&d| d as usize).collect(),
data: model.entry_bytes(entry),
})
.collect();
CmfModel::write_ref(
output,
&model.header,
&specs,
Some(&catalog),
model.vocab.as_deref(),
)?;
println!(
"moe-mask: '{name}' added ({masked_layers} layers, expert sparsity {:.0}%)\nactivate with: cortiq run {output} --task {name} …",
sparsity * 100.0
);
Ok(())
}
pub fn cmd_moe_defrag(
model_path: &str,
stats_path: Option<&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_text, stats_origin): (String, String) = match stats_path {
Some(p) => (
std::fs::read_to_string(p).with_context(|| format!("reading {p}"))?,
p.to_string(),
),
None => {
let counts = model
.header
.provenance
.as_ref()
.and_then(|p| p.get("moe_defrag"))
.and_then(|d| d.get("routing_counts"))
.ok_or_else(|| {
anyhow::anyhow!(
"--stats not given and the file carries no embedded \
routing counts (provenance.moe_defrag.routing_counts)"
)
})?;
(counts.to_string(), "embedded provenance".to_string())
}
};
let stats: HashMap<String, Vec<u64>> =
serde_json::from_str(&stats_text).with_context(|| format!("parsing {stats_origin}"))?;
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_origin}");
}
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");
}
let mut header = model.header.clone();
let mut kept_per_layer: Vec<(usize, usize)> =
remap.iter().map(|(&li, m)| (li, m.len())).collect();
kept_per_layer.sort_unstable();
let mut routing_counts = serde_json::Map::new();
for (k, counts) in &stats {
let Ok(li) = k.parse::<usize>() else { continue };
let Some(map) = remap.get(&li) else { continue };
let mut remapped = vec![0u64; map.len()];
for (&old, &new) in map {
remapped[new] = counts[old];
}
routing_counts.insert(k.clone(), serde_json::json!(remapped));
}
let prov = serde_json::json!({
"tool": format!("cortiq moe-defrag {}", env!("CARGO_PKG_VERSION")),
"cover": cover,
"stats_hash64": format!("{:016x}", cortiq_core::hash64(stats_text.as_bytes())),
"num_experts_pre": header
.arch
.moe
.as_ref()
.map(|m| m.num_experts)
.unwrap_or(0),
"kept_per_layer": kept_per_layer
.iter()
.map(|&(_, k)| k)
.collect::<Vec<_>>(),
"routing_counts": serde_json::Value::Object(routing_counts),
});
match header.provenance.as_mut() {
Some(serde_json::Value::Object(map)) => {
map.insert("moe_defrag".into(), prov);
}
_ => {
header.provenance = Some(serde_json::json!({ "moe_defrag": prov }));
}
}
CmfModel::write_ref(
output,
&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(())
}