use super::base::{
Error, FORMAT_VERSION, HEADER_SIZE, Header, MAGIC_NUMBER, Metadata, Scalar, TENSOR_ALIGNMENT,
TensorDescriptor, aligned_data_section_start,
};
use super::tensor::Tensor;
use alloc::collections::BTreeMap;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec;
use alloc::vec::Vec;
use burn_std::Bytes;
#[cfg(feature = "std")]
use std::fs::File;
#[cfg(feature = "std")]
use std::io::{Read, Write};
#[cfg(feature = "std")]
use std::path::Path;
#[inline]
const fn align_offset(offset: u64, alignment: u64) -> u64 {
offset.div_ceil(alignment) * alignment
}
const WRITE_CHUNK_SIZE: usize = 8 * 1024 * 1024;
pub struct Writer {
pub(crate) tensors: Vec<Tensor>,
pub(crate) metadata: BTreeMap<String, String>,
pub(crate) scalars: BTreeMap<String, Scalar>,
}
impl Writer {
pub fn new(tensors: Vec<Tensor>) -> Self {
Self {
tensors,
metadata: BTreeMap::new(),
scalars: BTreeMap::new(),
}
}
pub fn with_metadata(mut self, key: &str, value: &str) -> Self {
self.metadata.insert(key.to_string(), value.to_string());
self
}
pub fn with_scalar(mut self, key: &str, value: Scalar) -> Self {
self.scalars.insert(key.to_string(), value);
self
}
pub fn size(&self) -> Result<usize, Error> {
Ok(self.plan()?.total_size())
}
pub fn write_into(self, buffer: &mut [u8]) -> Result<(), Error> {
let layout = self.plan()?;
let total_size = layout.total_size();
if buffer.len() < total_size {
return Err(Error::IoError(format!(
"Buffer too small: need {} bytes, got {} bytes",
total_size,
buffer.len()
)));
}
let mut sink = BufferSink { buffer, offset: 0 };
self.write_container(&layout, &mut sink)
}
pub fn into_bytes(self) -> Result<Bytes, Error> {
let layout = self.plan()?;
let mut buffer = vec![0u8; layout.total_size()];
let mut sink = BufferSink {
buffer: &mut buffer,
offset: 0,
};
self.write_container(&layout, &mut sink)?;
Ok(Bytes::from_bytes_vec(buffer))
}
#[cfg(feature = "std")]
pub fn write_to_file<P: AsRef<Path>>(self, path: P) -> Result<(), Error> {
let path = path.as_ref();
let path = if path.extension().is_none() {
path.with_extension(crate::EXTENSION)
} else {
path.to_path_buf()
};
let layout = self.plan()?;
let file = File::create(path).map_err(|e| Error::IoError(e.to_string()))?;
let mut sink = FileSink { file };
self.write_container(&layout, &mut sink)?;
sink.file.flush().map_err(|e| Error::IoError(e.to_string()))
}
fn plan(&self) -> Result<Layout, Error> {
let (metadata, metadata_bytes, data_size) = self.build_metadata()?;
let metadata_size: u32 = metadata_bytes.len().try_into().map_err(|_| {
Error::IoError(format!(
"Metadata size {} exceeds maximum of {} bytes",
metadata_bytes.len(),
u32::MAX
))
})?;
let header = Header {
magic: MAGIC_NUMBER,
version: FORMAT_VERSION,
metadata_size,
};
let data_section_start = aligned_data_section_start(metadata_bytes.len());
Ok(Layout {
metadata,
metadata_bytes,
header,
data_section_start,
data_size,
})
}
fn build_metadata(&self) -> Result<(Metadata, Vec<u8>, usize), Error> {
let (tensors, data_size) = self.build_descriptors()?;
let metadata = Metadata {
tensors,
metadata: self.metadata.clone(),
scalars: self.scalars.clone(),
};
let mut metadata_bytes = Vec::new();
ciborium::ser::into_writer(&metadata, &mut metadata_bytes)
.map_err(|e| Error::MetadataSerializationError(e.to_string()))?;
Ok((metadata, metadata_bytes, data_size))
}
fn build_descriptors(&self) -> Result<(BTreeMap<String, TensorDescriptor>, usize), Error> {
let mut tensors = BTreeMap::new();
let mut current_offset = 0u64;
for tensor in &self.tensors {
let data_len = tensor.bytes.len() as u64;
let aligned_start = align_offset(current_offset, TENSOR_ALIGNMENT);
let end = aligned_start.checked_add(data_len).ok_or_else(|| {
Error::IoError(format!(
"Tensor offset overflow: {} + {} exceeds maximum",
aligned_start, data_len
))
})?;
if tensors
.insert(
tensor.name.clone(),
TensorDescriptor {
dtype: tensor.dtype,
shape: tensor.shape.iter().map(|&s| s as u64).collect(),
data_offsets: (aligned_start, end),
param_id: tensor.param_id,
},
)
.is_some()
{
return Err(Error::ValidationError(format!(
"Duplicate tensor name '{}'",
tensor.name
)));
}
current_offset = end;
}
Ok((tensors, current_offset as usize))
}
fn write_container(self, layout: &Layout, sink: &mut impl Sink) -> Result<(), Error> {
sink.write(&layout.header.into_bytes())?;
sink.write(&layout.metadata_bytes)?;
let unaligned_data_start = HEADER_SIZE + layout.metadata_bytes.len();
if layout.data_section_start > unaligned_data_start {
sink.pad(layout.data_section_start - unaligned_data_start)?;
}
self.write_tensors(&layout.metadata, sink)
}
fn write_tensors(self, metadata: &Metadata, sink: &mut impl Sink) -> Result<(), Error> {
let mut data_offset = 0usize;
for tensor in self.tensors.into_iter() {
let (aligned_offset, data) = Self::resolve_tensor(tensor, metadata)?;
if aligned_offset > data_offset {
sink.pad(aligned_offset - data_offset)?;
data_offset = aligned_offset;
}
Self::write_tensor_data(&data, sink)?;
data_offset += data.len();
}
Ok(())
}
fn write_tensor_data(data: &Bytes, sink: &mut impl Sink) -> Result<(), Error> {
let len = data.len();
let mut offset = 0;
while offset < len {
let end = (offset + WRITE_CHUNK_SIZE).min(len);
match data.view(offset, end) {
Ok(chunk) => {
sink.write(&chunk)?;
offset = end;
}
Err(_) => {
sink.write(&data[offset..])?;
break;
}
}
}
Ok(())
}
fn resolve_tensor(tensor: Tensor, metadata: &Metadata) -> Result<(usize, Bytes), Error> {
let descriptor = metadata.tensors.get(&tensor.name).ok_or_else(|| {
Error::IoError(format!(
"Internal error: tensor '{}' not found in metadata",
tensor.name
))
})?;
let (start, end) = descriptor.data_offsets;
let declared_len = (end - start) as usize;
let actual_len = tensor.bytes.len();
if actual_len != declared_len {
return Err(Error::TensorBytesSizeMismatch(format!(
"tensor '{}' has inconsistent length (expected {}, got {})",
tensor.name, declared_len, actual_len
)));
}
Ok((start as usize, tensor.bytes))
}
}
struct Layout {
metadata: Metadata,
metadata_bytes: Vec<u8>,
header: Header,
data_section_start: usize,
data_size: usize,
}
impl Layout {
fn total_size(&self) -> usize {
self.data_section_start + self.data_size
}
}
trait Sink {
fn pad(&mut self, count: usize) -> Result<(), Error>;
fn write(&mut self, data: &[u8]) -> Result<(), Error>;
}
struct BufferSink<'a> {
buffer: &'a mut [u8],
offset: usize,
}
impl Sink for BufferSink<'_> {
fn pad(&mut self, count: usize) -> Result<(), Error> {
self.buffer[self.offset..self.offset + count].fill(0);
self.offset += count;
Ok(())
}
fn write(&mut self, data: &[u8]) -> Result<(), Error> {
self.buffer[self.offset..self.offset + data.len()].copy_from_slice(data);
self.offset += data.len();
Ok(())
}
}
#[cfg(feature = "std")]
struct FileSink {
file: File,
}
#[cfg(feature = "std")]
impl Sink for FileSink {
fn pad(&mut self, count: usize) -> Result<(), Error> {
std::io::copy(&mut std::io::repeat(0).take(count as u64), &mut self.file)
.map(|_| ())
.map_err(|e| Error::IoError(e.to_string()))
}
fn write(&mut self, data: &[u8]) -> Result<(), Error> {
self.file
.write_all(data)
.map_err(|e| Error::IoError(e.to_string()))
}
}