use heapless::Vec as HVec;
use crate::quantized::QuantParams;
pub const MAX_LORA_RANK: usize = 2;
pub const MAX_LORA_DIM: usize = 64;
#[derive(Debug, Clone, Copy)]
pub struct LoRAConfig {
pub rank: usize,
pub dim: usize,
pub scale: i8,
pub frozen: bool,
}
impl Default for LoRAConfig {
fn default() -> Self {
Self {
rank: 1,
dim: 32,
scale: 8, frozen: true,
}
}
}
pub struct MicroLoRA {
a_weights: HVec<i8, { MAX_LORA_DIM * MAX_LORA_RANK }>,
b_weights: HVec<i8, { MAX_LORA_RANK * MAX_LORA_DIM }>,
config: LoRAConfig,
a_params: QuantParams,
b_params: QuantParams,
intermediate: [i32; MAX_LORA_RANK],
}
impl MicroLoRA {
pub fn new(config: LoRAConfig, seed: u32) -> crate::Result<Self> {
if config.rank > MAX_LORA_RANK || config.dim > MAX_LORA_DIM {
return Err(crate::Error::InvalidModel("LoRA dimensions too large"));
}
let mut a_weights = HVec::new();
let mut b_weights = HVec::new();
let mut rng_state = seed;
let mut next_rand = || {
rng_state = rng_state.wrapping_mul(1103515245).wrapping_add(12345);
(((rng_state >> 16) & 0x3F) as i16 - 32) as i8 };
for _ in 0..(config.dim * config.rank) {
a_weights.push(next_rand()).map_err(|_| crate::Error::BufferOverflow)?;
}
for _ in 0..(config.rank * config.dim) {
b_weights.push(0).map_err(|_| crate::Error::BufferOverflow)?;
}
Ok(Self {
a_weights,
b_weights,
config,
a_params: QuantParams::default(),
b_params: QuantParams::default(),
intermediate: [0; MAX_LORA_RANK],
})
}
pub fn from_weights(
config: LoRAConfig,
a_weights: &[i8],
b_weights: &[i8],
) -> crate::Result<Self> {
if a_weights.len() != config.dim * config.rank {
return Err(crate::Error::InvalidModel("A weights size mismatch"));
}
if b_weights.len() != config.rank * config.dim {
return Err(crate::Error::InvalidModel("B weights size mismatch"));
}
let mut a_vec = HVec::new();
let mut b_vec = HVec::new();
for &w in a_weights {
a_vec.push(w).map_err(|_| crate::Error::BufferOverflow)?;
}
for &w in b_weights {
b_vec.push(w).map_err(|_| crate::Error::BufferOverflow)?;
}
Ok(Self {
a_weights: a_vec,
b_weights: b_vec,
config,
a_params: QuantParams::default(),
b_params: QuantParams::default(),
intermediate: [0; MAX_LORA_RANK],
})
}
#[inline]
pub fn apply(&mut self, input: &[i8], output: &mut [i32]) {
let dim = self.config.dim;
let rank = self.config.rank;
let scale = self.config.scale as i32;
for i in 0..rank {
self.intermediate[i] = 0;
}
for r in 0..rank {
let mut sum: i32 = 0;
for d in 0..dim {
sum += input[d] as i32 * self.a_weights[d * rank + r] as i32;
}
self.intermediate[r] = sum >> 4; }
for d in 0..dim {
let mut sum: i32 = 0;
for r in 0..rank {
sum += self.intermediate[r] * self.b_weights[r * dim + d] as i32;
}
output[d] += (sum * scale) >> 8;
}
}
pub fn apply_inplace(&mut self, data: &mut [i32], input: &[i8]) {
self.apply(input, data);
}
pub fn memory_size(&self) -> usize {
self.a_weights.len() + self.b_weights.len()
}
#[cfg(not(feature = "frozen"))]
pub fn update(&mut self, input: &[i8], grad_output: &[i32], learning_rate: i8) {
let dim = self.config.dim;
let rank = self.config.rank;
let lr = learning_rate as i32;
let mut grad_intermediate = [0i32; MAX_LORA_RANK];
for r in 0..rank {
let mut sum: i32 = 0;
for d in 0..dim {
sum += grad_output[d] * self.b_weights[r * dim + d] as i32;
}
grad_intermediate[r] = sum >> 8;
}
for d in 0..dim {
for r in 0..rank {
let grad = (input[d] as i32 * grad_intermediate[r] * lr) >> 12;
let idx = d * rank + r;
let new_val = self.a_weights[idx] as i32 + grad;
self.a_weights[idx] = new_val.clamp(-127, 127) as i8;
}
}
for r in 0..rank {
for d in 0..dim {
let grad = (self.intermediate[r] * grad_output[d] * lr) >> 12;
let idx = r * dim + d;
let new_val = self.b_weights[idx] as i32 + grad;
self.b_weights[idx] = new_val.clamp(-127, 127) as i8;
}
}
}
}
pub struct LoRAStack<const NUM_LAYERS: usize> {
adapters: [Option<MicroLoRA>; NUM_LAYERS],
active_count: usize,
}
impl<const NUM_LAYERS: usize> LoRAStack<NUM_LAYERS> {
pub fn new() -> Self {
Self {
adapters: core::array::from_fn(|_| None),
active_count: 0,
}
}
pub fn add_adapter(&mut self, layer_idx: usize, adapter: MicroLoRA) -> crate::Result<()> {
if layer_idx >= NUM_LAYERS {
return Err(crate::Error::InvalidModel("Layer index out of range"));
}
self.adapters[layer_idx] = Some(adapter);
self.active_count += 1;
Ok(())
}
pub fn get(&mut self, layer_idx: usize) -> Option<&mut MicroLoRA> {
self.adapters.get_mut(layer_idx).and_then(|a| a.as_mut())
}
pub fn total_memory(&self) -> usize {
self.adapters.iter()
.filter_map(|a| a.as_ref())
.map(|a| a.memory_size())
.sum()
}
}
impl<const N: usize> Default for LoRAStack<N> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_micro_lora_creation() {
let config = LoRAConfig {
rank: 2,
dim: 32,
scale: 8,
frozen: true,
};
let lora = MicroLoRA::new(config, 42).unwrap();
assert_eq!(lora.memory_size(), 128);
}
#[test]
fn test_lora_apply() {
let config = LoRAConfig {
rank: 1,
dim: 4,
scale: 64, frozen: true,
};
let a_weights = [16i8, 32, 48, 64]; let b_weights = [64i8, 64, 64, 64];
let mut lora = MicroLoRA::from_weights(config, &a_weights, &b_weights).unwrap();
let input = [64i8, 64, 64, 64];
let mut output = [0i32; 4];
lora.apply(&input, &mut output);
let non_zero_count = output.iter().filter(|&&o| o != 0).count();
assert!(non_zero_count > 0, "At least some outputs should be non-zero, got {:?}", output);
}
#[test]
fn test_lora_stack() {
let mut stack = LoRAStack::<4>::new();
let config = LoRAConfig::default();
let adapter = MicroLoRA::new(config, 42).unwrap();
stack.add_adapter(0, adapter).unwrap();
assert!(stack.get(0).is_some());
assert!(stack.get(1).is_none());
assert!(stack.total_memory() > 0);
}
}