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#[derive(Debug)]
13struct TensorInfo {
14 name: String,
15 shape: Vec<u64>,
16 dtype: GGMLType,
17 data: Vec<u8>,
19}
20
21pub 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 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 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 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 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 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 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 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}