use alloc::boxed::Box;
use alloc::vec::Vec;
use super::applier::{Applier, ApplyResult};
use crate::collector::Collector;
use crate::{ModuleAdapter, PathFilter, TensorSnapshot};
use burn_core::module::Module;
use burn_tensor::backend::Backend;
pub trait ModuleSnapshot<B: Backend>: Module<B> {
fn collect(
&self,
filter: Option<PathFilter>,
adapter: Option<Box<dyn ModuleAdapter>>,
) -> Vec<TensorSnapshot> {
let mut collector = Collector::new(filter, adapter);
self.visit(&mut collector);
collector.into_tensors()
}
fn apply(
&mut self,
snapshots: Vec<TensorSnapshot>,
filter: Option<PathFilter>,
adapter: Option<Box<dyn ModuleAdapter>>,
) -> ApplyResult
where
Self: Sized,
{
let mut applier = Applier::new(snapshots, filter, adapter);
unsafe {
let module = core::ptr::read(self as *const Self);
let new_module = module.map(&mut applier);
core::ptr::write(self as *mut Self, new_module);
}
applier.into_result()
}
fn save_into<P>(&self, store: &mut P) -> Result<(), P::Error>
where
P: ModuleStore,
{
store.collect_from(self)
}
fn load_from<P>(&mut self, store: &mut P) -> Result<ApplyResult, P::Error>
where
P: ModuleStore,
{
store.apply_to(self)
}
}
pub trait ModuleStore {
type Error: core::fmt::Debug + core::fmt::Display;
fn collect_from<B: Backend, M: ModuleSnapshot<B>>(
&mut self,
module: &M,
) -> Result<(), Self::Error>;
fn apply_to<B: Backend, M: ModuleSnapshot<B>>(
&mut self,
module: &mut M,
) -> Result<ApplyResult, Self::Error>;
}
impl<B: Backend, M: Module<B>> ModuleSnapshot<B> for M {}