use std::io::{Seek, Write};
use crate::backends::gguf::types::MetaValue;
use crate::backends::gguf::writer::GgufWriter;
use crate::quantize::ggml_quants::GgmlType;
use super::accumulator::{Accumulator, AccumulatorRegistry};
use super::error::ImatrixError;
pub const KV_KEY_TYPE: &str = "general.type";
pub const KV_KEY_DATASETS: &str = "imatrix.datasets";
pub const KV_KEY_CHUNK_COUNT: &str = "imatrix.chunk_count";
pub const KV_KEY_CHUNK_SIZE: &str = "imatrix.chunk_size";
pub const KV_VALUE_TYPE: &str = "imatrix";
pub fn write_imatrix<W: Write + Seek>(
sink: W,
registry: &AccumulatorRegistry,
datasets: &[String],
chunk_count: u32,
chunk_size: u32,
) -> Result<W, ImatrixError> {
let mut w = GgufWriter::new(sink);
let storable: Vec<(&str, &Accumulator)> =
registry.iter().filter(|(_, acc)| acc.has_data()).collect();
let tensor_count = storable.len() as u64 * 2; let kv_count: u64 = 4;
w.write_header(tensor_count, kv_count)?;
w.write_metadata_kv(KV_KEY_TYPE, &MetaValue::String(KV_VALUE_TYPE.to_string()))?;
w.write_metadata_kv(KV_KEY_DATASETS, &MetaValue::ArrayString(datasets.to_vec()))?;
w.write_metadata_kv(KV_KEY_CHUNK_COUNT, &MetaValue::U32(chunk_count))?;
w.write_metadata_kv(KV_KEY_CHUNK_SIZE, &MetaValue::U32(chunk_size))?;
let mut payload_idx: Vec<(usize, usize)> = Vec::with_capacity(storable.len());
for (_, acc) in &storable {
let in_sum2_name = format!("{}.in_sum2", acc.name);
let counts_name = format!("{}.counts", acc.name);
let in_sum2_idx = w.reserve_tensor_info(
&in_sum2_name,
&[acc.n_per_row as u64, acc.n_mat as u64],
GgmlType::F32,
)?;
let counts_idx =
w.reserve_tensor_info(&counts_name, &[1u64, acc.n_mat as u64], GgmlType::F32)?;
payload_idx.push((in_sum2_idx, counts_idx));
}
w.pad_to_alignment()?;
for ((_, acc), (in_sum2_idx, counts_idx)) in storable.iter().zip(payload_idx.iter()) {
let mut in_sum2_bytes = Vec::with_capacity(acc.values.len() * 4);
for v in &acc.values {
in_sum2_bytes.extend_from_slice(&v.to_le_bytes());
}
w.stream_tensor_payload(*in_sum2_idx, &in_sum2_bytes)?;
let mut counts_bytes = Vec::with_capacity(acc.counts.len() * 4);
for &c in &acc.counts {
counts_bytes.extend_from_slice(&(c as f32).to_le_bytes());
}
w.stream_tensor_payload(*counts_idx, &counts_bytes)?;
}
w.finalize()?;
Ok(w.into_inner())
}
pub fn write_imatrix_to_path(
path: &std::path::Path,
registry: &AccumulatorRegistry,
datasets: &[String],
chunk_count: u32,
chunk_size: u32,
) -> Result<(), ImatrixError> {
let f = std::fs::File::create(path)?;
write_imatrix(f, registry, datasets, chunk_count, chunk_size)?;
Ok(())
}
pub fn estimate_kv_bytes(datasets: &[String]) -> usize {
fn str_size(s: &str) -> usize {
8 + s.len()
}
let mut bytes = 0;
bytes += str_size(KV_KEY_TYPE) + 4 + str_size(KV_VALUE_TYPE);
bytes += str_size(KV_KEY_DATASETS) + 4 + 4 + 8;
for d in datasets {
bytes += str_size(d);
}
bytes += str_size(KV_KEY_CHUNK_COUNT) + 4 + 4;
bytes += str_size(KV_KEY_CHUNK_SIZE) + 4 + 4;
bytes
}
#[cfg(test)]
fn write_kv_to_vec(key: &str, value: &MetaValue) -> Vec<u8> {
use crate::backends::gguf::types::write_metadata_kv;
let mut buf = Vec::new();
write_metadata_kv(&mut buf, key, value).expect("in-memory write cannot fail");
buf
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn round_trip_minimal_imatrix() {
let mut reg = AccumulatorRegistry::new();
let acc = reg.register("blk.0.attn_q.weight", 8, 1).unwrap();
for j in 0..8 {
let row: Vec<f32> = (0..8).map(|k| (j + k) as f32).collect();
acc.absorb_dense(&row).unwrap();
}
let buf = Cursor::new(Vec::new());
let inner = write_imatrix(
buf,
®,
&["cdv3".to_string()],
1,
512,
)
.unwrap();
let bytes = inner.into_inner();
let tmp = tempfile::NamedTempFile::new().unwrap();
{
use std::io::Write;
let mut f = std::fs::File::create(tmp.path()).unwrap();
f.write_all(&bytes).unwrap();
f.flush().unwrap();
}
let gguf = mlx_native::gguf::GgufFile::open(tmp.path()).expect("parse imatrix gguf");
assert_eq!(gguf.metadata_string("general.type"), Some("imatrix"));
assert_eq!(gguf.metadata_u32("imatrix.chunk_count"), Some(1));
assert_eq!(gguf.metadata_u32("imatrix.chunk_size"), Some(512));
assert_eq!(gguf.tensor_count(), 2);
assert!(gguf.tensor_info("blk.0.attn_q.weight.in_sum2").is_some());
assert!(gguf.tensor_info("blk.0.attn_q.weight.counts").is_some());
let info = gguf
.tensor_info("blk.0.attn_q.weight.in_sum2")
.expect("in_sum2 present");
assert_eq!(info.shape, vec![1, 8]);
let counts_info = gguf
.tensor_info("blk.0.attn_q.weight.counts")
.expect("counts present");
assert_eq!(counts_info.shape, vec![1, 1]);
}
#[test]
fn empty_registry_writes_valid_gguf() {
let reg = AccumulatorRegistry::new();
let buf = Cursor::new(Vec::new());
let inner = write_imatrix(buf, ®, &["cdv3".to_string()], 0, 512).unwrap();
let bytes = inner.into_inner();
let tmp = tempfile::NamedTempFile::new().unwrap();
{
use std::io::Write;
let mut f = std::fs::File::create(tmp.path()).unwrap();
f.write_all(&bytes).unwrap();
f.flush().unwrap();
}
let gguf = mlx_native::gguf::GgufFile::open(tmp.path()).expect("parse empty imatrix");
assert_eq!(gguf.metadata_string("general.type"), Some("imatrix"));
assert_eq!(gguf.tensor_count(), 0);
}
#[test]
fn moe_accumulator_writes_per_expert_shape() {
let mut reg = AccumulatorRegistry::new();
let acc = reg
.register("blk.0.ffn_gate_exps.weight", 4, 8)
.unwrap();
acc.absorb_moe(0, &[1.0, 2.0, 3.0, 4.0]).unwrap();
acc.absorb_moe(3, &[0.5, 0.5, 0.5, 0.5]).unwrap();
let buf = Cursor::new(Vec::new());
let inner = write_imatrix(buf, ®, &["cdv3".to_string()], 1, 512).unwrap();
let bytes = inner.into_inner();
let tmp = tempfile::NamedTempFile::new().unwrap();
{
use std::io::Write;
let mut f = std::fs::File::create(tmp.path()).unwrap();
f.write_all(&bytes).unwrap();
f.flush().unwrap();
}
let gguf = mlx_native::gguf::GgufFile::open(tmp.path()).expect("parse moe imatrix");
let info = gguf
.tensor_info("blk.0.ffn_gate_exps.weight.in_sum2")
.unwrap();
assert_eq!(info.shape, vec![8, 4]);
let counts = gguf
.tensor_info("blk.0.ffn_gate_exps.weight.counts")
.unwrap();
assert_eq!(counts.shape, vec![8, 1]);
}
#[test]
fn estimate_kv_bytes_is_accurate() {
let datasets = vec!["cdv3".to_string(), "mudler".to_string()];
let mut actual = 0;
actual += write_kv_to_vec(KV_KEY_TYPE, &MetaValue::String(KV_VALUE_TYPE.to_string())).len();
actual += write_kv_to_vec(KV_KEY_DATASETS, &MetaValue::ArrayString(datasets.clone())).len();
actual += write_kv_to_vec(KV_KEY_CHUNK_COUNT, &MetaValue::U32(42)).len();
actual += write_kv_to_vec(KV_KEY_CHUNK_SIZE, &MetaValue::U32(512)).len();
let estimate = estimate_kv_bytes(&datasets);
assert_eq!(estimate, actual, "estimate {estimate} vs actual {actual}");
}
}