use crate::{Dropout, DropoutConfig, Linear, LinearConfig};
use ruda_model::{
config::Config,
module::{Initializer, Module},
tensor::{Tensor, backend::Backend},
};
#[derive(Config, Debug)]
pub struct LoRALinearConfig {
pub rank: usize,
pub alpha: f64,
#[config(default = 0.0)]
pub dropout: f64,
}
#[derive(Module, Debug)]
pub struct LoRALinear<B: Backend> {
pub base: Linear<B>,
pub adapter_a: Linear<B>,
pub adapter_b: Linear<B>,
pub dropout: Dropout,
pub scale: f64,
}
impl LoRALinearConfig {
pub fn init<B: Backend>(&self, base: Linear<B>) -> LoRALinear<B> {
assert!(self.rank > 0, "LoRA rank must be positive");
assert!(self.alpha.is_finite(), "LoRA alpha must be finite");
assert!(
self.dropout.is_finite() && (0.0..1.0).contains(&self.dropout),
"LoRA dropout must be in [0, 1)"
);
let [input, output] = base.weight.val().dims();
let device = base.weight.val().device();
let dtype = base.weight.val().dtype();
let mut adapter_a = LinearConfig::new(input, self.rank)
.with_bias(false)
.init(&device);
let mut adapter_b = LinearConfig::new(self.rank, output)
.with_bias(false)
.with_initializer(Initializer::Zeros)
.init(&device);
adapter_a.weight = adapter_a
.weight
.map(|value| value.cast(dtype).detach().require_grad());
adapter_b.weight = adapter_b
.weight
.map(|value| value.cast(dtype).detach().require_grad());
LoRALinear {
base: base.no_grad(),
adapter_a,
adapter_b,
dropout: DropoutConfig::new(self.dropout).init(),
scale: self.alpha / self.rank as f64,
}
}
}
impl<B: Backend> LoRALinear<B> {
pub fn forward<const D: usize>(&self, input: Tensor<B, D>) -> Tensor<B, D> {
let update = self
.adapter_b
.forward(self.adapter_a.forward(self.dropout.forward(input.clone())));
self.base.forward(input) + update.mul_scalar(self.scale)
}
pub fn merge(self) -> Linear<B> {
let update = self
.adapter_a
.weight
.val()
.matmul(self.adapter_b.weight.val())
.mul_scalar(self.scale)
.detach();
let mut base = self.base;
base.weight = base
.weight
.map(|weight| (weight + update).detach().set_require_grad(false));
base
}
}
#[cfg(test)]
mod tests;