use std::cell::RefCell;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
#[derive(Debug)]
pub struct LoraScale(AtomicU32);
impl LoraScale {
pub fn new(scale: f32) -> Arc<Self> {
Arc::new(Self(AtomicU32::new(scale.to_bits())))
}
pub fn get(&self) -> f32 {
f32::from_bits(self.0.load(Ordering::Relaxed))
}
pub fn set(&self, scale: f32) {
self.0.store(scale.to_bits(), Ordering::Relaxed);
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoraShapeError {
pub rows: usize,
pub cols: usize,
pub rank: usize,
pub a_len: usize,
pub b_len: usize,
}
impl std::fmt::Display for LoraShapeError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"LoRA pair does not fit a [{} x {}] base at rank {}: lora_a has {} values \
(want rank x cols = {}), lora_b has {} values (want rows x rank = {})",
self.rows,
self.cols,
self.rank,
self.a_len,
self.rank * self.cols,
self.b_len,
self.rows * self.rank
)
}
}
impl std::error::Error for LoraShapeError {}
#[derive(Debug)]
pub struct LoraDelta {
a: Vec<f32>,
b: Vec<f32>,
rank: usize,
rows: usize,
cols: usize,
alpha_over_rank: f32,
scale: Arc<LoraScale>,
}
impl LoraDelta {
pub fn new(
a: Vec<f32>,
b: Vec<f32>,
rank: usize,
rows: usize,
cols: usize,
alpha: f32,
scale: Arc<LoraScale>,
) -> Result<Self, LoraShapeError> {
if rank == 0 || a.len() != rank * cols || b.len() != rows * rank {
return Err(LoraShapeError {
rows,
cols,
rank,
a_len: a.len(),
b_len: b.len(),
});
}
let alpha_over_rank = if alpha != 0.0 {
alpha / rank as f32
} else {
1.0
};
Ok(Self {
a,
b,
rank,
rows,
cols,
alpha_over_rank,
scale,
})
}
pub fn rank(&self) -> usize {
self.rank
}
pub fn rows(&self) -> usize {
self.rows
}
pub fn cols(&self) -> usize {
self.cols
}
pub fn scale_handle(&self) -> &Arc<LoraScale> {
&self.scale
}
pub fn resident_bytes(&self) -> usize {
(self.a.len() + self.b.len()) * 4
}
#[inline]
fn effective_scale(&self) -> f32 {
self.scale.get() * self.alpha_over_rank
}
#[inline]
fn project(&self, x: &[f32], y: &mut [f32]) {
for (k, yk) in y.iter_mut().enumerate() {
let row = &self.a[k * self.cols..(k + 1) * self.cols];
*yk = dot(row, x);
}
}
#[inline]
fn accumulate(&self, s: f32, y: &[f32], out: &mut [f32]) {
for (r, o) in out.iter_mut().enumerate() {
let brow = &self.b[r * self.rank..(r + 1) * self.rank];
*o += s * dot(brow, y);
}
}
pub fn add_to(&self, x: &[f32], out: &mut [f32], scratch: &mut Vec<f32>) {
debug_assert_eq!(x.len(), self.cols);
debug_assert_eq!(out.len(), self.rows);
let s = self.effective_scale();
if s == 0.0 {
return;
}
scratch.clear();
scratch.resize(self.rank, 0.0);
self.project(x, scratch);
self.accumulate(s, scratch, out);
}
pub fn add_row_to(&self, r: usize, out_row: &mut [f32]) {
debug_assert!(r < self.rows);
debug_assert_eq!(out_row.len(), self.cols);
let s = self.effective_scale();
if s == 0.0 {
return;
}
let brow = &self.b[r * self.rank..(r + 1) * self.rank];
for (k, &bk) in brow.iter().enumerate() {
let arow = &self.a[k * self.cols..(k + 1) * self.cols];
let sb = s * bk;
for (o, &a) in out_row.iter_mut().zip(arow) {
*o += sb * a;
}
}
}
pub fn add_batch_to(&self, x_batch: &[f32], batch: usize, out: &mut [f32]) {
debug_assert_eq!(x_batch.len(), batch * self.cols);
debug_assert_eq!(out.len(), batch * self.rows);
let s = self.effective_scale();
if s == 0.0 || batch == 0 {
return;
}
let rows = self.rows;
let cols = self.cols;
crate::par::chunks_mut(out, rows, 1, |b, out_b| {
SCRATCH.with(|cell| {
let mut y = cell.borrow_mut();
y.clear();
y.resize(self.rank, 0.0);
self.project(&x_batch[b * cols..(b + 1) * cols], &mut y);
self.accumulate(s, &y, out_b);
});
});
}
}
thread_local! {
static SCRATCH: RefCell<Vec<f32>> = const { RefCell::new(Vec::new()) };
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
#[derive(Debug, Default)]
pub struct LoraStack {
deltas: Vec<LoraDelta>,
}
impl LoraStack {
pub fn new(delta: LoraDelta) -> Self {
Self {
deltas: vec![delta],
}
}
pub fn push(&mut self, delta: LoraDelta) {
self.deltas.push(delta);
}
pub fn len(&self) -> usize {
self.deltas.len()
}
pub fn is_empty(&self) -> bool {
self.deltas.is_empty()
}
pub fn deltas(&self) -> &[LoraDelta] {
&self.deltas
}
pub fn resident_bytes(&self) -> usize {
self.deltas.iter().map(LoraDelta::resident_bytes).sum()
}
pub fn add_to(&self, x: &[f32], out: &mut [f32]) {
SCRATCH.with(|cell| {
let mut scratch = cell.borrow_mut();
for d in &self.deltas {
d.add_to(x, out, &mut scratch);
}
});
}
pub fn add_row_to(&self, r: usize, out_row: &mut [f32]) {
for d in &self.deltas {
d.add_row_to(r, out_row);
}
}
pub fn add_batch_to(&self, x_batch: &[f32], batch: usize, out: &mut [f32]) {
for d in &self.deltas {
d.add_batch_to(x_batch, batch, out);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn delta(rows: usize, cols: usize, rank: usize, alpha: f32, scale: f32) -> LoraDelta {
let a: Vec<f32> = (0..rank * cols).map(|i| (i as f32 * 0.37).sin()).collect();
let b: Vec<f32> = (0..rows * rank).map(|i| (i as f32 * 0.53).cos()).collect();
LoraDelta::new(a, b, rank, rows, cols, alpha, LoraScale::new(scale)).unwrap()
}
fn dense_delta(d: &LoraDelta, x: &[f32]) -> Vec<f32> {
let s = d.effective_scale();
(0..d.rows)
.map(|r| {
let mut acc = 0.0;
for k in 0..d.rank {
let bk = d.b[r * d.rank + k];
for (c, &xc) in x.iter().enumerate() {
acc += s * bk * d.a[k * d.cols + c] * xc;
}
}
acc
})
.collect()
}
fn close(a: &[f32], b: &[f32]) {
assert_eq!(a.len(), b.len());
for (x, y) in a.iter().zip(b) {
assert!((x - y).abs() < 1e-4, "{x} vs {y}");
}
}
#[test]
fn matvec_delta_matches_the_materialised_product() {
let d = delta(7, 5, 3, 6.0, 0.8);
let x: Vec<f32> = (0..5).map(|i| i as f32 - 2.0).collect();
let mut out = vec![0.0; 7];
d.add_to(&x, &mut out, &mut Vec::new());
close(&out, &dense_delta(&d, &x));
}
#[test]
fn alpha_over_rank_is_llama_cpp_s_get_scale() {
let with = delta(4, 4, 3, 6.0, 0.5);
let without = delta(4, 4, 3, 0.0, 0.5);
assert!((with.effective_scale() - 1.0).abs() < 1e-7);
assert!((without.effective_scale() - 0.5).abs() < 1e-7);
}
#[test]
fn row_delta_agrees_with_the_matvec_on_a_unit_vector() {
let d = delta(6, 8, 2, 4.0, 1.3);
for r in 0..6 {
let mut row = vec![0.0; 8];
d.add_row_to(r, &mut row);
for c in 0..8 {
let mut e = vec![0.0; 8];
e[c] = 1.0;
let mut out = vec![0.0; 6];
d.add_to(&e, &mut out, &mut Vec::new());
assert!((out[r] - row[c]).abs() < 1e-5);
}
}
}
#[test]
fn batch_delta_is_the_matvec_delta_per_position() {
let d = delta(5, 6, 4, 8.0, 0.25);
let batch = 3;
let x: Vec<f32> = (0..batch * 6).map(|i| (i as f32 * 0.11).cos()).collect();
let mut out = vec![0.0; batch * 5];
d.add_batch_to(&x, batch, &mut out);
for b in 0..batch {
let mut one = vec![0.0; 5];
d.add_to(&x[b * 6..(b + 1) * 6], &mut one, &mut Vec::new());
close(&out[b * 5..(b + 1) * 5], &one);
}
}
#[test]
fn a_zero_scale_adds_nothing_bit_for_bit() {
let d = delta(5, 6, 4, 8.0, 0.0);
let x = vec![1.0; 6];
let mut out = vec![0.1, 0.2, 0.3, 0.4, 0.5];
let before = out.clone();
d.add_to(&x, &mut out, &mut Vec::new());
assert_eq!(out, before);
let mut batch = vec![0.7; 10];
d.add_batch_to(&[x.clone(), x.clone()].concat(), 2, &mut batch);
assert_eq!(batch, vec![0.7; 10]);
}
#[test]
fn the_scale_is_read_at_apply_time() {
let d = delta(5, 6, 4, 0.0, 1.0);
let x = vec![1.0; 6];
let mut at_one = vec![0.0; 5];
d.add_to(&x, &mut at_one, &mut Vec::new());
d.scale_handle().set(0.5);
let mut at_half = vec![0.0; 5];
d.add_to(&x, &mut at_half, &mut Vec::new());
let halved: Vec<f32> = at_one.iter().map(|v| v * 0.5).collect();
close(&at_half, &halved);
}
#[test]
fn a_stack_sums_its_adapters() {
let d1 = delta(5, 6, 2, 0.0, 1.0);
let d2 = delta(5, 6, 3, 0.0, 0.5);
let x: Vec<f32> = (0..6).map(|i| i as f32 * 0.3 - 1.0).collect();
let mut want = vec![0.0; 5];
d1.add_to(&x, &mut want, &mut Vec::new());
d2.add_to(&x, &mut want, &mut Vec::new());
let mut stack = LoraStack::new(d1);
stack.push(d2);
let mut got = vec![0.0; 5];
stack.add_to(&x, &mut got);
close(&got, &want);
}
#[test]
fn a_pair_of_the_wrong_shape_is_refused_with_every_dimension() {
let err = LoraDelta::new(
vec![0.0; 6],
vec![0.0; 5],
2,
3,
4,
1.0,
LoraScale::new(1.0),
)
.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("[3 x 4]"), "{msg}");
assert!(msg.contains("rank 2"), "{msg}");
assert!(msg.contains("want rank x cols = 8"), "{msg}");
assert!(msg.contains("want rows x rank = 6"), "{msg}");
}
}