use super::base::{
BurnpackError, BurnpackHeader, BurnpackMetadata, FORMAT_VERSION, HEADER_SIZE, MAGIC_NUMBER,
TensorDescriptor,
};
use crate::TensorSnapshot;
use alloc::collections::BTreeMap;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec;
use alloc::vec::Vec;
use burn_tensor::Bytes;
#[cfg(feature = "std")]
use std::fs::File;
#[cfg(feature = "std")]
use std::io::Write;
#[cfg(feature = "std")]
use std::path::Path;
pub struct BurnpackWriter {
pub(crate) snapshots: Vec<TensorSnapshot>,
pub(crate) metadata: BTreeMap<String, String>,
}
impl BurnpackWriter {
pub fn new(snapshots: Vec<TensorSnapshot>) -> Self {
Self {
snapshots,
metadata: BTreeMap::new(),
}
}
pub fn with_metadata(mut self, key: &str, value: &str) -> Self {
self.metadata.insert(key.to_string(), value.to_string());
self
}
fn build_metadata(&self) -> Result<(BurnpackMetadata, Vec<u8>), BurnpackError> {
let mut tensors = BTreeMap::new();
let mut current_offset = 0u64;
for snapshot in &self.snapshots {
let data_len = snapshot.data_len() as u64;
let start = current_offset;
let end = start.checked_add(data_len).ok_or_else(|| {
BurnpackError::IoError(format!(
"Tensor offset overflow: {} + {} exceeds maximum",
start, data_len
))
})?;
tensors.insert(
snapshot.full_path(),
TensorDescriptor {
dtype: snapshot.dtype,
shape: snapshot.shape.iter().map(|&s| s as u64).collect(),
data_offsets: (start, end),
param_id: snapshot.tensor_id.map(|id| id.val()),
},
);
current_offset = end;
}
let metadata = BurnpackMetadata {
tensors,
metadata: self.metadata.clone(),
};
let mut metadata_bytes = Vec::new();
ciborium::ser::into_writer(&metadata, &mut metadata_bytes)
.map_err(|e| BurnpackError::IoError(e.to_string()))?;
Ok((metadata, metadata_bytes))
}
pub fn size(&self) -> Result<usize, BurnpackError> {
let (_, metadata_bytes) = self.build_metadata()?;
let data_size = self.snapshots.iter().map(|s| s.data_len()).sum::<usize>();
Ok(HEADER_SIZE + metadata_bytes.len() + data_size)
}
pub fn write_into(&self, buffer: &mut [u8]) -> Result<(), BurnpackError> {
let (_, metadata_bytes) = self.build_metadata()?;
let metadata_size: u32 = metadata_bytes.len().try_into().map_err(|_| {
BurnpackError::IoError(format!(
"Metadata size {} exceeds maximum of {} bytes",
metadata_bytes.len(),
u32::MAX
))
})?;
let header = BurnpackHeader {
magic: MAGIC_NUMBER,
version: FORMAT_VERSION,
metadata_size,
};
let data_size = self.snapshots.iter().map(|s| s.data_len()).sum::<usize>();
let total_size = HEADER_SIZE + metadata_bytes.len() + data_size;
if buffer.len() < total_size {
return Err(BurnpackError::IoError(format!(
"Buffer too small: need {} bytes, got {} bytes",
total_size,
buffer.len()
)));
}
let mut offset = 0;
let header_bytes = header.into_bytes();
buffer[offset..offset + HEADER_SIZE].copy_from_slice(&header_bytes);
offset += HEADER_SIZE;
buffer[offset..offset + metadata_bytes.len()].copy_from_slice(&metadata_bytes);
offset += metadata_bytes.len();
for snapshot in &self.snapshots {
let expected_len = snapshot.data_len();
let data = snapshot.to_data().map_err(|e| {
BurnpackError::IoError(format!("Failed to get tensor data: {:?}", e))
})?;
let actual_len = data.bytes.len();
if actual_len != expected_len {
return Err(BurnpackError::IoError(format!(
"Data corruption: tensor '{}' has inconsistent length (expected {}, got {})",
snapshot.full_path(),
expected_len,
actual_len
)));
}
buffer[offset..offset + actual_len].copy_from_slice(&data.bytes);
offset += actual_len;
}
Ok(())
}
pub fn to_bytes(&self) -> Result<Bytes, BurnpackError> {
let size = self.size()?;
let mut buffer = vec![0u8; size];
self.write_into(&mut buffer)?;
Ok(Bytes::from_bytes_vec(buffer))
}
#[cfg(feature = "std")]
pub fn write_to_file<P: AsRef<Path>>(&self, path: P) -> Result<(), BurnpackError> {
let mut file = File::create(path).map_err(|e| BurnpackError::IoError(e.to_string()))?;
let (_, metadata_bytes) = self.build_metadata()?;
let metadata_size: u32 = metadata_bytes.len().try_into().map_err(|_| {
BurnpackError::IoError(format!(
"Metadata size {} exceeds maximum of {} bytes",
metadata_bytes.len(),
u32::MAX
))
})?;
let header = BurnpackHeader {
magic: MAGIC_NUMBER,
version: FORMAT_VERSION,
metadata_size,
};
file.write_all(&header.into_bytes())
.map_err(|e| BurnpackError::IoError(e.to_string()))?;
file.write_all(&metadata_bytes)
.map_err(|e| BurnpackError::IoError(e.to_string()))?;
for snapshot in &self.snapshots {
let expected_len = snapshot.data_len();
let data = snapshot.to_data().map_err(|e| {
BurnpackError::IoError(format!("Failed to get tensor data: {:?}", e))
})?;
let actual_len = data.bytes.len();
if actual_len != expected_len {
return Err(BurnpackError::IoError(format!(
"Data corruption: tensor '{}' has inconsistent length (expected {}, got {})",
snapshot.full_path(),
expected_len,
actual_len
)));
}
file.write_all(&data.bytes)
.map_err(|e| BurnpackError::IoError(e.to_string()))?;
}
file.flush()
.map_err(|e| BurnpackError::IoError(e.to_string()))?;
Ok(())
}
}