1use super::{
8 AprV2Flags, AprV2Header, TensorDType, V2FormatError, HEADER_SIZE_V2, MAGIC_V2, VERSION_V2,
9};
10use crate::crc32::crc32;
11use serde::{Deserialize, Serialize};
12use std::collections::HashMap;
13
14impl AprV2Header {
15 #[must_use]
17 pub fn new() -> Self {
18 Self {
19 magic: MAGIC_V2,
20 version: VERSION_V2,
21 flags: AprV2Flags::new(),
22 tensor_count: 0,
23 metadata_offset: HEADER_SIZE_V2 as u64,
24 metadata_size: 0,
25 tensor_index_offset: 0,
26 data_offset: 0,
27 checksum: 0,
28 reserved: [0u8; 20],
29 }
30 }
31
32 #[must_use]
34 pub fn is_valid(&self) -> bool {
35 self.magic == MAGIC_V2
36 }
37
38 #[must_use]
40 pub fn to_bytes(&self) -> [u8; HEADER_SIZE_V2] {
41 let mut buf = [0u8; HEADER_SIZE_V2];
42
43 buf[0..4].copy_from_slice(&self.magic);
44 buf[4] = self.version.0;
45 buf[5] = self.version.1;
46 buf[6..8].copy_from_slice(&self.flags.bits().to_le_bytes());
47 buf[8..12].copy_from_slice(&self.tensor_count.to_le_bytes());
48 buf[12..20].copy_from_slice(&self.metadata_offset.to_le_bytes());
49 buf[20..24].copy_from_slice(&self.metadata_size.to_le_bytes());
50 buf[24..32].copy_from_slice(&self.tensor_index_offset.to_le_bytes());
51 buf[32..40].copy_from_slice(&self.data_offset.to_le_bytes());
52 buf[40..44].copy_from_slice(&self.checksum.to_le_bytes());
53 buf[44..64].copy_from_slice(&self.reserved);
54
55 buf
56 }
57
58 pub fn from_bytes(buf: &[u8]) -> Result<Self, V2FormatError> {
63 if buf.len() < HEADER_SIZE_V2 {
64 return Err(V2FormatError::InvalidHeader("buffer too small".to_string()));
65 }
66
67 let magic: [u8; 4] = buf[0..4]
68 .try_into()
69 .map_err(|_| V2FormatError::InvalidHeader("failed to read magic".to_string()))?;
70
71 if magic != MAGIC_V2 {
73 return Err(V2FormatError::InvalidMagic(magic));
74 }
75
76 let version = (buf[4], buf[5]);
77 let flags = AprV2Flags::from_bits(u16::from_le_bytes([buf[6], buf[7]]));
78 let tensor_count = u32::from_le_bytes([buf[8], buf[9], buf[10], buf[11]]);
79 let metadata_offset = u64::from_le_bytes(buf[12..20].try_into().unwrap_or([0; 8]));
80 let metadata_size = u32::from_le_bytes([buf[20], buf[21], buf[22], buf[23]]);
81 let tensor_index_offset = u64::from_le_bytes(buf[24..32].try_into().unwrap_or([0; 8]));
82 let data_offset = u64::from_le_bytes(buf[32..40].try_into().unwrap_or([0; 8]));
83 let checksum = u32::from_le_bytes([buf[40], buf[41], buf[42], buf[43]]);
84
85 let mut reserved = [0u8; 20];
86 reserved.copy_from_slice(buf.get(44..64).unwrap_or(&[0u8; 20]));
87
88 Ok(Self {
89 magic,
90 version,
91 flags,
92 tensor_count,
93 metadata_offset,
94 metadata_size,
95 tensor_index_offset,
96 data_offset,
97 checksum,
98 reserved,
99 })
100 }
101
102 #[must_use]
104 pub fn compute_checksum(&self) -> u32 {
105 let bytes = self.to_bytes();
106 let mut data = Vec::with_capacity(60);
109 data.extend_from_slice(bytes.get(0..40).unwrap_or(&[]));
110 data.extend_from_slice(bytes.get(44..64).unwrap_or(&[]));
111 crc32(&data)
112 }
113
114 pub fn update_checksum(&mut self) {
116 self.checksum = self.compute_checksum();
117 }
118
119 #[must_use]
121 pub fn verify_checksum(&self) -> bool {
122 self.checksum == self.compute_checksum()
123 }
124}
125
126#[derive(Debug, Clone, Default, Serialize, Deserialize)]
132pub struct AprV2Metadata {
133 #[serde(default)]
135 pub model_type: String,
136
137 #[serde(default, skip_serializing_if = "Option::is_none")]
139 pub name: Option<String>,
140
141 #[serde(default, skip_serializing_if = "Option::is_none")]
143 pub description: Option<String>,
144
145 #[serde(default, skip_serializing_if = "Option::is_none")]
147 pub author: Option<String>,
148
149 #[serde(default)]
155 pub license: Option<String>,
156
157 #[serde(default)]
161 pub data_source: Option<String>,
162
163 #[serde(default)]
167 pub data_license: Option<String>,
168
169 #[serde(default, skip_serializing_if = "Option::is_none")]
171 pub version: Option<String>,
172
173 #[serde(default, skip_serializing_if = "Option::is_none")]
176 pub source: Option<String>,
177
178 #[serde(default, skip_serializing_if = "Option::is_none")]
181 pub original_format: Option<String>,
182
183 #[serde(default, skip_serializing_if = "Option::is_none")]
185 pub created_at: Option<String>,
186
187 #[serde(default)]
189 pub total_size: u64,
190
191 #[serde(default)]
193 pub param_count: u64,
194
195 #[serde(default, skip_serializing_if = "Option::is_none")]
197 pub quantization: Option<QuantizationMetadata>,
198
199 #[serde(default, skip_serializing_if = "Option::is_none")]
201 pub sharding: Option<ShardingMetadata>,
202
203 #[serde(default, skip_serializing_if = "Option::is_none")]
206 pub chat_template: Option<String>,
207
208 #[serde(default, skip_serializing_if = "Option::is_none")]
212 pub chat_format: Option<String>,
213
214 #[serde(default, skip_serializing_if = "Option::is_none")]
217 pub special_tokens: Option<ChatSpecialTokens>,
218
219 #[serde(default, skip_serializing_if = "Option::is_none")]
224 pub architecture: Option<String>,
225
226 #[serde(default, skip_serializing_if = "Option::is_none")]
232 pub hf_architecture: Option<String>,
233
234 #[serde(default, skip_serializing_if = "Option::is_none")]
238 pub hf_model_type: Option<String>,
239
240 #[serde(default, skip_serializing_if = "Option::is_none")]
242 pub hidden_size: Option<usize>,
243
244 #[serde(default, skip_serializing_if = "Option::is_none")]
246 pub num_layers: Option<usize>,
247
248 #[serde(default, skip_serializing_if = "Option::is_none")]
250 pub num_heads: Option<usize>,
251
252 #[serde(default, skip_serializing_if = "Option::is_none")]
254 pub num_kv_heads: Option<usize>,
255
256 #[serde(default, skip_serializing_if = "Option::is_none")]
258 pub vocab_size: Option<usize>,
259
260 #[serde(default, skip_serializing_if = "Option::is_none")]
262 pub intermediate_size: Option<usize>,
263
264 #[serde(default, skip_serializing_if = "Option::is_none")]
266 pub max_position_embeddings: Option<usize>,
267
268 #[serde(default, skip_serializing_if = "Option::is_none")]
270 pub rope_theta: Option<f32>,
271
272 #[serde(default, skip_serializing_if = "Option::is_none")]
275 pub rope_type: Option<u32>,
276
277 #[serde(default, skip_serializing_if = "Option::is_none")]
279 pub rms_norm_eps: Option<f32>,
280
281 #[serde(default, skip_serializing_if = "Option::is_none")]
283 pub head_dim: Option<usize>,
284
285 #[serde(default, skip_serializing_if = "Option::is_none")]
287 pub num_experts: Option<usize>,
288
289 #[serde(default, skip_serializing_if = "Option::is_none")]
291 pub num_experts_per_tok: Option<usize>,
292
293 #[serde(default, skip_serializing_if = "Option::is_none")]
295 pub moe_intermediate_size: Option<usize>,
296
297 #[serde(default, flatten)]
299 pub custom: HashMap<String, serde_json::Value>,
300}
301
302#[derive(Debug, Clone, Default, Serialize, Deserialize)]
304pub struct ChatSpecialTokens {
305 #[serde(default)]
307 pub bos_token: Option<String>,
308
309 #[serde(default)]
311 pub eos_token: Option<String>,
312
313 #[serde(default)]
315 pub unk_token: Option<String>,
316
317 #[serde(default)]
319 pub pad_token: Option<String>,
320
321 #[serde(default)]
323 pub im_start_token: Option<String>,
324
325 #[serde(default)]
327 pub im_end_token: Option<String>,
328}
329
330impl AprV2Metadata {
331 #[must_use]
333 pub fn new(model_type: impl Into<String>) -> Self {
334 Self {
335 model_type: model_type.into(),
336 ..Default::default()
337 }
338 }
339
340 pub fn to_json(&self) -> Result<Vec<u8>, V2FormatError> {
345 serde_json::to_vec(self).map_err(|e| V2FormatError::MetadataError(e.to_string()))
346 }
347
348 pub fn to_json_pretty(&self) -> Result<String, V2FormatError> {
353 serde_json::to_string_pretty(self).map_err(|e| V2FormatError::MetadataError(e.to_string()))
354 }
355
356 pub fn canonicalize_hf_aliases(&mut self) {
374 fn take_usize(
375 custom: &mut HashMap<String, serde_json::Value>,
376 keys: &[&str],
377 ) -> Option<usize> {
378 let mut found = None;
379 for k in keys {
380 if let Some(v) = custom.remove(*k) {
381 if found.is_none() {
382 found = v.as_u64().and_then(|n| usize::try_from(n).ok());
383 }
384 }
385 }
386 found
387 }
388 fn take_f32(custom: &mut HashMap<String, serde_json::Value>, keys: &[&str]) -> Option<f32> {
389 let mut found = None;
390 for k in keys {
391 if let Some(v) = custom.remove(*k) {
392 if found.is_none() {
393 #[allow(clippy::cast_possible_truncation)]
394 {
395 found = v.as_f64().map(|n| n as f32);
396 }
397 }
398 }
399 }
400 found
401 }
402
403 let v = take_usize(&mut self.custom, &["hidden_dim", "d_model", "n_embd"]);
405 self.hidden_size = self.hidden_size.or(v);
406 let v = take_usize(
407 &mut self.custom,
408 &["num_hidden_layers", "n_layers", "n_layer"],
409 );
410 self.num_layers = self.num_layers.or(v);
411 let v = take_usize(
412 &mut self.custom,
413 &["num_attention_heads", "n_heads", "n_head"],
414 );
415 self.num_heads = self.num_heads.or(v);
416 let v = take_usize(&mut self.custom, &["num_key_value_heads", "n_kv_heads"]);
417 self.num_kv_heads = self.num_kv_heads.or(v);
418 let v = take_usize(&mut self.custom, &["n_vocab"]);
419 self.vocab_size = self.vocab_size.or(v);
420 let v = take_usize(
421 &mut self.custom,
422 &["ffn_dim", "intermediate_dim", "n_inner"],
423 );
424 self.intermediate_size = self.intermediate_size.or(v);
425 let v = take_usize(
426 &mut self.custom,
427 &["max_seq_len", "context_length", "n_ctx"],
428 );
429 self.max_position_embeddings = self.max_position_embeddings.or(v);
430 let v = take_f32(&mut self.custom, &["layer_norm_eps", "norm_eps"]);
431 self.rms_norm_eps = self.rms_norm_eps.or(v);
432 }
433
434 pub fn from_json(data: &[u8]) -> Result<Self, V2FormatError> {
439 let value: serde_json::Value = serde_json::from_slice(data)
444 .map_err(|e| V2FormatError::MetadataError(e.to_string()))?;
445 serde_json::from_value(value).map_err(|e| V2FormatError::MetadataError(e.to_string()))
446 }
447}
448
449#[derive(Debug, Clone, Default, Serialize, Deserialize)]
451pub struct QuantizationMetadata {
452 pub quant_type: String,
454 pub bits: u8,
456 pub block_size: Option<usize>,
458 pub symmetric: bool,
460}
461
462#[derive(Debug, Clone, Default, Serialize, Deserialize)]
464pub struct ShardingMetadata {
465 pub shard_count: usize,
467 pub shard_index: usize,
469 pub total_size: u64,
471 pub pattern: Option<String>,
473}
474
475#[derive(Debug, Clone)]
481pub struct TensorIndexEntry {
482 pub name: String,
484 pub dtype: TensorDType,
486 pub shape: Vec<usize>,
488 pub offset: u64,
490 pub size: u64,
492}