use crate::candle::fast_index::gather_rows;
use crate::candle::optim::RowAdagrad;
use crate::matrix::rand_util::{mix_seed, normal_f32_seeded};
use candle_core::backprop::GradStore;
use candle_core::{DType, Device, Result, Tensor, Var};
use candle_nn::{AdamW, Optimizer, VarBuilder, VarMap};
const ROW_FACTOR_SALT: u64 = 0x4c4f_5241;
pub const U_VAR_NAME: &str = "lora_u";
pub const V_VAR_NAME: &str = "lora_v";
pub struct LoraFactors {
pub u: Tensor,
pub v: Tensor,
}
impl LoraFactors {
pub fn new(n_rows: usize, dim: usize, rank: usize, vs: VarBuilder) -> Result<Self> {
if rank == 0 {
candle_core::bail!("a LoRA residual needs rank ≥ 1");
}
Ok(Self {
u: vs.get_with_hints(
(n_rows, rank),
U_VAR_NAME,
candle_nn::Init::Randn {
mean: 0.0,
stdev: (1.0 / rank as f64).sqrt(),
},
)?,
v: vs.get_with_hints((rank, dim), V_VAR_NAME, candle_nn::Init::Const(0.0))?,
})
}
#[must_use]
pub fn from_parts(u: Tensor, v: Tensor) -> Self {
Self { u, v }
}
#[must_use]
pub fn rank(&self) -> usize {
self.v.dims()[0]
}
pub fn residual_rows(&self, ids: &Tensor) -> Result<Tensor> {
gather_rows(&self.u, ids)?.matmul(&self.v)
}
pub fn project_dims(&self, v_hc: &Tensor) -> Result<Tensor> {
self.u.matmul(&self.v.matmul(v_hc)?)
}
pub fn residual(&self) -> Result<Tensor> {
self.u.matmul(&self.v)
}
pub fn ridge(&self) -> Result<Tensor> {
let uu = self.u.t()?.matmul(&self.u)?;
let vv = self.v.matmul(&self.v.t()?)?;
(uu * vv)?.sum_all()
}
}
pub struct PinnedLora {
pub u: Var,
pub v: Var,
pub u_mask: Tensor,
pub n_pinned: usize,
}
impl PinnedLora {
pub fn new(
n_rows: usize,
dim: usize,
rank: usize,
pinned: &[u32],
seed: u64,
dev: &Device,
) -> Result<Self> {
if rank == 0 {
candle_core::bail!("a LoRA residual needs rank ≥ 1");
}
let draw = normal_f32_seeded(
pinned.len() * rank,
(1.0 / rank as f32).sqrt(),
mix_seed(seed, ROW_FACTOR_SALT),
);
let mut u = vec![0f32; n_rows * rank];
let mut mask = vec![0f32; n_rows];
for (i, &g) in pinned.iter().enumerate() {
let g = g as usize;
if g >= n_rows {
candle_core::bail!("pinned row {g} is outside the {n_rows}-row table");
}
u[g * rank..(g + 1) * rank].copy_from_slice(&draw[i * rank..(i + 1) * rank]);
mask[g] = 1.0;
}
Ok(Self {
u: Var::from_tensor(&Tensor::from_vec(u, (n_rows, rank), dev)?)?,
v: Var::zeros((rank, dim), DType::F32, dev)?,
u_mask: Tensor::from_vec(mask, (n_rows, 1), dev)?,
n_pinned: pinned.len(),
})
}
#[must_use]
pub fn factors(&self) -> LoraFactors {
LoraFactors::from_parts(self.u.as_tensor().clone(), self.v.as_tensor().clone())
}
pub fn residual_rows_masked(&self, ids: &Tensor) -> Result<Tensor> {
gather_rows(&self.u, ids)?
.broadcast_mul(&gather_rows(&self.u_mask, ids)?)?
.matmul(&self.v)
}
pub fn residual_masked(&self) -> Result<Tensor> {
self.u.broadcast_mul(&self.u_mask)?.matmul(&self.v)
}
pub fn residual_rows(&self, ids: &Tensor) -> Result<Tensor> {
self.factors().residual_rows(ids)
}
pub fn residual(&self) -> Result<Tensor> {
self.factors().residual()
}
pub fn ridge(&self) -> Result<Tensor> {
self.factors().ridge()
}
pub fn optimizers(&self, lr: f64, lr_ratio: f32, dev: &Device) -> Result<PinnedLoraOpt> {
Ok(PinnedLoraOpt {
u: RowAdagrad::new(self.u.dims()[0], lr, dev)?,
v: RowAdagrad::new(self.v.dims()[0], lr * f64::from(lr_ratio), dev)?,
})
}
pub fn step(&self, opt: &mut PinnedLoraOpt, grads: &GradStore) -> Result<()> {
if let Some(g) = grads.get(&self.u) {
opt.u.step(&self.u, &g.broadcast_mul(&self.u_mask)?)?;
}
if let Some(g) = grads.get(&self.v) {
opt.v.step(&self.v, g)?;
}
Ok(())
}
}
pub struct PinnedLoraOpt {
pub u: RowAdagrad,
pub v: RowAdagrad,
}
#[derive(Clone, Copy, Debug)]
pub struct LoraPlus<'a> {
pub v_var: &'a str,
pub lr_ratio: f32,
pub ridge: f32,
}
impl LoraPlus<'_> {
pub fn optimizer(&self, varmap: &VarMap, lr: f32) -> anyhow::Result<AdamW> {
let v = crate::candle::frozen_features::trainable_only(varmap, &[self.v_var]);
anyhow::ensure!(
v.len() == 1,
"LoRA+ names `{}` but the model has no such Var",
self.v_var
);
Ok(AdamW::new(
v,
candle_nn::ParamsAdamW {
lr: f64::from(lr * self.lr_ratio),
weight_decay: 0.0,
..Default::default()
},
)?)
}
}
#[must_use]
pub fn join(prefix: &str, slot: &str) -> String {
if prefix.is_empty() {
slot.to_string()
} else {
format!("{prefix}.{slot}")
}
}
#[must_use]
pub fn factor_names(prefix: &str) -> (String, String) {
(join(prefix, U_VAR_NAME), join(prefix, V_VAR_NAME))
}
pub fn fold(varmap: &VarMap, base_name: &str, prefix: &str) -> Result<()> {
let (u_name, v_name) = factor_names(prefix);
let mut tbl = varmap.data().lock().unwrap();
let (Some(u), Some(v)) = (tbl.remove(&u_name), tbl.remove(&v_name)) else {
return Ok(());
};
let base = tbl.get(base_name).ok_or_else(|| {
candle_core::Error::Msg(format!("no {base_name} to fold the LoRA residual into"))
})?;
let factors = LoraFactors::from_parts(u.as_tensor().clone(), v.as_tensor().clone());
let folded = (base.as_tensor() + factors.residual()?)?;
base.set(&folded.detach())
}
#[cfg(test)]
#[path = "lora_tests.rs"]
mod lora_tests;