Skip to main content

zsfm_gguf/
writer.rs

1use std::io::{self, Seek, Write};
2
3use byteorder::{LittleEndian, WriteBytesExt};
4
5use super::types::{GGMLType, GGUFMetaValue, GGUFValueType};
6
7const GGUF_MAGIC: &[u8; 4] = b"GGUF";
8const GGUF_VERSION: u32 = 3;
9const ALIGNMENT: u64 = 32;
10
11/// Describes one tensor's position in the GGUF data section.
12#[derive(Debug)]
13struct TensorInfo {
14    name: String,
15    shape: Vec<u64>,
16    dtype: GGMLType,
17    /// Raw tensor bytes (row-major, little-endian).
18    data: Vec<u8>,
19}
20
21/// Streaming GGUF v3 writer.
22///
23/// Call [`Self::add_metadata`] for every key-value pair, then [`Self::add_tensor`] for
24/// every tensor, then [`Self::write_to`] to flush the complete file.
25pub struct GGUFWriter {
26    metadata: Vec<(String, GGUFMetaValue)>,
27    tensors: Vec<TensorInfo>,
28}
29
30impl Default for GGUFWriter {
31    fn default() -> Self {
32        Self::new()
33    }
34}
35
36impl GGUFWriter {
37    pub fn new() -> Self {
38        Self {
39            metadata: Vec::new(),
40            tensors: Vec::new(),
41        }
42    }
43
44    pub fn add_metadata(&mut self, key: impl Into<String>, value: GGUFMetaValue) {
45        self.metadata.push((key.into(), value));
46    }
47
48    /// Buffer a tensor. `data` must already be in the target dtype byte layout.
49    pub fn add_tensor(
50        &mut self,
51        name: impl Into<String>,
52        shape: Vec<u64>,
53        dtype: GGMLType,
54        data: Vec<u8>,
55    ) {
56        self.tensors.push(TensorInfo {
57            name: name.into(),
58            shape,
59            dtype,
60            data,
61        });
62    }
63
64    /// Serialize the complete GGUF file to `writer`.
65    ///
66    /// Tensors are laid out in sorted-name order regardless of insertion order, so
67    /// converting the same source twice yields byte-identical files (several converters
68    /// feed tensors in `safetensors`' HashMap iteration order, which is randomized per
69    /// process).
70    pub fn write_to<W: Write + Seek>(&self, writer: &mut W) -> anyhow::Result<()> {
71        let mut order: Vec<&TensorInfo> = self.tensors.iter().collect();
72        order.sort_by(|a, b| a.name.cmp(&b.name));
73
74        // --- header ---
75        writer.write_all(GGUF_MAGIC)?;
76        writer.write_u32::<LittleEndian>(GGUF_VERSION)?;
77        writer.write_u64::<LittleEndian>(self.tensors.len() as u64)?;
78        writer.write_u64::<LittleEndian>(self.metadata.len() as u64)?;
79
80        // --- metadata key-value pairs ---
81        for (key, value) in &self.metadata {
82            write_string(writer, key)?;
83            writer.write_u32::<LittleEndian>(value.value_type() as u32)?;
84            write_value(writer, value)?;
85        }
86
87        // --- tensor info ---
88        let mut data_offset = 0u64;
89        for t in &order {
90            write_string(writer, &t.name)?;
91            writer.write_u32::<LittleEndian>(t.shape.len() as u32)?;
92            for &dim in &t.shape {
93                writer.write_u64::<LittleEndian>(dim)?;
94            }
95            writer.write_u32::<LittleEndian>(t.dtype as u32)?;
96            writer.write_u64::<LittleEndian>(data_offset)?;
97            data_offset += round_up(t.data.len() as u64, ALIGNMENT);
98        }
99
100        // --- align to ALIGNMENT before tensor data ---
101        let pos = writer.stream_position()?;
102        let aligned = round_up(pos, ALIGNMENT);
103        if aligned > pos {
104            let pad = vec![0u8; (aligned - pos) as usize];
105            writer.write_all(&pad)?;
106        }
107
108        // --- tensor data (each padded to ALIGNMENT) ---
109        for t in &order {
110            writer.write_all(&t.data)?;
111            let remainder = t.data.len() as u64 % ALIGNMENT;
112            if remainder != 0 {
113                let pad = vec![0u8; (ALIGNMENT - remainder) as usize];
114                writer.write_all(&pad)?;
115            }
116        }
117
118        Ok(())
119    }
120}
121
122fn round_up(value: u64, align: u64) -> u64 {
123    (value + align - 1) / align * align
124}
125
126fn write_string<W: Write>(writer: &mut W, s: &str) -> io::Result<()> {
127    writer.write_u64::<LittleEndian>(s.len() as u64)?;
128    writer.write_all(s.as_bytes())
129}
130
131fn write_value<W: Write>(writer: &mut W, value: &GGUFMetaValue) -> anyhow::Result<()> {
132    match value {
133        GGUFMetaValue::Uint8(v) => writer.write_u8(*v)?,
134        GGUFMetaValue::Int8(v) => writer.write_i8(*v)?,
135        GGUFMetaValue::Uint16(v) => writer.write_u16::<LittleEndian>(*v)?,
136        GGUFMetaValue::Int16(v) => writer.write_i16::<LittleEndian>(*v)?,
137        GGUFMetaValue::Uint32(v) => writer.write_u32::<LittleEndian>(*v)?,
138        GGUFMetaValue::Int32(v) => writer.write_i32::<LittleEndian>(*v)?,
139        GGUFMetaValue::Float32(v) => writer.write_f32::<LittleEndian>(*v)?,
140        GGUFMetaValue::Bool(v) => writer.write_u8(*v as u8)?,
141        GGUFMetaValue::String(v) => write_string(writer, v)?,
142        GGUFMetaValue::Uint64(v) => writer.write_u64::<LittleEndian>(*v)?,
143        GGUFMetaValue::Int64(v) => writer.write_i64::<LittleEndian>(*v)?,
144        GGUFMetaValue::Float64(v) => writer.write_f64::<LittleEndian>(*v)?,
145        GGUFMetaValue::ArrayUint32(arr) => {
146            writer.write_u32::<LittleEndian>(GGUFValueType::Uint32 as u32)?;
147            writer.write_u64::<LittleEndian>(arr.len() as u64)?;
148            for v in arr {
149                writer.write_u32::<LittleEndian>(*v)?;
150            }
151        }
152        GGUFMetaValue::ArrayString(arr) => {
153            writer.write_u32::<LittleEndian>(GGUFValueType::String as u32)?;
154            writer.write_u64::<LittleEndian>(arr.len() as u64)?;
155            for s in arr {
156                write_string(writer, s)?;
157            }
158        }
159        GGUFMetaValue::ArrayFloat32(arr) => {
160            writer.write_u32::<LittleEndian>(GGUFValueType::Float32 as u32)?;
161            writer.write_u64::<LittleEndian>(arr.len() as u64)?;
162            for v in arr {
163                writer.write_f32::<LittleEndian>(*v)?;
164            }
165        }
166    }
167    Ok(())
168}