1use std::collections::HashMap;
12use std::io::{Read, Write};
13use std::path::Path;
14
15use anyhow::{Context, Result};
16use serde::{Deserialize, Serialize};
17
18use crate::model::{DataType, Model, TensorData};
19
20const MAGIC: &[u8] = b"MODELC";
21const VERSION: u32 = 2;
22const FLAG_COMPRESSED: u32 = 1;
23
24#[derive(Debug, Clone, Serialize, Deserialize)]
25pub struct ArtifactHeader {
26 pub name: String,
27 pub architecture: String,
28 pub metadata: HashMap<String, String>,
29 pub tensors: Vec<ArtifactTensor>,
30}
31
32#[derive(Debug, Clone, Serialize, Deserialize)]
33pub struct ArtifactTensor {
34 pub name: String,
35 pub shape: Vec<usize>,
36 pub dtype: String,
37 pub offset: u64,
38 pub length: u64,
39}
40
41pub fn pack(model: &Model, path: &Path, compress: bool) -> Result<()> {
43 let mut file = std::fs::File::create(path)
44 .with_context(|| format!("failed to create artifact at {:?}", path))?;
45
46 let mut tensors = Vec::with_capacity(model.tensors.len());
48 let mut data_offset: u64 = 0;
49
50 let mut names: Vec<&String> = model.tensors.keys().collect();
52 names.sort();
53
54 for name in &names {
55 let td = &model.tensors[*name];
56 let length = td.data.len() as u64;
57 tensors.push(ArtifactTensor {
58 name: name.to_string(),
59 shape: td.shape.clone(),
60 dtype: data_type_to_str(td.dtype),
61 offset: data_offset,
62 length,
63 });
64 data_offset += length;
65 }
66
67 let header = ArtifactHeader {
68 name: model.name.clone(),
69 architecture: model.architecture.clone(),
70 metadata: model.metadata.clone(),
71 tensors,
72 };
73
74 let header_json = serde_json::to_vec(&header).context("failed to serialize artifact header")?;
75 let header_len = header_json.len() as u64;
76
77 let mut raw_blob = Vec::with_capacity(data_offset as usize);
79 for name in &names {
80 let td = &model.tensors[*name];
81 raw_blob.extend_from_slice(&td.data);
82 }
83
84 let flags = if compress { FLAG_COMPRESSED } else { 0 };
85 let data_blob: Vec<u8> = if compress {
86 zstd::encode_all(&raw_blob[..], 3).context("failed to compress artifact data")?
87 } else {
88 raw_blob
89 };
90
91 file.write_all(MAGIC)?;
93 file.write_all(&VERSION.to_le_bytes())?;
94 file.write_all(&flags.to_le_bytes())?;
95 file.write_all(&header_len.to_le_bytes())?;
96 file.write_all(&header_json)?;
97 file.write_all(&data_blob)?;
98
99 file.flush().context("failed to flush artifact file")?;
100 Ok(())
101}
102
103pub fn read_header(path: &Path) -> Result<ArtifactHeader> {
105 let mut file = std::fs::File::open(path)
106 .with_context(|| format!("failed to open artifact at {:?}", path))?;
107
108 let mut magic = [0u8; 6];
109 file.read_exact(&mut magic)
110 .context("artifact file too short (magic)")?;
111 if magic != MAGIC {
112 anyhow::bail!("invalid artifact magic bytes (not a .modelc file)");
113 }
114
115 let mut version_bytes = [0u8; 4];
116 file.read_exact(&mut version_bytes)
117 .context("artifact file too short (version)")?;
118 let version = u32::from_le_bytes(version_bytes);
119
120 let _flags = if version >= 2 {
121 let mut flags_bytes = [0u8; 4];
122 file.read_exact(&mut flags_bytes)
123 .context("artifact file too short (flags)")?;
124 u32::from_le_bytes(flags_bytes)
125 } else {
126 0
127 };
128
129 if version != 1 && version != 2 {
130 anyhow::bail!("unsupported artifact version {} (expected 1 or 2)", version);
131 }
132
133 let mut header_len_bytes = [0u8; 8];
134 file.read_exact(&mut header_len_bytes)
135 .context("artifact file too short (header length)")?;
136 let header_len = u64::from_le_bytes(header_len_bytes);
137
138 let mut header_json = vec![0u8; header_len as usize];
139 file.read_exact(&mut header_json)
140 .context("artifact file too short (header)")?;
141 let header: ArtifactHeader =
142 serde_json::from_slice(&header_json).context("failed to deserialize artifact header")?;
143
144 Ok(header)
145}
146
147pub fn unpack(path: &Path) -> Result<Model> {
149 let mut file = std::fs::File::open(path)
150 .with_context(|| format!("failed to open artifact at {:?}", path))?;
151
152 let mut magic = [0u8; 6];
154 file.read_exact(&mut magic)
155 .context("artifact file too short (magic)")?;
156 if magic != MAGIC {
157 anyhow::bail!("invalid artifact magic bytes (not a .modelc file)");
158 }
159
160 let mut version_bytes = [0u8; 4];
162 file.read_exact(&mut version_bytes)
163 .context("artifact file too short (version)")?;
164 let version = u32::from_le_bytes(version_bytes);
165
166 let flags = if version >= 2 {
168 let mut flags_bytes = [0u8; 4];
169 file.read_exact(&mut flags_bytes)
170 .context("artifact file too short (flags)")?;
171 u32::from_le_bytes(flags_bytes)
172 } else {
173 0
174 };
175
176 if version != 1 && version != 2 {
177 anyhow::bail!("unsupported artifact version {} (expected 1 or 2)", version);
178 }
179
180 let mut header_len_bytes = [0u8; 8];
182 file.read_exact(&mut header_len_bytes)
183 .context("artifact file too short (header length)")?;
184 let header_len = u64::from_le_bytes(header_len_bytes);
185
186 let mut header_json = vec![0u8; header_len as usize];
188 file.read_exact(&mut header_json)
189 .context("artifact file too short (header)")?;
190 let header: ArtifactHeader =
191 serde_json::from_slice(&header_json).context("failed to deserialize artifact header")?;
192
193 let mut data_blob = Vec::new();
195 file.read_to_end(&mut data_blob)
196 .context("failed to read artifact data blob")?;
197
198 let data_blob = if flags & FLAG_COMPRESSED != 0 {
200 zstd::decode_all(&data_blob[..]).context("failed to decompress artifact data")?
201 } else {
202 data_blob
203 };
204
205 let mut tensors = HashMap::new();
207 for at in &header.tensors {
208 let start = at.offset as usize;
209 let end = start + at.length as usize;
210 if end > data_blob.len() {
211 anyhow::bail!(
212 "artifact corrupted: tensor {} claims offset {} + length {} exceeds blob size {}",
213 at.name,
214 at.offset,
215 at.length,
216 data_blob.len()
217 );
218 }
219 let data = data_blob[start..end].to_vec();
220 tensors.insert(
221 at.name.clone(),
222 TensorData {
223 shape: at.shape.clone(),
224 dtype: str_to_data_type(&at.dtype)?,
225 data,
226 },
227 );
228 }
229
230 Ok(Model {
231 name: header.name,
232 architecture: header.architecture,
233 tensors,
234 metadata: header.metadata,
235 })
236}
237
238fn data_type_to_str(dt: DataType) -> String {
239 match dt {
240 DataType::F32 => "f32".to_string(),
241 DataType::F16 => "f16".to_string(),
242 DataType::BF16 => "bf16".to_string(),
243 DataType::I64 => "i64".to_string(),
244 DataType::I32 => "i32".to_string(),
245 DataType::I16 => "i16".to_string(),
246 DataType::I8 => "i8".to_string(),
247 DataType::U8 => "u8".to_string(),
248 DataType::Bool => "bool".to_string(),
249 DataType::Q4_0 => "q4_0".to_string(),
250 DataType::Q5_0 => "q5_0".to_string(),
251 DataType::Q8_0 => "q8_0".to_string(),
252 DataType::Q4_K => "q4_k".to_string(),
253 DataType::Q6_K => "q6_k".to_string(),
254 }
255}
256
257fn str_to_data_type(s: &str) -> Result<DataType> {
258 match s {
259 "f32" => Ok(DataType::F32),
260 "f16" => Ok(DataType::F16),
261 "bf16" => Ok(DataType::BF16),
262 "i64" => Ok(DataType::I64),
263 "i32" => Ok(DataType::I32),
264 "i16" => Ok(DataType::I16),
265 "i8" => Ok(DataType::I8),
266 "u8" => Ok(DataType::U8),
267 "bool" => Ok(DataType::Bool),
268 "q4_0" => Ok(DataType::Q4_0),
269 "q5_0" => Ok(DataType::Q5_0),
270 "q8_0" => Ok(DataType::Q8_0),
271 "q4_k" => Ok(DataType::Q4_K),
272 "q6_k" => Ok(DataType::Q6_K),
273 _ => anyhow::bail!("unknown dtype '{}' in artifact", s),
274 }
275}
276
277pub fn verify(path: &Path) -> Result<()> {
279 let mut file = std::fs::File::open(path)
280 .with_context(|| format!("failed to open artifact at {:?}", path))?;
281
282 let mut magic = [0u8; 6];
283 file.read_exact(&mut magic)
284 .context("artifact file too short (magic)")?;
285 if magic != MAGIC {
286 anyhow::bail!("invalid artifact magic bytes");
287 }
288
289 let mut version_bytes = [0u8; 4];
290 file.read_exact(&mut version_bytes)
291 .context("artifact file too short (version)")?;
292 let version = u32::from_le_bytes(version_bytes);
293
294 let flags = if version >= 2 {
295 let mut flags_bytes = [0u8; 4];
296 file.read_exact(&mut flags_bytes)
297 .context("artifact file too short (flags)")?;
298 u32::from_le_bytes(flags_bytes)
299 } else {
300 0
301 };
302
303 if version != 1 && version != 2 {
304 anyhow::bail!("unsupported artifact version {}", version);
305 }
306
307 let mut header_len_bytes = [0u8; 8];
308 file.read_exact(&mut header_len_bytes)
309 .context("artifact file too short (header length)")?;
310 let header_len = u64::from_le_bytes(header_len_bytes);
311
312 let mut header_json = vec![0u8; header_len as usize];
313 file.read_exact(&mut header_json)
314 .context("artifact file too short (header)")?;
315 let header: ArtifactHeader =
316 serde_json::from_slice(&header_json).context("failed to deserialize artifact header")?;
317
318 let mut data_blob = Vec::new();
319 file.read_to_end(&mut data_blob)
320 .context("failed to read artifact data blob")?;
321
322 let data_blob = if flags & FLAG_COMPRESSED != 0 {
323 zstd::decode_all(&data_blob[..]).context("failed to decompress artifact data")?
324 } else {
325 data_blob
326 };
327
328 let mut total_claimed: u64 = 0;
329 let mut max_end: u64 = 0;
330 for at in &header.tensors {
331 let end = at.offset + at.length;
332 if end > max_end {
333 max_end = end;
334 }
335 total_claimed += at.length;
336 }
337
338 if max_end > data_blob.len() as u64 {
339 anyhow::bail!(
340 "artifact corrupted: tensor data claims {} bytes but blob is {} bytes",
341 max_end,
342 data_blob.len()
343 );
344 }
345
346 println!("Verify OK: {}", path.display());
347 println!(" version: {}", version);
348 println!(" compressed: {}", flags & FLAG_COMPRESSED != 0);
349 println!(" name: {}", header.name);
350 println!(" architecture: {}", header.architecture);
351 println!(" tensors: {}", header.tensors.len());
352 println!(" total tensor bytes: {}", total_claimed);
353 println!(" blob size: {} bytes", data_blob.len());
354
355 Ok(())
356}
357
358pub fn export_to_safetensors(path: &Path, output: &Path) -> Result<()> {
360 let model = unpack(path)?;
361
362 let mut metadata = std::collections::HashMap::new();
363 for (k, v) in &model.metadata {
364 metadata.insert(k.clone(), v.clone());
365 }
366 if !model.architecture.is_empty() {
367 metadata.insert("architecture".to_string(), model.architecture.clone());
368 }
369
370 let mut tensors: Vec<(String, SafetensorsView)> = Vec::new();
371 for (name, td) in &model.tensors {
372 tensors.push((
373 name.clone(),
374 SafetensorsView {
375 data: td.data.clone(),
376 dtype: data_type_to_safetensors_dtype(td.dtype),
377 shape: td.shape.clone(),
378 },
379 ));
380 }
381
382 safetensors::serialize_to_file(tensors, &Some(metadata), output)
383 .with_context(|| format!("failed to write safetensors to {:?}", output))?;
384
385 eprintln!("Exported {} tensors -> {:?}", model.tensors.len(), output);
386 Ok(())
387}
388
389fn data_type_to_safetensors_dtype(dt: DataType) -> safetensors::Dtype {
390 match dt {
391 DataType::F32 => safetensors::Dtype::F32,
392 DataType::F16 => safetensors::Dtype::F16,
393 DataType::BF16 => safetensors::Dtype::BF16,
394 DataType::I64 => safetensors::Dtype::I64,
395 DataType::I32 => safetensors::Dtype::I32,
396 DataType::I16 => safetensors::Dtype::I16,
397 DataType::I8 => safetensors::Dtype::I8,
398 DataType::U8 => safetensors::Dtype::U8,
399 DataType::Bool => safetensors::Dtype::BOOL,
400 DataType::Q4_0 | DataType::Q5_0 | DataType::Q8_0 | DataType::Q4_K | DataType::Q6_K => {
401 panic!("cannot export GGUF-quantized tensor to safetensors; dequantize first")
402 }
403 }
404}
405
406struct SafetensorsView {
407 data: Vec<u8>,
408 dtype: safetensors::Dtype,
409 shape: Vec<usize>,
410}
411
412impl safetensors::View for SafetensorsView {
413 fn dtype(&self) -> safetensors::Dtype {
414 self.dtype
415 }
416
417 fn shape(&self) -> &[usize] {
418 &self.shape
419 }
420
421 fn data(&self) -> std::borrow::Cow<'_, [u8]> {
422 std::borrow::Cow::Borrowed(&self.data)
423 }
424
425 fn data_len(&self) -> usize {
426 self.data.len()
427 }
428}