use alloc::collections::BTreeMap;
use alloc::string::String;
use alloc::vec::Vec;
use burn::store::RecordError;
use burn::tensor::Bytes;
use burn_core as burn;
use burn_pack::{Reader, Scalar, Writer};
#[derive(Default)]
pub struct OptimizerRecord {
pub(crate) tensors: Vec<burn_pack::Tensor>,
pub(crate) scalars: BTreeMap<String, Scalar>,
pub(crate) paths: BTreeMap<String, String>,
}
impl core::fmt::Debug for OptimizerRecord {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("OptimizerRecord")
.field("num_tensors", &self.tensors.len())
.field("num_scalars", &self.scalars.len())
.finish()
}
}
impl OptimizerRecord {
pub fn len(&self) -> usize {
self.tensors.len()
}
pub fn is_empty(&self) -> bool {
self.tensors.is_empty()
}
pub fn into_bytes(self) -> Result<Bytes, RecordError> {
Ok(self.into_writer().into_bytes()?)
}
pub fn from_bytes(bytes: 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> {
self.into_writer().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 into_writer(self) -> Writer {
let mut writer = Writer::new(self.tensors);
for (key, value) in &self.scalars {
writer = writer.with_scalar(key, *value);
}
for (key, value) in &self.paths {
writer = writer.with_metadata(key, value);
}
writer
}
fn from_reader(reader: Reader) -> Result<Self, RecordError> {
let scalars = reader.scalars().clone();
let paths = reader.metadata().clone();
let tensors = reader.into_tensors()?;
Ok(Self {
tensors,
scalars,
paths,
})
}
}