1use super::base::{
2 Error, FORMAT_VERSION, HEADER_SIZE, Header, MAGIC_NUMBER, Metadata, Scalar, TENSOR_ALIGNMENT,
3 TensorDescriptor, aligned_data_section_start,
4};
5use super::tensor::Tensor;
6use alloc::collections::BTreeMap;
7use alloc::format;
8use alloc::string::{String, ToString};
9use alloc::vec;
10use alloc::vec::Vec;
11use burn_std::Bytes;
12
13#[cfg(feature = "std")]
14use std::fs::File;
15#[cfg(feature = "std")]
16use std::io::{Read, Write};
17#[cfg(feature = "std")]
18use std::path::Path;
19
20#[inline]
24const fn align_offset(offset: u64, alignment: u64) -> u64 {
25 offset.div_ceil(alignment) * alignment
26}
27
28const WRITE_CHUNK_SIZE: usize = 8 * 1024 * 1024;
37
38pub struct Writer {
40 pub(crate) tensors: Vec<Tensor>,
42 pub(crate) metadata: BTreeMap<String, String>,
44 pub(crate) scalars: BTreeMap<String, Scalar>,
46}
47
48impl Writer {
49 pub fn new(tensors: Vec<Tensor>) -> Self {
51 Self {
52 tensors,
53 metadata: BTreeMap::new(),
54 scalars: BTreeMap::new(),
55 }
56 }
57
58 pub fn with_metadata(mut self, key: &str, value: &str) -> Self {
60 self.metadata.insert(key.to_string(), value.to_string());
61 self
62 }
63
64 pub fn with_scalar(mut self, key: &str, value: Scalar) -> Self {
66 self.scalars.insert(key.to_string(), value);
67 self
68 }
69
70 pub fn size(&self) -> Result<usize, Error> {
75 Ok(self.plan()?.total_size())
76 }
77
78 pub fn write_into(self, buffer: &mut [u8]) -> Result<(), Error> {
92 let layout = self.plan()?;
93 let total_size = layout.total_size();
94
95 if buffer.len() < total_size {
96 return Err(Error::IoError(format!(
97 "Buffer too small: need {} bytes, got {} bytes",
98 total_size,
99 buffer.len()
100 )));
101 }
102
103 let mut sink = BufferSink { buffer, offset: 0 };
104 self.write_container(&layout, &mut sink)
105 }
106
107 pub fn into_bytes(self) -> Result<Bytes, Error> {
112 let layout = self.plan()?;
113 let mut buffer = vec![0u8; layout.total_size()];
114
115 let mut sink = BufferSink {
116 buffer: &mut buffer,
117 offset: 0,
118 };
119 self.write_container(&layout, &mut sink)?;
120
121 Ok(Bytes::from_bytes_vec(buffer))
122 }
123
124 #[cfg(feature = "std")]
128 pub fn write_to_file<P: AsRef<Path>>(self, path: P) -> Result<(), Error> {
129 let path = path.as_ref();
130 let path = if path.extension().is_none() {
131 path.with_extension(crate::EXTENSION)
132 } else {
133 path.to_path_buf()
134 };
135
136 let layout = self.plan()?;
137 let file = File::create(path).map_err(|e| Error::IoError(e.to_string()))?;
138
139 let mut sink = FileSink { file };
140 self.write_container(&layout, &mut sink)?;
141
142 sink.file.flush().map_err(|e| Error::IoError(e.to_string()))
143 }
144
145 fn plan(&self) -> Result<Layout, Error> {
148 let (metadata, metadata_bytes, data_size) = self.build_metadata()?;
149
150 let metadata_size: u32 = metadata_bytes.len().try_into().map_err(|_| {
151 Error::IoError(format!(
152 "Metadata size {} exceeds maximum of {} bytes",
153 metadata_bytes.len(),
154 u32::MAX
155 ))
156 })?;
157
158 let header = Header {
159 magic: MAGIC_NUMBER,
160 version: FORMAT_VERSION,
161 metadata_size,
162 };
163
164 let data_section_start = aligned_data_section_start(metadata_bytes.len());
165
166 Ok(Layout {
167 metadata,
168 metadata_bytes,
169 header,
170 data_section_start,
171 data_size,
172 })
173 }
174
175 fn build_metadata(&self) -> Result<(Metadata, Vec<u8>, usize), Error> {
179 let (tensors, data_size) = self.build_descriptors()?;
180 let metadata = Metadata {
181 tensors,
182 metadata: self.metadata.clone(),
183 scalars: self.scalars.clone(),
184 };
185
186 let mut metadata_bytes = Vec::new();
187 ciborium::ser::into_writer(&metadata, &mut metadata_bytes)
188 .map_err(|e| Error::MetadataSerializationError(e.to_string()))?;
189
190 Ok((metadata, metadata_bytes, data_size))
191 }
192
193 fn build_descriptors(&self) -> Result<(BTreeMap<String, TensorDescriptor>, usize), Error> {
199 let mut tensors = BTreeMap::new();
200 let mut current_offset = 0u64;
201
202 for tensor in &self.tensors {
203 let data_len = tensor.bytes.len() as u64;
204
205 let aligned_start = align_offset(current_offset, TENSOR_ALIGNMENT);
207 let end = aligned_start.checked_add(data_len).ok_or_else(|| {
208 Error::IoError(format!(
209 "Tensor offset overflow: {} + {} exceeds maximum",
210 aligned_start, data_len
211 ))
212 })?;
213
214 if tensors
218 .insert(
219 tensor.name.clone(),
220 TensorDescriptor {
221 dtype: tensor.dtype,
222 shape: tensor.shape.iter().map(|&s| s as u64).collect(),
223 data_offsets: (aligned_start, end),
224 param_id: tensor.param_id,
225 },
226 )
227 .is_some()
228 {
229 return Err(Error::ValidationError(format!(
230 "Duplicate tensor name '{}'",
231 tensor.name
232 )));
233 }
234
235 current_offset = end;
236 }
237
238 Ok((tensors, current_offset as usize))
239 }
240
241 fn write_container(self, layout: &Layout, sink: &mut impl Sink) -> Result<(), Error> {
244 sink.write(&layout.header.into_bytes())?;
245 sink.write(&layout.metadata_bytes)?;
246
247 let unaligned_data_start = HEADER_SIZE + layout.metadata_bytes.len();
249 if layout.data_section_start > unaligned_data_start {
250 sink.pad(layout.data_section_start - unaligned_data_start)?;
251 }
252
253 self.write_tensors(&layout.metadata, sink)
254 }
255
256 fn write_tensors(self, metadata: &Metadata, sink: &mut impl Sink) -> Result<(), Error> {
259 let mut data_offset = 0usize;
261
262 for tensor in self.tensors.into_iter() {
263 let (aligned_offset, data) = Self::resolve_tensor(tensor, metadata)?;
264
265 if aligned_offset > data_offset {
266 sink.pad(aligned_offset - data_offset)?;
267 data_offset = aligned_offset;
268 }
269
270 Self::write_tensor_data(&data, sink)?;
271 data_offset += data.len();
272 }
273
274 Ok(())
275 }
276
277 fn write_tensor_data(data: &Bytes, sink: &mut impl Sink) -> Result<(), Error> {
291 let len = data.len();
292 let mut offset = 0;
293
294 while offset < len {
295 let end = (offset + WRITE_CHUNK_SIZE).min(len);
296 match data.view(offset, end) {
297 Ok(chunk) => {
298 sink.write(&chunk)?;
299 offset = end;
300 }
301 Err(_) => {
305 sink.write(&data[offset..])?;
306 break;
307 }
308 }
309 }
310
311 Ok(())
312 }
313
314 fn resolve_tensor(tensor: Tensor, metadata: &Metadata) -> Result<(usize, Bytes), Error> {
317 let descriptor = metadata.tensors.get(&tensor.name).ok_or_else(|| {
318 Error::IoError(format!(
319 "Internal error: tensor '{}' not found in metadata",
320 tensor.name
321 ))
322 })?;
323
324 let (start, end) = descriptor.data_offsets;
325 let declared_len = (end - start) as usize;
326 let actual_len = tensor.bytes.len();
327 if actual_len != declared_len {
328 return Err(Error::TensorBytesSizeMismatch(format!(
329 "tensor '{}' has inconsistent length (expected {}, got {})",
330 tensor.name, declared_len, actual_len
331 )));
332 }
333
334 Ok((start as usize, tensor.bytes))
335 }
336}
337
338struct Layout {
345 metadata: Metadata,
346 metadata_bytes: Vec<u8>,
347 header: Header,
348 data_section_start: usize,
349 data_size: usize,
350}
351
352impl Layout {
353 fn total_size(&self) -> usize {
355 self.data_section_start + self.data_size
356 }
357}
358
359trait Sink {
365 fn pad(&mut self, count: usize) -> Result<(), Error>;
367 fn write(&mut self, data: &[u8]) -> Result<(), Error>;
369}
370
371struct BufferSink<'a> {
373 buffer: &'a mut [u8],
374 offset: usize,
375}
376
377impl Sink for BufferSink<'_> {
378 fn pad(&mut self, count: usize) -> Result<(), Error> {
379 self.buffer[self.offset..self.offset + count].fill(0);
380 self.offset += count;
381 Ok(())
382 }
383
384 fn write(&mut self, data: &[u8]) -> Result<(), Error> {
385 self.buffer[self.offset..self.offset + data.len()].copy_from_slice(data);
386 self.offset += data.len();
387 Ok(())
388 }
389}
390
391#[cfg(feature = "std")]
393struct FileSink {
394 file: File,
395}
396
397#[cfg(feature = "std")]
398impl Sink for FileSink {
399 fn pad(&mut self, count: usize) -> Result<(), Error> {
400 std::io::copy(&mut std::io::repeat(0).take(count as u64), &mut self.file)
402 .map(|_| ())
403 .map_err(|e| Error::IoError(e.to_string()))
404 }
405
406 fn write(&mut self, data: &[u8]) -> Result<(), Error> {
407 self.file
408 .write_all(data)
409 .map_err(|e| Error::IoError(e.to_string()))
410 }
411}