1use super::{
9 align_64, AprV2Flags, AprV2Header, AprV2Metadata, ShardingMetadata, TensorDType,
10 TensorIndexEntry, V2FormatError, HEADER_SIZE_V2,
11};
12use crate::crc32::crc32;
13use crate::f16::f32_to_f16;
14use std::io::Write;
15
16#[derive(Debug)]
18pub struct AprV2Writer {
19 header: AprV2Header,
20 metadata: AprV2Metadata,
21 tensors: Vec<(TensorIndexEntry, Vec<u8>)>,
22}
23
24impl AprV2Writer {
25 #[must_use]
30 pub fn new(metadata: AprV2Metadata) -> Self {
31 let mut header = AprV2Header::new();
32 header.flags = header.flags.with(AprV2Flags::LAYOUT_ROW_MAJOR);
34 Self {
35 header,
36 metadata,
37 tensors: Vec::new(),
38 }
39 }
40
41 pub fn add_tensor(
43 &mut self,
44 name: impl Into<String>,
45 dtype: TensorDType,
46 shape: Vec<usize>,
47 data: Vec<u8>,
48 ) {
49 let entry = TensorIndexEntry::new(name, dtype, shape, 0, data.len() as u64);
50 self.tensors.push((entry, data));
51 }
52
53 pub fn add_f32_tensor(&mut self, name: impl Into<String>, shape: Vec<usize>, data: &[f32]) {
55 let bytes: Vec<u8> = data.iter().flat_map(|f| f.to_le_bytes()).collect();
56 self.add_tensor(name, TensorDType::F32, shape, bytes);
57 }
58
59 pub fn add_tensor_f32_owned(
62 &mut self,
63 name: impl Into<String>,
64 shape: Vec<usize>,
65 data: Vec<f32>,
66 ) {
67 let mut bytes = Vec::with_capacity(data.len() * 4);
68 for &f in &data {
69 bytes.extend_from_slice(&f.to_le_bytes());
70 }
71 drop(data);
72 self.add_tensor(name, TensorDType::F32, shape, bytes);
73 }
74
75 pub fn add_f16_tensor(&mut self, name: impl Into<String>, shape: Vec<usize>, data: &[f32]) {
80 let bytes: Vec<u8> = data
81 .iter()
82 .flat_map(|&f| f32_to_f16(f).to_le_bytes())
83 .collect();
84 self.add_tensor(name, TensorDType::F16, shape, bytes);
85 }
86
87 pub fn add_q8_tensor(&mut self, name: impl Into<String>, shape: Vec<usize>, data: &[f32]) {
93 let name = name.into();
94 if data.is_empty() {
95 self.add_tensor(name, TensorDType::AprQ8, shape, Vec::new());
96 return;
97 }
98
99 let max_abs = data.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
101 let scale = if max_abs == 0.0 { 1.0 } else { max_abs / 127.0 };
102
103 let mut bytes = Vec::with_capacity(4 + data.len());
105 bytes.extend_from_slice(&scale.to_le_bytes());
106
107 for &v in data {
108 let q = (v / scale).round().clamp(-127.0, 127.0) as i8;
109 bytes.push(q as u8);
110 }
111
112 let element_count: usize = shape.iter().product();
114 assert_eq!(
115 bytes.len(),
116 4 + element_count,
117 "Q8 CONTRACT VIOLATION: tensor '{}' packed {} bytes, expected {} (4 + {})",
118 name,
119 bytes.len(),
120 4 + element_count,
121 element_count
122 );
123
124 #[allow(clippy::naive_bytecount)]
130 if element_count >= 1024 {
131 let zero_count = bytes[4..].iter().filter(|&&b| b == 0).count();
132 let zero_pct = zero_count as f64 / element_count as f64;
133 if zero_pct > 0.995 {
134 eprintln!(
135 "[F-DATA-QUALITY-001] WARNING: tensor '{}' Q8 has {:.1}% zeros (global-scale Q8 precision loss)",
136 name,
137 zero_pct * 100.0
138 );
139 }
140 }
141
142 self.add_tensor(name, TensorDType::AprQ8, shape, bytes);
143 }
144
145 pub fn add_q4_tensor(&mut self, name: impl Into<String>, shape: Vec<usize>, data: &[f32]) {
153 const BLOCK_SIZE: usize = 32;
154
155 let name = name.into();
156 if data.is_empty() {
157 self.add_tensor(name, TensorDType::AprQ4, shape, Vec::new());
158 return;
159 }
160
161 let num_blocks = data.len().div_ceil(BLOCK_SIZE);
163 let mut bytes = Vec::with_capacity(num_blocks * 18);
164
165 for block_start in (0..data.len()).step_by(BLOCK_SIZE) {
166 let block_end = (block_start + BLOCK_SIZE).min(data.len());
167 let block = &data[block_start..block_end];
168
169 let max_abs = block.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
171 let scale = if max_abs == 0.0 { 1.0 } else { max_abs / 7.0 };
172
173 bytes.extend_from_slice(&f32_to_f16(scale).to_le_bytes());
175
176 let mut packed_idx = 0;
178 let mut packed_buf = [0u8; 16];
179
180 for (i, &v) in block.iter().enumerate() {
181 let q = (v / scale).round().clamp(-8.0, 7.0) as i8;
183 let nibble = ((q + 8) as u8) & 0x0F;
185
186 if i % 2 == 0 {
187 packed_buf[packed_idx] = nibble;
188 } else {
189 packed_buf[packed_idx] |= nibble << 4;
190 packed_idx += 1;
191 }
192 }
193 bytes.extend_from_slice(&packed_buf);
197 }
198
199 let element_count: usize = shape.iter().product();
201 let expected_blocks = element_count.div_ceil(32);
202 assert_eq!(
203 bytes.len(),
204 expected_blocks * 18,
205 "Q4 CONTRACT VIOLATION: tensor '{}' packed {} bytes, expected {} ({} blocks * 18)",
206 name,
207 bytes.len(),
208 expected_blocks * 18,
209 expected_blocks
210 );
211
212 if element_count >= 1024 {
217 let mut zero_nibbles = 0usize;
218 let mut total_nibbles = 0usize;
219 for block_idx in 0..num_blocks {
220 let block_offset = block_idx * 18 + 2; let block_elem_count =
222 BLOCK_SIZE.min(element_count.saturating_sub(block_idx * BLOCK_SIZE));
223 for i in 0..block_elem_count {
224 let byte = bytes[block_offset + i / 2];
225 let nibble = if i % 2 == 0 {
226 byte & 0x0F
227 } else {
228 (byte >> 4) & 0x0F
229 };
230 if nibble == 8 {
231 zero_nibbles += 1;
232 }
233 total_nibbles += 1;
234 }
235 }
236 if total_nibbles > 0 {
237 let zero_pct = zero_nibbles as f64 / total_nibbles as f64;
238 assert!(
239 zero_pct <= 0.99,
240 "Q4 DENSITY VIOLATION: tensor '{}' has {:.1}% zeros (threshold 99%)",
241 name,
242 zero_pct * 100.0
243 );
244 }
245 }
246
247 self.add_tensor(name, TensorDType::AprQ4, shape, bytes);
248 }
249
250 pub fn add_q4k_raw_tensor(
259 &mut self,
260 name: impl Into<String>,
261 shape: Vec<usize>,
262 raw_data: Vec<u8>,
263 ) {
264 self.add_tensor(name, TensorDType::Q4K, shape, raw_data);
265 }
266
267 pub fn add_q6k_raw_tensor(
274 &mut self,
275 name: impl Into<String>,
276 shape: Vec<usize>,
277 raw_data: Vec<u8>,
278 ) {
279 self.add_tensor(name, TensorDType::Q6K, shape, raw_data);
280 }
281
282 pub fn with_lz4_compression(&mut self) -> &mut Self {
284 self.header.flags = self.header.flags.with(AprV2Flags::LZ4_COMPRESSED);
285 self
286 }
287
288 pub fn set_header_flags(&mut self, flags: AprV2Flags) {
298 self.header.flags = flags.with(AprV2Flags::LAYOUT_ROW_MAJOR);
299 }
300
301 pub fn with_sharding(&mut self, shard_count: usize, shard_index: usize) -> &mut Self {
303 self.header.flags = self.header.flags.with(AprV2Flags::SHARDED);
304 self.metadata.sharding = Some(ShardingMetadata {
305 shard_count,
306 shard_index,
307 total_size: 0,
308 pattern: None,
309 });
310 self
311 }
312
313 pub fn write(&mut self) -> Result<Vec<u8>, V2FormatError> {
318 self.tensors.sort_by(|a, b| a.0.name.cmp(&b.0.name));
320
321 let metadata_bytes = self.metadata.to_json()?;
323 let metadata_padded_size = align_64(metadata_bytes.len());
324
325 let mut tensor_index_bytes = Vec::new();
327 let mut data_offset = 0_u64;
328
329 for (entry, data) in &mut self.tensors {
330 entry.offset = data_offset;
331 entry.size = data.len() as u64;
332 tensor_index_bytes.extend_from_slice(&entry.to_bytes());
333 data_offset += align_64(data.len()) as u64;
334 }
335 let tensor_index_padded_size = align_64(tensor_index_bytes.len());
336
337 let metadata_offset = HEADER_SIZE_V2;
339 let tensor_index_offset = metadata_offset + metadata_padded_size;
340 let data_section_offset = tensor_index_offset + tensor_index_padded_size;
341
342 self.header.tensor_count = self.tensors.len() as u32;
344 self.header.metadata_offset = metadata_offset as u64;
345 self.header.metadata_size = metadata_bytes.len() as u32;
346 self.header.tensor_index_offset = tensor_index_offset as u64;
347 self.header.data_offset = data_section_offset as u64;
348 self.header.update_checksum();
349
350 let total_data_size: usize = self.tensors.iter().map(|(_, d)| align_64(d.len())).sum();
352 let total_size = data_section_offset + total_data_size + 4; let mut output = Vec::with_capacity(total_size);
354
355 output.extend_from_slice(&self.header.to_bytes());
357
358 output.extend_from_slice(&metadata_bytes);
360 output.resize(metadata_offset + metadata_padded_size, 0);
361
362 output.extend_from_slice(&tensor_index_bytes);
364 output.resize(tensor_index_offset + tensor_index_padded_size, 0);
365
366 for (_, data) in &self.tensors {
368 let start = output.len();
369 output.extend_from_slice(data);
370 let padded_size = align_64(data.len());
371 output.resize(start + padded_size, 0);
372 }
373
374 let footer_checksum = crc32(&output);
376 output.extend_from_slice(&footer_checksum.to_le_bytes());
377
378 Ok(output)
379 }
380
381 pub fn write_to<W: Write>(&mut self, writer: &mut W) -> Result<(), V2FormatError> {
386 let bytes = self.write()?;
387 writer
388 .write_all(&bytes)
389 .map_err(|e| V2FormatError::IoError(e.to_string()))
390 }
391
392 pub fn write_into(mut self, path: impl AsRef<std::path::Path>) -> Result<(), V2FormatError> {
397 let mut file =
398 std::fs::File::create(path).map_err(|e| V2FormatError::IoError(e.to_string()))?;
399 self.write_to(&mut file)
400 }
401}
402
403