use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec::Vec;
use hashbrown::HashMap;
use crate::module::{Module, ModuleMapper, ModuleVisitor, Param, ParamId};
use crate::tensor::{Bool, DType, Device, Float, Int, Shape, Tensor, TensorData, kind::Basic};
use burn_pack::{Reader, Writer};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum DTypePolicy {
#[default]
FromRecord,
CastToModule,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RecordError {
Io(String),
Validation(String),
}
impl core::fmt::Display for RecordError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
RecordError::Io(msg) => write!(f, "Record I/O error: {msg}"),
RecordError::Validation(msg) => write!(f, "Record validation error: {msg}"),
}
}
}
#[cfg(feature = "std")]
impl std::error::Error for RecordError {}
impl From<burn_pack::Error> for RecordError {
fn from(err: burn_pack::Error) -> Self {
RecordError::Io(err.to_string())
}
}
#[derive(Clone)]
struct RecordTensor {
path: String,
id: ParamId,
data: TensorData,
}
#[derive(Clone)]
pub struct ModuleRecord {
tensors: Vec<RecordTensor>,
dtype_policy: DTypePolicy,
allow_partial: bool,
validate: bool,
}
impl core::fmt::Debug for ModuleRecord {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("ModuleRecord")
.field("num_tensors", &self.tensors.len())
.field("dtype_policy", &self.dtype_policy)
.field("allow_partial", &self.allow_partial)
.field("validate", &self.validate)
.finish()
}
}
impl ModuleRecord {
fn from_tensors(tensors: Vec<RecordTensor>) -> Self {
Self {
tensors,
dtype_policy: DTypePolicy::default(),
allow_partial: false,
validate: true,
}
}
pub fn len(&self) -> usize {
self.tensors.len()
}
pub fn is_empty(&self) -> bool {
self.tensors.is_empty()
}
pub fn with_dtype_policy(mut self, policy: DTypePolicy) -> Self {
self.dtype_policy = policy;
self
}
pub fn cast_to_module_dtype(self) -> Self {
self.with_dtype_policy(DTypePolicy::CastToModule)
}
pub fn allow_partial(mut self, allow: bool) -> Self {
self.allow_partial = allow;
self
}
pub fn validate(mut self, validate: bool) -> Self {
self.validate = validate;
self
}
pub fn into_bytes(self) -> Result<crate::tensor::Bytes, RecordError> {
Ok(Writer::new(self.pack_tensors()).into_bytes()?)
}
pub fn from_bytes(bytes: crate::tensor::Bytes) -> Result<Self, RecordError> {
Self::from_reader(Reader::from_bytes(bytes)?)
}
#[cfg(feature = "std")]
pub fn save<P: AsRef<std::path::Path>>(self, path: P) -> Result<(), RecordError> {
Writer::new(self.pack_tensors()).write_to_file(path)?;
Ok(())
}
#[cfg(feature = "std")]
pub fn load<P: AsRef<std::path::Path>>(path: P) -> Result<Self, RecordError> {
Self::from_reader(Reader::from_file(path)?)
}
fn pack_tensors(self) -> Vec<burn_pack::Tensor> {
self.tensors
.into_iter()
.map(|t| {
burn_pack::Tensor::new(
t.path,
t.data.dtype,
t.data.shape,
Some(t.id.val()),
t.data.bytes,
)
})
.collect()
}
fn from_reader(reader: Reader) -> Result<Self, RecordError> {
let tensors = reader
.into_tensors()?
.into_iter()
.map(|t| {
let id = t.param_id.map(ParamId::from).unwrap_or_else(ParamId::new);
let data = TensorData::from_bytes(t.bytes, t.shape, t.dtype);
RecordTensor {
path: t.name,
id,
data,
}
})
.collect();
Ok(Self::from_tensors(tensors))
}
pub(crate) fn from_module<M: Module>(module: M) -> Self {
let mut collector = Collector::default();
module.visit(&mut collector);
ModuleRecord::from_tensors(collector.tensors)
}
pub(crate) fn apply<M: Module>(self, module: M) -> Result<M, RecordError> {
let validate = self.validate;
let allow_partial = self.allow_partial;
let mut mapper = ModuleRecordMapper::new(self);
let module = module.map(&mut mapper);
if validate && !mapper.errors.is_empty() {
return Err(RecordError::Validation(format!(
"Apply errors: {:?}",
mapper.errors
)));
}
if !allow_partial && !mapper.missing.is_empty() {
return Err(RecordError::Validation(format!(
"Missing tensors: {:?}",
mapper.missing
)));
}
Ok(module)
}
}
#[derive(Default)]
struct Collector {
path: Vec<String>,
tensors: Vec<RecordTensor>,
}
impl Collector {
fn record(&mut self, id: ParamId, data: TensorData) {
self.tensors.push(RecordTensor {
path: self.path.join("."),
id,
data,
});
}
}
impl ModuleVisitor for Collector {
fn enter_module(&mut self, name: &str, _container_type: &str) {
self.path.push(name.to_string());
}
fn exit_module(&mut self, _name: &str, _container_type: &str) {
self.path.pop();
}
fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<D>>) {
self.record(param.id, param.transform_for_save().val().into_data());
}
fn visit_int<const D: usize>(&mut self, param: &Param<Tensor<D, Int>>) {
self.record(param.id, param.transform_for_save().val().into_data());
}
fn visit_bool<const D: usize>(&mut self, param: &Param<Tensor<D, Bool>>) {
self.record(param.id, param.transform_for_save().val().into_data());
}
}
struct ModuleRecordMapper {
path: Vec<String>,
tensors: HashMap<String, (ParamId, TensorData)>,
dtype_policy: DTypePolicy,
missing: Vec<String>,
errors: Vec<String>,
}
impl ModuleRecordMapper {
fn new(record: ModuleRecord) -> Self {
let tensors = record
.tensors
.into_iter()
.map(|t| (t.path, (t.id, t.data)))
.collect();
Self {
path: Vec::new(),
tensors,
dtype_policy: record.dtype_policy,
missing: Vec::new(),
errors: Vec::new(),
}
}
fn take<const D: usize, K: Basic>(
&mut self,
device: &Device,
target_shape: Shape,
module_dtype: impl FnOnce() -> DType,
) -> Option<(Tensor<D, K>, ParamId)> {
let path = self.path.join(".");
let (id, data) = match self.tensors.remove_entry(&path) {
Some(entry) => entry.1,
None => {
self.missing.push(path);
return None;
}
};
let dtype = match self.dtype_policy {
DTypePolicy::FromRecord => data.dtype,
DTypePolicy::CastToModule => module_dtype(),
};
if data.shape != target_shape {
self.errors.push(format!(
"{path}: shape mismatch, expected {:?} but record has {:?}",
target_shape, data.shape
));
return None;
}
Some((Tensor::from_data(data, (device, dtype)), id))
}
}
macro_rules! map_kind {
($method:ident, $kind:ty) => {
fn $method<const D: usize>(
&mut self,
param: Param<Tensor<D, $kind>>,
) -> Param<Tensor<D, $kind>> {
let device = param.lazy_device();
let shape = param.lazy_shape();
match self.take(&device, shape, || param.val().dtype()) {
Some((tensor, record_id)) => param.transform_for_load(tensor, record_id),
None => param,
}
}
};
}
impl ModuleMapper for ModuleRecordMapper {
fn enter_module(&mut self, name: &str, _container_type: &str) {
self.path.push(name.to_string());
}
fn exit_module(&mut self, _name: &str, _container_type: &str) {
self.path.pop();
}
map_kind!(map_float, Float);
map_kind!(map_int, Int);
map_kind!(map_bool, Bool);
}
#[cfg(all(test, feature = "std"))]
mod tests {
use super::*;
use crate as burn;
use crate::module::{Module, Param};
use crate::tensor::Tensor;
use burn_tensor::Device;
#[derive(Module, Debug)]
struct Tiny {
weight: Param<Tensor<2>>,
bias: Param<Tensor<1>>,
}
impl Tiny {
fn new(weight: [[f32; 2]; 2], bias: [f32; 2], device: &Device) -> Self {
Self {
weight: Param::from_data(weight, device),
bias: Param::from_data(bias, device),
}
}
}
#[derive(Module, Debug)]
struct TinyWide {
weight: Param<Tensor<2>>,
bias: Param<Tensor<1>>,
gamma: Param<Tensor<1>>,
}
impl TinyWide {
fn zeros(device: &Device) -> Self {
Self {
weight: Param::from_data([[0.0, 0.0], [0.0, 0.0]], device),
bias: Param::from_data([0.0, 0.0], device),
gamma: Param::from_data([0.0, 0.0], device),
}
}
}
fn weights(model: &Tiny) -> (Vec<f32>, Vec<f32>) {
(
model.weight.val().to_data().to_vec().unwrap(),
model.bias.val().to_data().to_vec().unwrap(),
)
}
#[test]
fn round_trip_in_memory() {
let device = Default::default();
let model = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device);
let bytes = model.into_record().into_bytes().unwrap();
let record = ModuleRecord::from_bytes(bytes).unwrap();
assert_eq!(record.len(), 2);
let loaded = Tiny::new([[0.0; 2]; 2], [0.0; 2], &device).load_record(record);
let (w, b) = weights(&loaded);
assert_eq!(w, vec![1.0, 2.0, 3.0, 4.0]);
assert_eq!(b, vec![5.0, 6.0]);
}
#[test]
fn round_trip_file() {
let device = Default::default();
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("tiny.bpk");
Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device)
.into_record()
.save(&path)
.unwrap();
let record = ModuleRecord::load(&path).unwrap();
let loaded = Tiny::new([[0.0; 2]; 2], [0.0; 2], &device).load_record(record);
let (w, b) = weights(&loaded);
assert_eq!(w, vec![1.0, 2.0, 3.0, 4.0]);
assert_eq!(b, vec![5.0, 6.0]);
}
#[test]
fn missing_tensor_requires_allow_partial() {
let device = Default::default();
let record = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device).into_record();
let strict = TinyWide::zeros(&device).try_load_record(record.clone());
assert!(matches!(strict, Err(RecordError::Validation(_))));
let partial = TinyWide::zeros(&device).try_load_record(record.allow_partial(true));
assert!(partial.is_ok());
let loaded = partial.unwrap();
assert_eq!(
loaded.weight.val().to_data().to_vec::<f32>().unwrap(),
vec![1.0, 2.0, 3.0, 4.0]
);
assert_eq!(
loaded.gamma.val().to_data().to_vec::<f32>().unwrap(),
vec![0.0, 0.0]
);
}
fn tiny_with_dtype(device: &Device, dtype: DType) -> Tiny {
Tiny {
weight: Param::from_tensor(
Tensor::<2>::from_data([[0.0, 0.0], [0.0, 0.0]], device).cast(dtype),
),
bias: Param::from_tensor(Tensor::<1>::from_data([0.0, 0.0], device).cast(dtype)),
}
}
#[test]
fn dtype_policy_from_record_keeps_record_dtype() {
let device = Default::default();
let record = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device).into_record();
let loaded = tiny_with_dtype(&device, DType::F64).load_record(record);
assert_eq!(loaded.weight.val().dtype(), DType::F32);
assert_eq!(loaded.bias.val().dtype(), DType::F32);
}
#[test]
fn dtype_policy_cast_to_module_uses_module_dtype() {
let device = Default::default();
let record = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device).into_record();
let loaded =
tiny_with_dtype(&device, DType::F64).load_record(record.cast_to_module_dtype());
assert_eq!(loaded.weight.val().dtype(), DType::F64);
assert_eq!(loaded.bias.val().dtype(), DType::F64);
assert_eq!(
loaded.weight.val().to_data().to_vec::<f64>().unwrap(),
vec![1.0, 2.0, 3.0, 4.0]
);
}
fn tiny_wrong_bias_shape(device: &Device) -> Tiny {
Tiny {
weight: Param::from_data([[0.0, 0.0], [0.0, 0.0]], device),
bias: Param::from_data([0.0, 0.0, 0.0], device),
}
}
#[test]
fn shape_mismatch_fails_validation() {
let device = Default::default();
let record = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device).into_record();
let result = tiny_wrong_bias_shape(&device).try_load_record(record);
assert!(matches!(result, Err(RecordError::Validation(_))));
}
#[test]
fn load_record_preserves_param_id() {
let device = Default::default();
let model = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device);
let weight_id = model.weight.id;
let bias_id = model.bias.id;
let bytes = model.into_record().into_bytes().unwrap();
let record = ModuleRecord::from_bytes(bytes).unwrap();
let loaded = Tiny::new([[0.0; 2]; 2], [0.0; 2], &device).load_record(record);
assert_eq!(
loaded.weight.id, weight_id,
"weight ParamId should be restored from record"
);
assert_eq!(
loaded.bias.id, bias_id,
"bias ParamId should be restored from record"
);
}
#[test]
fn validate_false_ignores_shape_mismatch() {
let device = Default::default();
let record = Tiny::new([[1.0, 2.0], [3.0, 4.0]], [5.0, 6.0], &device).into_record();
let loaded = tiny_wrong_bias_shape(&device)
.try_load_record(record.validate(false))
.unwrap();
assert_eq!(
loaded.weight.val().to_data().to_vec::<f32>().unwrap(),
vec![1.0, 2.0, 3.0, 4.0]
);
assert_eq!(
loaded.bias.val().to_data().to_vec::<f32>().unwrap(),
vec![0.0, 0.0, 0.0]
);
}
#[derive(Module, Debug)]
struct ColLike {
weight: Param<Tensor<2>>,
}
impl ColLike {
fn new(seed: f32, device: &Device) -> Self {
let init_device = device.clone();
let weight = Param::uninitialized(
crate::module::ParamId::new(),
move |device, _| Tensor::<2>::full([3, 2], seed, device),
init_device,
true,
[3, 2].into(),
)
.init_mapper(|t: Tensor<2>| t.transpose())
.save_mapper(|t: Tensor<2>| t.transpose())
.load_mapper(|t: Tensor<2>| t.transpose());
Self { weight }
}
}
#[test]
fn round_trip_a_shape_mapped_param() {
let device = Default::default();
let saved = ColLike::new(1.0, &device);
assert_eq!(saved.weight.val().dims(), [2, 3]);
let record = saved.into_record();
assert_eq!(
record.tensors[0].data.shape,
Shape::from([3, 2]),
"the record must hold the save form, not the live form"
);
let record = ModuleRecord::from_bytes(record.into_bytes().unwrap()).unwrap();
let loaded = ColLike::new(0.0, &device).load_record(record);
assert_eq!(loaded.weight.val().dims(), [2, 3]);
assert_eq!(
loaded.weight.val().to_data().to_vec::<f32>().unwrap(),
vec![1.0; 6],
"the recorded values must land, mapped back to the live form"
);
}
}