use super::{Param, Reparameterization};
use crate as burn;
use crate::module::Module;
use burn_tensor::Tensor;
#[derive(Debug, Module)]
pub struct LoraAdapter {
pub a: Param<Tensor<2>>,
pub b: Param<Tensor<2>>,
pub scale: f64,
}
impl Reparameterization for LoraAdapter {
const NAME: &'static str = "lora";
fn materialize<const D: usize>(&self, base: Tensor<D>) -> Tensor<D> {
let delta = self.delta().reshape(base.shape());
let base = if base.dtype().is_float() && base.dtype() != delta.dtype() {
base.cast(delta.dtype())
} else {
base
};
base + delta
}
}
impl LoraAdapter {
pub fn delta(&self) -> Tensor<2> {
self.a.val().matmul(self.b.val()).mul_scalar(self.scale)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_device;
use burn_tensor::DType;
#[test]
fn materialize_casts_a_dense_base_to_the_factor_dtype() {
let device = test_device();
let base = Tensor::<2>::ones([4, 4], (&device, DType::F32));
let adapter = LoraAdapter {
a: Param::from_tensor(Tensor::<2>::ones([4, 2], (&device, DType::F16))),
b: Param::from_tensor(Tensor::<2>::zeros([2, 4], (&device, DType::F16))),
scale: 1.0,
};
let materialized = adapter.materialize(base);
assert_eq!(materialized.dtype(), DType::F16);
}
}