use std::io::{Seek, SeekFrom, Write};
use super::types::{
align_up, write_gguf_string, write_metadata_kv, MetaValue, ALIGNMENT, GGUF_MAGIC, GGUF_VERSION,
};
use crate::quantize::ggml_quants::GgmlType;
#[derive(Debug)]
pub enum WriterError {
Io(std::io::Error),
UnknownTensorIndex { tensor_idx: usize, reserved: usize },
DuplicateTensorPayload { tensor_idx: usize },
MissingTensorPayloads { reserved: usize, streamed: usize },
PayloadSizeMismatch {
tensor_idx: usize,
expected: usize,
actual: usize,
},
PayloadAlreadyActive {
active_tensor_idx: usize,
requested_tensor_idx: usize,
},
NoActivePayload { tensor_idx: usize },
WrongActivePayload {
active_tensor_idx: usize,
requested_tensor_idx: usize,
},
}
impl std::fmt::Display for WriterError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
WriterError::Io(e) => write!(f, "gguf writer I/O: {e}"),
WriterError::UnknownTensorIndex { tensor_idx, reserved } => write!(
f,
"stream_tensor_payload tensor_idx {tensor_idx} >= reserved count {reserved}"
),
WriterError::DuplicateTensorPayload { tensor_idx } => write!(
f,
"stream_tensor_payload called twice for tensor_idx {tensor_idx}"
),
WriterError::MissingTensorPayloads { reserved, streamed } => write!(
f,
"finalize: only {streamed} / {reserved} tensors had payloads streamed"
),
WriterError::PayloadSizeMismatch { tensor_idx, expected, actual } => write!(
f,
"tensor_idx {tensor_idx}: payload size mismatch (expected {expected} bytes, got {actual})"
),
WriterError::PayloadAlreadyActive {
active_tensor_idx,
requested_tensor_idx,
} => write!(
f,
"cannot begin tensor_idx {requested_tensor_idx}: tensor_idx {active_tensor_idx} payload is still active"
),
WriterError::NoActivePayload { tensor_idx } => {
write!(f, "tensor_idx {tensor_idx}: no active chunked payload")
}
WriterError::WrongActivePayload {
active_tensor_idx,
requested_tensor_idx,
} => write!(
f,
"chunk targets tensor_idx {requested_tensor_idx}, but tensor_idx {active_tensor_idx} is active"
),
}
}
}
impl std::error::Error for WriterError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
WriterError::Io(e) => Some(e),
_ => None,
}
}
}
impl From<std::io::Error> for WriterError {
fn from(e: std::io::Error) -> Self {
WriterError::Io(e)
}
}
type Result<T> = std::result::Result<T, WriterError>;
#[derive(Debug, Clone, Copy)]
struct OffsetFixup {
file_pos_of_offset_field: u64,
expected_byte_len: usize,
}
#[derive(Debug, Clone, Copy)]
struct ActivePayload {
tensor_idx: usize,
relative_offset: u64,
written: usize,
}
pub struct GgufWriter<W: Write + Seek> {
writer: W,
fixups: Vec<OffsetFixup>,
tensor_offsets: Vec<Option<u64>>,
tensor_data_start: Option<u64>,
active_payload: Option<ActivePayload>,
}
impl<W: Write + Seek> GgufWriter<W> {
pub fn new(writer: W) -> Self {
Self {
writer,
fixups: Vec::new(),
tensor_offsets: Vec::new(),
tensor_data_start: None,
active_payload: None,
}
}
pub fn write_header(&mut self, tensor_count: u64, kv_count: u64) -> Result<()> {
self.writer.write_all(&GGUF_MAGIC)?;
self.writer.write_all(&GGUF_VERSION.to_le_bytes())?;
self.writer.write_all(&tensor_count.to_le_bytes())?;
self.writer.write_all(&kv_count.to_le_bytes())?;
Ok(())
}
pub fn write_metadata_kv(&mut self, key: &str, value: &MetaValue) -> Result<()> {
write_metadata_kv(&mut self.writer, key, value)?;
Ok(())
}
pub fn reserve_tensor_info(
&mut self,
name: &str,
dims: &[u64],
ggml_type: GgmlType,
) -> Result<usize> {
write_gguf_string(&mut self.writer, name)?;
self.writer.write_all(&(dims.len() as u32).to_le_bytes())?;
for &d in dims {
self.writer.write_all(&d.to_le_bytes())?;
}
let ggml_code: u32 = ggml_type.into();
self.writer.write_all(&ggml_code.to_le_bytes())?;
let file_pos_of_offset_field = self.writer.stream_position()?;
self.writer.write_all(&[0u8; 8])?;
let (rows, n_per_row) = split_rows_and_cols(dims);
let block_size = ggml_type.block_size();
if n_per_row % block_size != 0 {
return Err(WriterError::PayloadSizeMismatch {
tensor_idx: self.fixups.len(),
expected: 0,
actual: 0,
});
}
let expected_byte_len = rows * ggml_type.row_size(n_per_row);
let idx = self.fixups.len();
self.fixups.push(OffsetFixup {
file_pos_of_offset_field,
expected_byte_len,
});
self.tensor_offsets.push(None);
Ok(idx)
}
pub fn pad_to_alignment(&mut self) -> Result<()> {
let cur = self.writer.stream_position()?;
let target = align_up(cur, ALIGNMENT);
let pad = (target - cur) as usize;
if pad > 0 {
let zeros = [0u8; 32];
self.writer.write_all(&zeros[..pad])?;
}
self.tensor_data_start = Some(target);
Ok(())
}
pub fn begin_tensor_payload(&mut self, tensor_idx: usize) -> Result<()> {
if tensor_idx >= self.fixups.len() {
return Err(WriterError::UnknownTensorIndex {
tensor_idx,
reserved: self.fixups.len(),
});
}
if self.tensor_offsets[tensor_idx].is_some() {
return Err(WriterError::DuplicateTensorPayload { tensor_idx });
}
if let Some(active) = self.active_payload {
return Err(WriterError::PayloadAlreadyActive {
active_tensor_idx: active.tensor_idx,
requested_tensor_idx: tensor_idx,
});
}
let data_start = self.tensor_data_start.expect(
"pad_to_alignment must be called before begin_tensor_payload (caller-side bug)",
);
let cur = self.writer.stream_position()?;
self.active_payload = Some(ActivePayload {
tensor_idx,
relative_offset: cur - data_start,
written: 0,
});
Ok(())
}
pub fn write_tensor_payload_chunk(&mut self, tensor_idx: usize, payload: &[u8]) -> Result<()> {
let active = self
.active_payload
.ok_or(WriterError::NoActivePayload { tensor_idx })?;
if active.tensor_idx != tensor_idx {
return Err(WriterError::WrongActivePayload {
active_tensor_idx: active.tensor_idx,
requested_tensor_idx: tensor_idx,
});
}
let expected = self.fixups[tensor_idx].expected_byte_len;
let actual = active
.written
.checked_add(payload.len())
.unwrap_or(usize::MAX);
if actual > expected {
return Err(WriterError::PayloadSizeMismatch {
tensor_idx,
expected,
actual,
});
}
self.writer.write_all(payload)?;
self.active_payload
.as_mut()
.expect("validated active payload")
.written = actual;
Ok(())
}
pub fn finish_tensor_payload(&mut self, tensor_idx: usize) -> Result<()> {
let active = self
.active_payload
.ok_or(WriterError::NoActivePayload { tensor_idx })?;
if active.tensor_idx != tensor_idx {
return Err(WriterError::WrongActivePayload {
active_tensor_idx: active.tensor_idx,
requested_tensor_idx: tensor_idx,
});
}
let expected = self.fixups[tensor_idx].expected_byte_len;
if active.written != expected {
return Err(WriterError::PayloadSizeMismatch {
tensor_idx,
expected,
actual: active.written,
});
}
let after = self.writer.stream_position()?;
let target = align_up(after, ALIGNMENT);
let pad = (target - after) as usize;
if pad > 0 {
let zeros = [0u8; 32];
self.writer.write_all(&zeros[..pad])?;
}
self.tensor_offsets[tensor_idx] = Some(active.relative_offset);
self.active_payload = None;
Ok(())
}
pub fn stream_tensor_payload(&mut self, tensor_idx: usize, payload: &[u8]) -> Result<()> {
if tensor_idx >= self.fixups.len() {
return Err(WriterError::UnknownTensorIndex {
tensor_idx,
reserved: self.fixups.len(),
});
}
if self.tensor_offsets[tensor_idx].is_some() {
return Err(WriterError::DuplicateTensorPayload { tensor_idx });
}
let expected = self.fixups[tensor_idx].expected_byte_len;
if payload.len() != expected {
return Err(WriterError::PayloadSizeMismatch {
tensor_idx,
expected,
actual: payload.len(),
});
}
self.begin_tensor_payload(tensor_idx)?;
self.write_tensor_payload_chunk(tensor_idx, payload)?;
self.finish_tensor_payload(tensor_idx)
}
pub fn finalize(&mut self) -> Result<()> {
if let Some(active) = self.active_payload {
return Err(WriterError::PayloadSizeMismatch {
tensor_idx: active.tensor_idx,
expected: self.fixups[active.tensor_idx].expected_byte_len,
actual: active.written,
});
}
let streamed = self.tensor_offsets.iter().filter(|o| o.is_some()).count();
if streamed != self.fixups.len() {
return Err(WriterError::MissingTensorPayloads {
reserved: self.fixups.len(),
streamed,
});
}
let eof = self.writer.stream_position()?;
for (idx, fixup) in self.fixups.iter().enumerate() {
let off = self.tensor_offsets[idx].expect("checked above");
self.writer
.seek(SeekFrom::Start(fixup.file_pos_of_offset_field))?;
self.writer.write_all(&off.to_le_bytes())?;
}
self.writer.seek(SeekFrom::Start(eof))?;
self.writer.flush()?;
Ok(())
}
pub fn into_inner(self) -> W {
self.writer
}
#[cfg(test)]
pub fn tensor_offsets(&self) -> &[Option<u64>] {
&self.tensor_offsets
}
}
fn split_rows_and_cols(dims: &[u64]) -> (usize, usize) {
match dims.len() {
0 => (0, 0),
1 => (1, dims[0] as usize),
_ => {
let n_per_row = dims[0] as usize;
let rows: usize = dims[1..].iter().map(|&d| d as usize).product();
(rows, n_per_row)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
fn build_tiny_gguf() -> Vec<u8> {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(1, 1).unwrap();
w.write_metadata_kv("general.architecture", &MetaValue::String("test".into()))
.unwrap();
let idx = w
.reserve_tensor_info("test.weight", &[32], GgmlType::Q4_0)
.unwrap();
assert_eq!(idx, 0);
w.pad_to_alignment().unwrap();
let payload: Vec<u8> = (0u8..18).collect();
w.stream_tensor_payload(0, &payload).unwrap();
w.finalize().unwrap();
w.into_inner().into_inner()
}
#[test]
fn round_trip_header_and_offset() {
let bytes = build_tiny_gguf();
assert_eq!(&bytes[0..4], &GGUF_MAGIC);
assert_eq!(
u32::from_le_bytes([bytes[4], bytes[5], bytes[6], bytes[7]]),
GGUF_VERSION
);
assert_eq!(u64::from_le_bytes(bytes[8..16].try_into().unwrap()), 1);
assert_eq!(u64::from_le_bytes(bytes[16..24].try_into().unwrap()), 1);
}
#[test]
fn round_trip_via_reader() {
use std::io::{Read, Seek, SeekFrom, Write};
let bytes = build_tiny_gguf();
let tmp = tempfile::NamedTempFile::new().unwrap();
{
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 hf2q GGUF");
assert_eq!(gguf.metadata_count(), 1);
assert_eq!(gguf.metadata_string("general.architecture"), Some("test"));
assert_eq!(gguf.tensor_count(), 1);
let info = gguf.tensor_info("test.weight").expect("tensor present");
assert_eq!(info.shape, vec![32]);
assert_eq!(info.byte_len, 18);
assert_eq!(info.ggml_type as u32, 2);
let mut f = std::fs::File::open(tmp.path()).unwrap();
let mut all = Vec::new();
f.seek(SeekFrom::Start(0)).unwrap();
f.read_to_end(&mut all).unwrap();
let expected: Vec<u8> = (0u8..18).collect();
let found = all
.windows(18)
.position(|w| w == expected.as_slice())
.expect("payload run 0..18 must appear exactly once in the file");
assert!(found > 24, "payload found inside header region: {found}");
assert_eq!(found as u64 % ALIGNMENT, 0);
}
#[test]
fn payload_size_mismatch_is_typed_error() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(1, 0).unwrap();
w.reserve_tensor_info("bad.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.pad_to_alignment().unwrap();
let bad_payload = vec![0u8; 17];
let err = w.stream_tensor_payload(0, &bad_payload).unwrap_err();
match err {
WriterError::PayloadSizeMismatch {
expected, actual, ..
} => {
assert_eq!(expected, 18);
assert_eq!(actual, 17);
}
other => panic!("expected PayloadSizeMismatch, got {other:?}"),
}
}
#[test]
fn chunked_payload_requires_exact_aggregate_length() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(1, 0).unwrap();
w.reserve_tensor_info("chunked.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.pad_to_alignment().unwrap();
w.begin_tensor_payload(0).unwrap();
w.write_tensor_payload_chunk(0, &[1; 8]).unwrap();
let err = w.finish_tensor_payload(0).unwrap_err();
assert!(matches!(
err,
WriterError::PayloadSizeMismatch {
tensor_idx: 0,
expected: 18,
actual: 8,
}
));
}
#[test]
fn chunked_payload_matches_single_payload_bytes() {
fn write(chunked: bool) -> Vec<u8> {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(1, 0).unwrap();
w.reserve_tensor_info("chunked.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.pad_to_alignment().unwrap();
let payload: Vec<u8> = (0..18).collect();
if chunked {
w.begin_tensor_payload(0).unwrap();
w.write_tensor_payload_chunk(0, &payload[..7]).unwrap();
w.write_tensor_payload_chunk(0, &payload[7..]).unwrap();
w.finish_tensor_payload(0).unwrap();
} else {
w.stream_tensor_payload(0, &payload).unwrap();
}
w.finalize().unwrap();
w.into_inner().into_inner()
}
assert_eq!(write(true), write(false));
}
#[test]
fn finalize_without_streaming_errors() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(1, 0).unwrap();
w.reserve_tensor_info("unstreamed.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.pad_to_alignment().unwrap();
let err = w.finalize().unwrap_err();
match err {
WriterError::MissingTensorPayloads { reserved, streamed } => {
assert_eq!(reserved, 1);
assert_eq!(streamed, 0);
}
other => panic!("expected MissingTensorPayloads, got {other:?}"),
}
}
#[test]
fn duplicate_stream_errors() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(1, 0).unwrap();
w.reserve_tensor_info("t.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.pad_to_alignment().unwrap();
let payload: Vec<u8> = vec![0u8; 18];
w.stream_tensor_payload(0, &payload).unwrap();
let err = w.stream_tensor_payload(0, &payload).unwrap_err();
match err {
WriterError::DuplicateTensorPayload { tensor_idx } => {
assert_eq!(tensor_idx, 0);
}
other => panic!("expected DuplicateTensorPayload, got {other:?}"),
}
}
#[test]
fn unknown_tensor_index_errors() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(0, 0).unwrap();
w.pad_to_alignment().unwrap();
let err = w.stream_tensor_payload(0, &[]).unwrap_err();
match err {
WriterError::UnknownTensorIndex {
tensor_idx,
reserved,
} => {
assert_eq!(tensor_idx, 0);
assert_eq!(reserved, 0);
}
other => panic!("expected UnknownTensorIndex, got {other:?}"),
}
}
#[test]
fn multi_tensor_offsets_seek_back_correctly() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(2, 0).unwrap();
w.reserve_tensor_info("a.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.reserve_tensor_info("b.weight", &[32], GgmlType::Q4_0)
.unwrap();
w.pad_to_alignment().unwrap();
let payload_a: Vec<u8> = (0u8..18).collect();
let payload_b: Vec<u8> = (100u8..118).collect();
w.stream_tensor_payload(0, &payload_a).unwrap();
w.stream_tensor_payload(1, &payload_b).unwrap();
assert_eq!(w.tensor_offsets()[0], Some(0));
assert_eq!(w.tensor_offsets()[1], Some(32));
w.finalize().unwrap();
}
#[test]
fn header_only_no_tensors() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
w.write_header(0, 0).unwrap();
w.pad_to_alignment().unwrap();
w.finalize().unwrap();
let bytes = w.into_inner().into_inner();
assert_eq!(&bytes[0..4], &GGUF_MAGIC);
assert_eq!(bytes.len() % ALIGNMENT as usize, 0);
assert_eq!(bytes.len(), 32);
}
#[test]
fn i32_hash_route_tensor_preserves_wire_type_and_payload() {
let buf = Cursor::new(Vec::new());
let mut w = GgufWriter::new(buf);
let name = "blk.0.ffn_gate_tid2eid.weight";
w.write_header(1, 0).unwrap();
w.reserve_tensor_info(name, &[3], GgmlType::I32).unwrap();
w.pad_to_alignment().unwrap();
let payload: Vec<u8> = [0_i32, 3, 255]
.into_iter()
.flat_map(i32::to_le_bytes)
.collect();
w.stream_tensor_payload(0, &payload).unwrap();
w.finalize().unwrap();
let bytes = w.into_inner().into_inner();
let type_offset = 24 + 8 + name.len() + 4 + 8;
assert_eq!(
u32::from_le_bytes(bytes[type_offset..type_offset + 4].try_into().unwrap()),
26
);
let data_start = align_up((type_offset + 4 + 8) as u64, ALIGNMENT) as usize;
assert_eq!(&bytes[data_start..data_start + payload.len()], payload);
}
}