use super::reparameterization_dyn::{self, DynReparameterization};
use super::{Param, ParamId, Parameter, ParameterValue, Reparameterization};
use crate::module::{
Content, Module, ModuleDisplay, ModuleDisplayDefault, ModuleMapper, ModuleVisitor,
};
use alloc::{boxed::Box, format, string::ToString, vec::Vec};
use burn_tensor::{Bool, Device, Float, Int, Tensor, TensorData};
impl<const D: usize> super::sealed::Sealed for Tensor<D, Float> {
fn is_active(&self) -> bool {
Tensor::is_require_grad(self)
}
fn materialize(self, reparameterization: &dyn DynReparameterization) -> Self {
*reparameterization
.materialize_dyn(Box::new(self))
.downcast::<Tensor<D>>()
.expect("Reparameterization should preserve tensor rank")
}
}
impl<const D: usize> super::sealed::Sealed for Tensor<D, Int> {
fn is_active(&self) -> bool {
false
}
}
impl<const D: usize> super::sealed::Sealed for Tensor<D, Bool> {
fn is_active(&self) -> bool {
false
}
}
impl<const D: usize> ParameterValue for Tensor<D, Float> {}
impl<const D: usize> Parameter for Tensor<D, Float> {
fn is_require_grad(&self) -> bool {
Tensor::is_require_grad(self)
}
fn set_require_grad(self, require_grad: bool) -> Self {
if require_grad && !self.is_autodiff() {
self
} else {
Tensor::set_require_grad(self, require_grad)
}
}
fn device(&self) -> Device {
Tensor::device(self)
}
fn shape(&self) -> burn_std::Shape {
Tensor::shape(self)
}
fn load_to_device(self, device: &Device) -> Self {
if self.device() != *device {
Tensor::to_device(self, device).detach()
} else {
self
}
}
}
impl<const D: usize> ParameterValue for Tensor<D, Int> {}
impl<const D: usize> Parameter for Tensor<D, Int> {
fn is_require_grad(&self) -> bool {
false
}
fn set_require_grad(self, _require_grad: bool) -> Self {
self
}
fn device(&self) -> Device {
Tensor::device(self)
}
fn shape(&self) -> burn_std::Shape {
Tensor::shape(self)
}
fn load_to_device(self, device: &Device) -> Self {
if self.device() != *device {
Tensor::to_device(self, device)
} else {
self
}
}
}
impl<const D: usize> ParameterValue for Tensor<D, Bool> {}
impl<const D: usize> Parameter for Tensor<D, Bool> {
fn is_require_grad(&self) -> bool {
false
}
fn set_require_grad(self, _require_grad: bool) -> Self {
self
}
fn device(&self) -> Device {
Tensor::device(self)
}
fn shape(&self) -> burn_std::Shape {
Tensor::shape(self)
}
fn load_to_device(self, device: &Device) -> Self {
if self.device() != *device {
Tensor::to_device(self, device)
} else {
self
}
}
}
impl<const D: usize> Param<Tensor<D>> {
pub fn from_tensor(value: Tensor<D>) -> Self {
let mut param =
Param::initialized(ParamId::new(), Parameter::set_require_grad(value, true));
param.is_active = true;
param
}
pub fn from_data<T>(data: T, device: &Device) -> Self
where
T: Into<TensorData>,
{
let data: TensorData = data.into();
device.memory_persistent_allocations(data, |data| {
let value = Tensor::from_data(data, device);
let mut param =
Param::initialized(ParamId::new(), Parameter::set_require_grad(value, true));
param.is_active = true;
param
})
}
pub(crate) fn with_reparameterization<R>(mut self, reparameterization: R) -> Self
where
R: Reparameterization,
{
self.reparameterization = Some(reparameterization_dyn::boxed::<R, D>(reparameterization));
self
}
}
impl<const D: usize> Module for Param<Tensor<D>> {
fn visit<V: ModuleVisitor>(&self, visitor: &mut V) {
match self.reparameterization_dyn() {
None => visitor.visit_float(self),
Some(reparameterization) => {
visitor.visit_float(&self.without_reparameterization());
visitor.enter_module(reparameterization.name(), "Reparameterization");
reparameterization_dyn::visit(reparameterization, visitor);
visitor.exit_module(reparameterization.name(), "Reparameterization");
}
}
}
fn map<M: ModuleMapper>(mut self, mapper: &mut M) -> Self {
match self.reparameterization.take() {
None => mapper.map_float(self),
Some(reparameterization) => {
let base = mapper.map_float(self);
mapper.enter_module(reparameterization.name(), "Reparameterization");
let reparameterization = reparameterization_dyn::map(reparameterization, mapper);
mapper.exit_module(reparameterization.name(), "Reparameterization");
base.with_dyn_reparameterization(Some(reparameterization))
}
}
}
fn to_device(mut self, device: &Device) -> Self {
let reparameterization = self.reparameterization.take();
let base = self.map_to_device(device, |tensor| tensor.to_device(device));
match reparameterization {
None => base,
Some(reparameterization) => {
base.with_dyn_reparameterization(Some(reparameterization.to_device_dyn(device)))
}
}
}
fn fork(mut self, device: &Device) -> Self {
let reparameterization = self.reparameterization.take();
let base = self.map_to_device(device, |tensor| {
let is_require_grad = tensor.is_require_grad();
let mut tensor = tensor.to_device(device).detach();
if is_require_grad {
tensor = tensor.require_grad();
}
tensor
});
match reparameterization {
None => base,
Some(reparameterization) => {
base.with_dyn_reparameterization(Some(reparameterization.fork_dyn(device)))
}
}
}
fn collect_devices(&self, mut devices: Vec<Device>) -> Vec<Device> {
let device = self.base().device();
if !devices.contains(&device) {
devices.push(device)
}
if let Some(reparameterization) = self.reparameterization_dyn() {
devices = reparameterization.collect_devices_dyn(devices);
}
devices
}
fn valid(&self) -> Self {
let is_active = self.is_active;
let mut param = Param::from_mapped_value(
self.id,
self.val().without_autodiff().set_require_grad(false),
self.param_mapper.clone(),
);
param.is_active = is_active;
param
}
fn train(mut self) -> Self {
let reparameterization = self.reparameterization.take();
let is_active = self.is_active;
let tensor = Tensor::from_inner(self.val()).set_require_grad(is_active);
let mut base = Param::from_mapped_value(self.id, tensor, self.param_mapper);
base.is_active = is_active;
match reparameterization {
None => base,
Some(reparameterization) => {
base.with_dyn_reparameterization(Some(reparameterization.train_dyn()))
}
}
}
}
impl<const D: usize> ModuleDisplayDefault for Param<Tensor<D>> {
fn content(&self, content: Content) -> Option<Content> {
let id = if content.display_settings.show_param_id() {
format!(", id: {}", self.id)
} else {
"".to_string()
};
let string = format!(
"ParamTensor {{rank: {D}, shape: {:?}, kind: float{id}}}",
self.shape().as_slice()
);
content.add_formatted(&string).optional()
}
}
impl<const D: usize> ModuleDisplay for Param<Tensor<D>> {}
impl<const D: usize> Module for Param<Tensor<D, Int>> {
fn visit<V: ModuleVisitor>(&self, visitor: &mut V) {
visitor.visit_int(self)
}
fn map<M: ModuleMapper>(self, mapper: &mut M) -> Self {
mapper.map_int(self)
}
fn to_device(self, device: &Device) -> Self {
self.map_to_device(device, |tensor| tensor.to_device(device))
}
fn fork(self, device: &Device) -> Self {
self.to_device(device) }
fn collect_devices(&self, mut devices: Vec<Device>) -> Vec<Device> {
let device = self.val().device();
if !devices.contains(&device) {
devices.push(device)
}
devices
}
fn valid(&self) -> Self {
Param::from_mapped_value(
self.id,
self.val().without_autodiff(),
self.param_mapper.clone(),
)
}
fn train(self) -> Self {
Param::from_mapped_value(self.id, Tensor::from_inner(self.val()), self.param_mapper)
}
}
impl<const D: usize> ModuleDisplayDefault for Param<Tensor<D, Int>> {
fn content(&self, content: Content) -> Option<Content> {
let id = if content.display_settings.show_param_id() {
format!(", id: {}", self.id)
} else {
"".to_string()
};
let string = format!(
"ParamTensor {{rank: {D}, shape: {:?}, kind: int{id}}}",
self.shape().as_slice()
);
content.add_formatted(&string).optional()
}
}
impl<const D: usize> ModuleDisplay for Param<Tensor<D, Int>> {}
impl<const D: usize> Module for Param<Tensor<D, Bool>> {
fn visit<V: ModuleVisitor>(&self, visitor: &mut V) {
visitor.visit_bool(self)
}
fn map<M: ModuleMapper>(self, mapper: &mut M) -> Self {
mapper.map_bool(self)
}
fn to_device(self, device: &Device) -> Self {
self.map_to_device(device, |tensor| tensor.to_device(device))
}
fn fork(self, device: &Device) -> Self {
self.to_device(device) }
fn collect_devices(&self, mut devices: Vec<Device>) -> Vec<Device> {
let device = self.val().device();
if !devices.contains(&device) {
devices.push(device)
}
devices
}
fn valid(&self) -> Self {
Param::from_mapped_value(
self.id,
self.val().without_autodiff(),
self.param_mapper.clone(),
)
}
fn train(self) -> Self {
Param::from_mapped_value(self.id, Tensor::from_inner(self.val()), self.param_mapper)
}
}
impl<const D: usize> ModuleDisplayDefault for Param<Tensor<D, Bool>> {
fn content(&self, content: Content) -> Option<Content> {
let id = if content.display_settings.show_param_id() {
format!(", id: {}", self.id)
} else {
"".to_string()
};
let string = format!(
"ParamTensor {{rank: {D}, shape: {:?}, kind: bool{id}}}",
self.shape().as_slice()
);
content.add_formatted(&string).optional()
}
}
impl<const D: usize> ModuleDisplay for Param<Tensor<D, Bool>> {}
#[cfg(all(test, feature = "std", feature = "autodiff"))]
mod tests {
use super::*;
use crate::{
module::{LoraAdapter, Module},
test_device,
};
use burn_tensor::Distribution;
#[test]
fn set_require_grad_updates_lazy_lifecycle_state() {
let device = test_device().autodiff();
let param: 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 param = param.set_require_grad(false);
assert!(!param.is_initialized());
assert!(!param.is_active);
assert!(!param.val().is_require_grad());
let param = param.valid().train();
assert!(!param.is_require_grad());
assert!(!param.is_active);
}
#[test]
fn set_require_grad_on_a_plain_tensor_is_applied_by_train() {
let device = test_device();
let param = Param::initialized(
ParamId::new(),
Tensor::<2>::ones([2, 3], &device).set_require_grad(false),
)
.set_require_grad(true);
assert!(!param.is_require_grad());
assert!(param.is_active);
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
}
#[test]
fn trainable_param_created_on_a_plain_device_is_applied_by_train() {
let device = test_device();
let param = Param::from_tensor(Tensor::<2>::ones([2, 3], &device));
assert!(!param.is_require_grad());
assert!(param.is_active);
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
}
#[test]
fn lazy_activation_setting_on_a_plain_device_is_applied_by_train() {
let device = test_device();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, require_grad| Tensor::ones([2, 3], device).set_require_grad(require_grad),
device,
false,
[2, 3].into(),
)
.set_require_grad(true);
assert!(!param.is_initialized());
assert!(param.is_active);
assert!(!param.val().is_require_grad());
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
}
#[test]
fn mapping_a_validation_param_preserves_what_train_restores() {
let device = test_device().autodiff();
let param = Param::from_tensor(Tensor::<2>::ones([2, 3], &device))
.valid()
.map(|tensor| tensor);
assert!(!param.is_require_grad());
assert!(param.is_active);
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
}
#[test]
fn mapped_value_reconstruction_preserves_what_train_restores() {
let device = test_device().autodiff();
let valid = Param::from_tensor(Tensor::<2>::ones([2, 3], &device)).valid();
let (id, tensor, mapper) = valid.consume();
let param = Param::from_mapped_value(id, tensor, mapper);
assert!(!param.is_require_grad());
assert!(param.is_active);
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
}
#[test]
fn loading_a_validation_param_preserves_what_train_restores() {
let device = test_device().autodiff();
let valid = Param::from_tensor(Tensor::<2>::ones([2, 3], &device)).valid();
let record = Tensor::<2>::zeros([2, 3], &test_device());
let param = valid.transform_for_load(record, ParamId::new());
assert!(!param.is_require_grad());
assert!(param.is_active);
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
}
#[test]
fn test_param_require_grad_stateful() {
let device = test_device().autodiff();
let tensor = Tensor::<2>::ones([3, 3], &device).require_grad();
let param = Param::initialized(ParamId::new(), tensor);
assert!(param.is_require_grad());
assert!(param.is_active);
let param = param.valid();
assert!(!param.is_require_grad());
assert!(param.is_active);
let param = param.train();
assert!(param.is_require_grad());
assert!(param.is_active);
let param = param.no_grad();
assert!(!param.is_require_grad());
assert!(!param.is_active);
let param = param.valid();
assert!(!param.is_require_grad()); assert!(!param.is_active);
let param = param.train();
assert!(!param.is_require_grad());
assert!(!param.is_active); }
#[test]
fn a_lazy_param_with_an_init_mapper_trains_on_an_autodiff_device() {
let device = test_device().autodiff();
let param: 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(),
)
.init_mapper(|tensor| tensor.mul_scalar(2.0));
let value = param.val();
let grads = value.clone().sum().backward();
value
.into_data()
.assert_eq(&TensorData::from([[2.0f32; 3]; 2]), false);
param
.grad(&grads)
.expect("the mapped value is the leaf that receives the gradient")
.into_data()
.assert_eq(&TensorData::from([[1.0f32; 3]; 2]), false);
}
#[test]
fn counting_a_lazy_param_leaves_it_uninitialized() {
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, require_grad| Tensor::ones([2, 3], device).set_require_grad(require_grad),
test_device(),
true,
[2, 3].into(),
);
assert_eq!(Module::num_params(¶m), 6);
assert!(!param.is_initialized());
}
#[test]
fn a_lazy_param_moved_then_loaded_never_initializes() {
let device = test_device();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|_, _| panic!("the moved parameter initialized before loading"),
device.clone(),
false,
[2, 3].into(),
);
let moved = param.to_device(&device.clone().autodiff());
assert!(!moved.is_initialized());
moved.transform_for_load(Tensor::ones([2, 3], &device), ParamId::new());
}
#[test]
fn a_moved_lazy_param_keeps_the_autodiff_context_it_was_built_with() {
let device = test_device();
let lazy_ones = |device: &Device| -> Param<Tensor<2>> {
Param::uninitialized(
ParamId::new(),
|device, require_grad| Tensor::ones([2, 3], device).set_require_grad(require_grad),
device.clone(),
true,
[2, 3].into(),
)
};
let onto_autodiff = lazy_ones(&device).to_device(&device.clone().autodiff());
let onto_plain = lazy_ones(&device.clone().autodiff()).fork(&device);
assert!(!onto_autodiff.lazy_device().is_autodiff());
assert!(onto_plain.lazy_device().is_autodiff());
}
#[test]
fn a_moved_param_keeps_its_reparameterization() {
let device = test_device();
let adapter = || LoraAdapter {
a: Param::from_tensor(Tensor::<2>::ones([3, 1], &device)),
b: Param::from_tensor(Tensor::<2>::zeros([1, 3], &device)),
scale: 1.0,
};
let lazy: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, _| Tensor::ones([3, 3], device),
device.clone(),
false,
[3, 3].into(),
)
.with_reparameterization(adapter());
let initialized = Param::from_tensor(Tensor::<2>::ones([3, 3], &device))
.with_reparameterization(adapter());
let target = device.clone().autodiff();
assert!(lazy.to_device(&target).adapter().is_some());
assert!(initialized.fork(&target).adapter().is_some());
}
#[test]
fn init_mapper_preserves_lora_after_a_clone_initializes_the_base() {
let device = test_device();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, _| Tensor::ones([3, 3], device),
device.clone(),
false,
[3, 3].into(),
)
.with_reparameterization(crate::module::LoraAdapter {
a: Param::from_tensor(Tensor::<2>::ones([3, 1], &device)),
b: Param::from_tensor(Tensor::<2>::ones([1, 3], &device)),
scale: 2.0,
});
let clone = param.clone();
assert!(!param.is_initialized());
let mapped = param.init_mapper(|value| value.mul_scalar(2.0));
clone
.val()
.into_data()
.assert_eq(&TensorData::from([[3.0f32; 3]; 3]), true);
drop(clone);
assert!(!mapped.is_initialized());
mapped
.val()
.into_data()
.assert_eq(&TensorData::from([[6.0f32; 3]; 3]), true);
}
#[test]
fn a_lazy_int_param_moved_never_initializes() {
let device = test_device();
let param: Param<Tensor<2, Int>> = Param::uninitialized(
ParamId::new(),
|_, _| panic!("the moved parameter initialized"),
device.clone(),
false,
[2, 3].into(),
);
let param = param.to_device(&device.clone().autodiff());
assert!(!param.is_initialized());
assert!(!param.lazy_device().is_autodiff());
}
#[test]
fn a_lazy_bool_param_moved_never_initializes() {
let device = test_device();
let param: Param<Tensor<2, Bool>> = Param::uninitialized(
ParamId::new(),
|_, _| panic!("the moved parameter initialized"),
device.clone(),
false,
[2, 3].into(),
);
let param = param.to_device(&device.clone().autodiff());
assert!(!param.is_initialized());
assert!(!param.lazy_device().is_autodiff());
}
#[test]
fn a_lazy_param_forked_keeps_its_gradient_requirement() {
let device = test_device();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, require_grad| Tensor::ones([2, 3], device).set_require_grad(require_grad),
device.clone().autodiff(),
true,
[2, 3].into(),
);
let param = param.fork(&device);
assert!(!param.is_initialized());
assert!(param.val().is_require_grad());
}
#[test]
fn a_lazy_param_with_an_init_mapper_follows_the_move() {
let device = test_device().autodiff();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, require_grad| Tensor::ones([2, 3], device).set_require_grad(require_grad),
device.clone(),
true,
[2, 3].into(),
)
.init_mapper(|tensor| tensor.mul_scalar(2.0));
let param = param.fork(&device);
assert!(!param.is_initialized());
let value = param.val();
assert!(value.device().is_autodiff());
assert!(value.is_require_grad());
}
#[test]
fn a_lazy_param_shared_with_a_clone_initializes_before_moving() {
let device = test_device();
let param: Param<Tensor<2>> = Param::uninitialized(
ParamId::new(),
|device, _| Tensor::random([2, 3], Distribution::Default, device),
device.clone(),
false,
[2, 3].into(),
);
let clone = param.clone();
let moved = param.to_device(&device.autodiff());
assert!(clone.is_initialized());
moved
.val()
.into_data()
.assert_eq(&clone.val().into_data(), true);
}
}