use alloc::collections::BTreeMap;
use alloc::format;
use alloc::string::String;
use alloc::vec::Vec;
use burn_core::tensor::{Device, TensorData};
use burn_pack::Scalar;
pub use burn_derive::RecordState;
pub fn join_path(prefix: &str, leaf: &str) -> String {
if prefix.is_empty() {
return String::from(leaf);
}
let mut path = String::with_capacity(prefix.len() + 1 + leaf.len());
path.push_str(prefix);
path.push('.');
path.push_str(leaf);
path
}
pub fn join_index(prefix: &str, index: usize) -> String {
format!("{prefix}.{index}")
}
#[derive(Default, Debug)]
pub struct StateSink {
pub tensors: Vec<(String, TensorData)>,
pub scalars: Vec<(String, Scalar)>,
}
impl StateSink {
pub fn push_tensor(&mut self, prefix: &str, leaf: &str, data: TensorData) {
self.tensors.push((join_path(prefix, leaf), data));
}
pub fn push_scalar(&mut self, prefix: &str, leaf: &str, value: Scalar) {
self.scalars.push((join_path(prefix, leaf), value));
}
}
#[derive(Default, Debug)]
pub struct StateSource {
tensors: BTreeMap<String, TensorData>,
scalars: BTreeMap<String, Scalar>,
}
impl StateSource {
pub fn new(scalars: BTreeMap<String, Scalar>) -> Self {
Self {
tensors: BTreeMap::new(),
scalars,
}
}
pub fn insert_tensor(&mut self, name: String, data: TensorData) {
self.tensors.insert(name, data);
}
pub fn take_tensor(&mut self, prefix: &str, leaf: &str) -> Option<TensorData> {
self.tensors.remove(&join_path(prefix, leaf))
}
pub fn take_scalar(&mut self, prefix: &str, leaf: &str) -> Option<Scalar> {
self.scalars.get(&join_path(prefix, leaf)).copied()
}
pub fn has_under(&self, prefix: &str) -> bool {
let pat = join_path(prefix, "");
self.tensors.keys().any(|k| k.starts_with(&pat))
|| self.scalars.keys().any(|k| k.starts_with(&pat))
}
}
pub trait RecordState: Sized + Send + Sync + 'static {
fn state_flatten(&self, prefix: &str, out: &mut StateSink);
fn state_unflatten(prefix: &str, src: &mut StateSource, device: &Device) -> Option<Self>;
}
impl RecordState for () {
fn state_flatten(&self, _prefix: &str, _out: &mut StateSink) {}
fn state_unflatten(_prefix: &str, _src: &mut StateSource, _device: &Device) -> Option<Self> {
Some(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::Tensor;
use burn_core as burn;
fn round_trip<T: RecordState>(state: &T) -> Option<T> {
let mut sink = StateSink::default();
state.state_flatten("p", &mut sink);
let scalars: BTreeMap<String, Scalar> = sink.scalars.into_iter().collect();
let mut source = StateSource::new(scalars);
for (name, data) in sink.tensors {
source.insert_tensor(name, data);
}
T::state_unflatten("p", &mut source, &Device::default())
}
fn tensor(values: &[f32]) -> Tensor<1> {
Tensor::from_data(TensorData::from(values), &Device::default())
}
fn data(t: &Tensor<1>) -> Vec<f32> {
t.clone().into_data().to_vec().unwrap()
}
#[derive(RecordState, Clone, Debug)]
struct Inner<const D: usize> {
weight: Tensor<D>,
step: i64,
}
#[derive(RecordState, Clone, Debug)]
struct Full<const D: usize> {
t: Tensor<D>,
opt_tensor: Option<Tensor<D>>,
history: Vec<Tensor<D>>,
count: usize,
opt_scalar: Option<f64>,
nested: Inner<D>,
opt_nested: Option<Inner<D>>,
}
#[test]
fn all_field_kinds_round_trip() {
let state = Full::<1> {
t: tensor(&[1.0, 2.0]),
opt_tensor: Some(tensor(&[3.0])),
history: vec![tensor(&[4.0]), tensor(&[5.0, 6.0])],
count: 7,
opt_scalar: Some(8.5),
nested: Inner {
weight: tensor(&[9.0]),
step: 10,
},
opt_nested: Some(Inner {
weight: tensor(&[11.0]),
step: 12,
}),
};
let out = round_trip(&state).unwrap();
assert_eq!(data(&out.t), vec![1.0, 2.0]);
assert_eq!(data(&out.opt_tensor.unwrap()), vec![3.0]);
assert_eq!(out.history.len(), 2);
assert_eq!(data(&out.history[0]), vec![4.0]);
assert_eq!(data(&out.history[1]), vec![5.0, 6.0]);
assert_eq!(out.count, 7);
assert_eq!(out.opt_scalar, Some(8.5));
assert_eq!(data(&out.nested.weight), vec![9.0]);
assert_eq!(out.nested.step, 10);
let opt_nested = out.opt_nested.unwrap();
assert_eq!(data(&opt_nested.weight), vec![11.0]);
assert_eq!(opt_nested.step, 12);
}
#[test]
fn absent_optionals_round_trip_to_none() {
let state = Full::<1> {
t: tensor(&[1.0]),
opt_tensor: None,
history: vec![],
count: 0,
opt_scalar: None,
nested: Inner {
weight: tensor(&[2.0]),
step: 0,
},
opt_nested: None,
};
let out = round_trip(&state).unwrap();
assert!(out.opt_tensor.is_none());
assert!(out.history.is_empty());
assert!(out.opt_scalar.is_none());
assert!(out.opt_nested.is_none());
}
#[derive(RecordState, Clone, Debug)]
struct AllOptional<const D: usize> {
x: Option<Tensor<D>>,
y: Option<f64>,
}
#[derive(RecordState, Clone, Debug)]
struct OuterOpt<const D: usize> {
inner: Option<AllOptional<D>>,
}
#[test]
fn optional_all_optional_nested_stays_none() {
let state = OuterOpt::<1> { inner: None };
let out = round_trip(&state).unwrap();
assert!(out.inner.is_none());
}
#[test]
fn optional_all_optional_nested_present_with_content() {
let state = OuterOpt::<1> {
inner: Some(AllOptional {
x: Some(tensor(&[1.0, 2.0])),
y: None,
}),
};
let out = round_trip(&state).unwrap();
let inner = out.inner.expect("present because a leaf was recorded");
assert_eq!(data(&inner.x.unwrap()), vec![1.0, 2.0]);
assert!(inner.y.is_none());
}
#[test]
fn missing_required_tensor_yields_none() {
let mut source = StateSource::new(BTreeMap::new());
assert!(Inner::<1>::state_unflatten("p", &mut source, &Device::default()).is_none());
}
}