use super::ParamId;
use super::lora::LoraAdapter;
use super::sync_once_cell::SyncOnceCell;
use alloc::format;
use alloc::boxed::Box;
use burn_std::stub::RwLock;
use burn_tensor::{Device, Shape};
use core::ops::Deref;
#[cfg(target_has_atomic = "ptr")]
use alloc::sync::Arc;
#[cfg(not(target_has_atomic = "ptr"))]
use portable_atomic_util::Arc;
#[cfg(target_has_atomic = "ptr")]
type Mapper<T> = Arc<dyn Fn(T) -> T + Send + Sync>;
#[cfg(not(target_has_atomic = "ptr"))]
type Mapper<T> = Arc<Box<dyn Fn(T) -> T + Send + Sync>>;
#[cfg(target_has_atomic = "ptr")]
fn new_mapper<T, F: Fn(T) -> T + Send + Sync + 'static>(func: F) -> Mapper<T> {
Arc::new(func)
}
#[cfg(not(target_has_atomic = "ptr"))]
fn new_mapper<T, F: Fn(T) -> T + Send + Sync + 'static>(func: F) -> Mapper<T> {
Arc::new(Box::new(func))
}
type InitFn<P> = Box<dyn FnOnce(&Device, bool) -> P + Send + Sync>;
fn new_init_fn<P: Parameter, F: FnOnce(&Device, bool) -> P + Send + Sync + 'static>(
func: F,
) -> InitFn<P> {
Box::new(func)
}
pub(crate) struct LazyInitState<T: Parameter> {
pub value: SyncOnceCell<T>,
pub initialization: Option<RwLock<Option<Uninitialized<T>>>>,
}
impl<T: Parameter> LazyInitState<T> {
fn initialized(value: T) -> Arc<Self> {
Arc::new(Self {
value: SyncOnceCell::initialized(value),
initialization: None,
})
}
fn uninitialized(uninit: Uninitialized<T>) -> Arc<Self> {
Arc::new(Self {
value: SyncOnceCell::new(),
initialization: Some(RwLock::new(Some(uninit))),
})
}
fn val(&self) -> &T {
self.value.get_or_init(|| {
let mut init = self
.initialization
.as_ref()
.expect("Should have an initialization when no state provided.")
.write()
.unwrap();
let state = init.take().expect("Should exist when not initialized");
state.initialize()
})
}
}
pub struct Param<T: Parameter> {
pub id: ParamId,
pub(crate) state: Arc<LazyInitState<T>>,
pub(crate) param_mapper: ParamMapper<T>,
pub(crate) require_grad: bool,
pub(crate) adapter: Option<Box<LoraAdapter>>,
}
#[derive(Clone)]
pub struct ParamMapper<T: Parameter> {
load: Option<Mapper<T>>,
save: Option<Mapper<T>>,
}
impl<T: Parameter> core::fmt::Debug for ParamMapper<T> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_fmt(format_args!(
"ParamMapper {{ load: {}, save: {} }}",
self.load.is_some(),
self.save.is_some(),
))
}
}
impl<T: Parameter> ParamMapper<T> {
pub fn on_load(&self, param: T) -> T {
match &self.load {
Some(mapper) => mapper(param),
None => param,
}
}
pub fn on_save(&self, param: T) -> T {
match &self.save {
Some(mapper) => mapper(param),
None => param,
}
}
}
impl<T: Parameter> Default for ParamMapper<T> {
fn default() -> Self {
Self {
load: None,
save: None,
}
}
}
impl<T: Parameter> core::fmt::Display for Param<T> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(format!("Param: {}", self.id).as_str())
}
}
impl<T: Parameter> core::fmt::Debug for Param<T> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(format!("Param: {} - {:?}", self.id, self.param_mapper).as_str())
}
}
pub(crate) mod sealed {
pub trait Sealed {}
}
pub trait Parameter: sealed::Sealed + Clone + core::fmt::Debug + Send {
fn device(&self) -> Device;
fn is_require_grad(&self) -> bool;
fn set_require_grad(self, require_grad: bool) -> Self;
fn shape(&self) -> Shape;
fn load_to_device(self, device: &Device) -> Self;
#[doc(hidden)]
fn compose_lora(self, adapter: &LoraAdapter) -> Self {
let _ = adapter;
self
}
}
#[allow(clippy::type_complexity)]
pub(crate) struct Uninitialized<P: Parameter> {
init: InitFn<P>,
pub(crate) device: Device,
pub(crate) is_require_grad: bool,
pub(crate) shape: Shape,
}
impl<P: Parameter> Uninitialized<P> {
fn initialize(self) -> P {
(self.init)(&self.device, self.is_require_grad)
}
}
impl<T: Parameter> Param<T> {
pub fn initialized(id: ParamId, value: T) -> Self {
let require_grad = value.is_require_grad();
Self {
id,
state: LazyInitState::initialized(value),
param_mapper: Default::default(),
require_grad,
adapter: None,
}
}
pub fn uninitialized<F>(
id: ParamId,
init: F,
device: Device,
is_require_grad: bool,
shape: Shape,
) -> Self
where
F: FnOnce(&Device, bool) -> T + Send + Sync + 'static,
{
Self {
id,
state: LazyInitState::uninitialized(Uninitialized {
init: new_init_fn(init),
device,
is_require_grad,
shape,
}),
param_mapper: Default::default(),
require_grad: is_require_grad,
adapter: None,
}
}
pub fn val(&self) -> T {
let base = self.deref().clone();
match &self.adapter {
Some(adapter) => base.compose_lora(adapter),
None => base,
}
}
pub fn base(&self) -> T {
self.deref().clone()
}
pub fn adapter(&self) -> Option<&LoraAdapter> {
self.adapter.as_deref()
}
pub(crate) fn without_adapter(&self) -> Self {
Self {
id: self.id,
state: self.state.clone(),
param_mapper: self.param_mapper.clone(),
require_grad: self.require_grad,
adapter: None,
}
}
pub(crate) fn with_adapter(mut self, adapter: Option<Box<LoraAdapter>>) -> Self {
self.adapter = adapter;
self
}
pub fn is_initialized(&self) -> bool {
self.state.value.get().is_some()
}
pub fn into_value(self) -> T {
self.consume().1
}
pub fn consume(self) -> (ParamId, T, ParamMapper<T>) {
let tensor = self.deref().clone();
core::mem::drop(self.state);
(self.id, tensor, self.param_mapper)
}
pub fn map<F: FnOnce(T) -> T>(self, func: F) -> Self {
let (id, tensor, param_mapper) = self.consume();
let tensor = func(tensor);
let require_grad = tensor.is_require_grad();
Self {
id,
state: LazyInitState::initialized(tensor),
param_mapper,
require_grad,
adapter: None,
}
}
pub fn from_mapped_value(id: ParamId, value: T, param_mapper: ParamMapper<T>) -> Self {
let require_grad = value.is_require_grad();
Self {
id,
state: LazyInitState::initialized(value),
param_mapper,
require_grad,
adapter: None,
}
}
pub fn load_mapper<F: Fn(T) -> T + Send + Sync + 'static>(mut self, func: F) -> Self {
self.param_mapper.load = Some(new_mapper(func));
self
}
pub fn save_mapper<F: Fn(T) -> T + Send + Sync + 'static>(mut self, func: F) -> Self {
self.param_mapper.save = Some(new_mapper(func));
self
}
pub fn init_mapper<F: Fn(T) -> T + Send + Sync + 'static>(self, func: F) -> Self
where
T: Sync + 'static,
{
let initialization = match &self.state.initialization {
Some(init) => init,
None => return self.map(func),
};
let mut init = initialization.write().unwrap();
match init.as_mut() {
Some(value) => {
let device = value.device.clone();
let is_require_grad = value.is_require_grad;
let shape = value.shape.clone();
core::mem::drop(init);
let base = self;
Self {
id: base.id,
param_mapper: base.param_mapper.clone(),
require_grad: base.require_grad,
adapter: None,
state: LazyInitState::uninitialized(Uninitialized {
init: new_init_fn(move |_a, b| func(base.val()).set_require_grad(b)),
device,
is_require_grad,
shape,
}),
}
}
None => {
core::mem::drop(init);
self.map(func)
}
}
}
pub fn lazy_device(&self) -> Device {
let initialization = match &self.state.initialization {
Some(init) => init,
None => return self.device(),
};
let init = initialization.read().unwrap();
match init.as_ref() {
Some(value) => value.device.clone(),
None => self.device(),
}
}
pub(crate) fn lazy_is_require_grad(&self) -> bool {
let initialization = match &self.state.initialization {
Some(init) => init,
None => return self.is_require_grad(),
};
let init = initialization.read().unwrap();
match init.as_ref() {
Some(value) => value.is_require_grad,
None => self.is_require_grad(),
}
}
pub fn set_require_grad(self, require_grad: bool) -> Self {
let initialization = match &self.state.initialization {
Some(init) => init,
None => return self.map(|tensor| tensor.set_require_grad(require_grad)),
};
let mut init = initialization.write().unwrap();
let mut is_lazy = false;
if let Some(value) = init.as_mut() {
is_lazy = true;
value.is_require_grad = require_grad;
};
core::mem::drop(init);
if is_lazy {
return self;
}
self.map(|tensor| tensor.set_require_grad(require_grad))
}
pub fn lazy_shape(&self) -> burn_tensor::Shape {
let initialization = match &self.state.initialization {
Some(init) => init,
None => return self.shape(),
};
let init = initialization.read().unwrap();
match init.as_ref() {
Some(value) => value.shape.clone(),
None => self.shape(),
}
}
pub fn transform_for_load(self, tensor: T, param_id: ParamId) -> Self {
let mut new_tensor = tensor;
let mapper = self.param_mapper.clone();
let expected_device = self.lazy_device();
let expected_require_grad = self.lazy_is_require_grad();
new_tensor = new_tensor.load_to_device(&expected_device);
new_tensor = mapper.on_load(new_tensor);
new_tensor = new_tensor.set_require_grad(expected_require_grad);
let mut loaded = Self::initialized(param_id, new_tensor);
loaded.param_mapper = mapper;
loaded
}
pub fn transform_for_save(&self) -> Self {
let mut tensor = self.val();
let mapper = self.param_mapper.clone();
tensor = mapper.on_save(tensor);
Self::initialized(self.id, tensor)
}
}
impl<T: Parameter> Clone for Param<T> {
fn clone(&self) -> Self {
Self {
id: self.id,
state: self.state.clone(),
param_mapper: self.param_mapper.clone(),
require_grad: self.require_grad,
adapter: self.adapter.clone(),
}
}
}
impl<T: Parameter> Deref for Param<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
self.state.val()
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn_tensor::Tensor;
fn _assert_sync<T: Sync>() {}
#[test]
fn param_is_sync() {
fn check() {
_assert_sync::<Param<Tensor<2>>>();
}
check();
}
#[cfg(feature = "std")]
#[test]
fn param_concurrent_lazy_init() {
use alloc::vec::Vec;
let device = Default::default();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, _require_grad| Tensor::random([2, 3], Default::default(), device),
device,
false,
[2, 3].into(),
);
std::thread::scope(|s| {
let handles: Vec<_> = (0..4).map(|_| s.spawn(|| param.val())).collect();
let results: Vec<_> = handles.into_iter().map(|h| h.join().unwrap()).collect();
let expected = results[0].to_data();
for result in &results[1..] {
assert_eq!(result.to_data(), expected);
}
});
}
#[test]
fn param_clones_share_lazy_initialization() {
let device = Default::default();
let param_original: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, _require_grad| Tensor::random([2, 3], Default::default(), device),
device,
false,
[2, 3].into(),
);
let param_clone = param_original.clone();
let tensor_original = param_original.val();
assert!(param_original.is_initialized());
assert!(param_clone.is_initialized());
let tensor_clone = param_clone.val();
tensor_original
.into_data()
.assert_eq(&tensor_clone.into_data(), true);
}
#[test]
fn param_set_require_grad_forks_from_shared_state() {
let device = Default::default();
let param1: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, require_grad| Tensor::ones([2, 3], device).set_require_grad(require_grad),
device,
true,
[2, 3].into(),
);
let param2 = param1.clone();
let _tensor1 = param1.val();
assert!(param1.is_initialized());
assert!(param2.is_initialized());
let param2 = param2.set_require_grad(false);
assert_eq!(param2.require_grad, false);
assert_eq!(param1.require_grad, true);
param1
.val()
.into_data()
.assert_eq(¶m2.val().into_data(), true);
}
}