use candle_core::{DType, Device, Tensor};
use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarBuilder, VarMap};
use peft_rs::training::{AdapterTrainingConfig, AdapterTrainingState, LrSchedule};
use std::collections::HashMap;
use crate::error::{QLoraError, Result};
use crate::qlora::QuantizedLinear;
#[derive(Debug, Clone)]
pub struct QLoraTrainingConfig {
pub adapter_config: AdapterTrainingConfig,
pub num_epochs: usize,
pub batch_size: usize,
pub log_every: usize,
pub save_every: Option<usize>,
pub warmup_steps: usize,
pub use_paged_optimizer: bool,
pub page_size: usize,
pub max_optimizer_memory: usize,
}
impl Default for QLoraTrainingConfig {
fn default() -> Self {
Self {
adapter_config: AdapterTrainingConfig {
learning_rate: 2e-4,
lr_schedule: LrSchedule::LinearWarmup { warmup_steps: 100 },
weight_decay: 0.01,
gradient_accumulation_steps: 4,
max_grad_norm: Some(1.0),
},
num_epochs: 3,
batch_size: 4,
log_every: 10,
save_every: Some(500),
warmup_steps: 100,
use_paged_optimizer: true,
page_size: 1024 * 1024, max_optimizer_memory: 0, }
}
}
#[derive(Debug)]
pub struct PagedAdamWState {
pub exp_avg: HashMap<String, Tensor>,
pub exp_avg_sq: HashMap<String, Tensor>,
pub steps: HashMap<String, usize>,
pub page_size: usize,
gpu_resident: std::collections::HashSet<String>,
access_order: Vec<String>,
pub max_gpu_memory: usize,
pub current_gpu_usage: usize,
}
impl PagedAdamWState {
#[must_use]
pub fn new(page_size: usize, max_gpu_memory: usize) -> Self {
Self {
exp_avg: HashMap::new(),
exp_avg_sq: HashMap::new(),
steps: HashMap::new(),
page_size,
gpu_resident: std::collections::HashSet::new(),
access_order: Vec::new(),
max_gpu_memory,
current_gpu_usage: 0,
}
}
pub fn init_param(&mut self, name: &str, shape: &[usize], _device: &Device) -> Result<()> {
let cpu_device = Device::Cpu;
let exp_avg = Tensor::zeros(shape, DType::F32, &cpu_device)?;
let exp_avg_sq = Tensor::zeros(shape, DType::F32, &cpu_device)?;
self.exp_avg.insert(name.to_string(), exp_avg);
self.exp_avg_sq.insert(name.to_string(), exp_avg_sq);
self.steps.insert(name.to_string(), 0);
Ok(())
}
#[allow(clippy::if_not_else, clippy::excessive_nesting)]
pub fn page_to_device(&mut self, name: &str, device: &Device) -> Result<(Tensor, Tensor)> {
let exp_avg = self
.exp_avg
.get(name)
.ok_or_else(|| QLoraError::InvalidConfig(format!("No state for param: {name}")))?;
let exp_avg_sq = self
.exp_avg_sq
.get(name)
.ok_or_else(|| QLoraError::InvalidConfig(format!("No state for param: {name}")))?;
if !self.gpu_resident.contains(name) {
let param_bytes = exp_avg.elem_count() * 4 * 2;
if self.max_gpu_memory > 0 {
while self.current_gpu_usage + param_bytes > self.max_gpu_memory
&& !self.access_order.is_empty()
{
if let Some(lru_name) = self.access_order.first().cloned() {
if lru_name != name {
self.gpu_resident.remove(&lru_name);
self.access_order.retain(|n| n != &lru_name);
let lru_bytes = self
.exp_avg
.get(&lru_name)
.map_or(0, |t| t.elem_count() * 4 * 2);
self.current_gpu_usage =
self.current_gpu_usage.saturating_sub(lru_bytes);
} else {
break; }
}
}
}
self.gpu_resident.insert(name.to_string());
self.current_gpu_usage += param_bytes;
}
self.access_order.retain(|n| n != name);
self.access_order.push(name.to_string());
Ok((exp_avg.to_device(device)?, exp_avg_sq.to_device(device)?))
}
pub fn page_to_cpu(&mut self, name: &str, exp_avg: &Tensor, exp_avg_sq: &Tensor) -> Result<()> {
if self.gpu_resident.remove(name) {
let param_bytes = exp_avg.elem_count() * 4 * 2; self.current_gpu_usage = self.current_gpu_usage.saturating_sub(param_bytes);
self.access_order.retain(|n| n != name);
}
self.exp_avg
.insert(name.to_string(), exp_avg.to_device(&Device::Cpu)?);
self.exp_avg_sq
.insert(name.to_string(), exp_avg_sq.to_device(&Device::Cpu)?);
Ok(())
}
pub fn increment_step(&mut self, name: &str) {
if let Some(step) = self.steps.get_mut(name) {
*step += 1;
}
}
#[must_use]
pub fn get_step(&self, name: &str) -> usize {
self.steps.get(name).copied().unwrap_or(0)
}
#[must_use]
pub fn is_gpu_resident(&self, name: &str) -> bool {
self.gpu_resident.contains(name)
}
#[must_use]
pub fn gpu_resident_count(&self) -> usize {
self.gpu_resident.len()
}
}
pub struct PagedAdamW {
lr: f64,
beta1: f64,
beta2: f64,
eps: f64,
weight_decay: f64,
state: PagedAdamWState,
initialized: bool,
}
impl PagedAdamW {
#[must_use]
pub fn new(lr: f64, weight_decay: f64, page_size: usize, max_gpu_memory: usize) -> Self {
Self {
lr,
beta1: 0.9,
beta2: 0.999,
eps: 1e-8,
weight_decay,
state: PagedAdamWState::new(page_size, max_gpu_memory),
initialized: false,
}
}
#[must_use]
pub fn with_betas(mut self, beta1: f64, beta2: f64) -> Self {
self.beta1 = beta1;
self.beta2 = beta2;
self
}
pub fn init(&mut self, params: &[(String, Tensor)]) -> Result<()> {
for (name, param) in params {
let shape = param.shape().dims();
self.state.init_param(name, shape, param.device())?;
}
self.initialized = true;
Ok(())
}
pub fn set_lr(&mut self, lr: f64) {
self.lr = lr;
}
#[must_use]
pub fn lr(&self) -> f64 {
self.lr
}
#[allow(clippy::cast_possible_truncation, clippy::cast_possible_wrap)]
pub fn step_param(&mut self, name: &str, param: &mut Tensor, grad: &Tensor) -> Result<()> {
let device = param.device().clone();
let (mut exp_avg, mut exp_avg_sq) = self.state.page_to_device(name, &device)?;
self.state.increment_step(name);
let step = self.state.get_step(name);
let beta1_tensor = Tensor::new(self.beta1 as f32, &device)?;
let one_minus_beta1 = Tensor::new((1.0 - self.beta1) as f32, &device)?;
exp_avg = exp_avg
.broadcast_mul(&beta1_tensor)?
.broadcast_add(&grad.broadcast_mul(&one_minus_beta1)?)?;
let beta2_tensor = Tensor::new(self.beta2 as f32, &device)?;
let one_minus_beta2 = Tensor::new((1.0 - self.beta2) as f32, &device)?;
let grad_sq = grad.sqr()?;
exp_avg_sq = exp_avg_sq
.broadcast_mul(&beta2_tensor)?
.broadcast_add(&grad_sq.broadcast_mul(&one_minus_beta2)?)?;
let bias_correction1 = 1.0 - self.beta1.powi(step as i32);
let bias_correction2 = 1.0 - self.beta2.powi(step as i32);
let bc1_tensor = Tensor::new(bias_correction1 as f32, &device)?;
let bc2_tensor = Tensor::new(bias_correction2 as f32, &device)?;
let exp_avg_corrected = exp_avg.broadcast_div(&bc1_tensor)?;
let exp_avg_sq_corrected = exp_avg_sq.broadcast_div(&bc2_tensor)?;
let denom = exp_avg_sq_corrected
.sqrt()?
.broadcast_add(&Tensor::new(self.eps as f32, &device)?)?;
let step_size = Tensor::new(self.lr as f32, &device)?;
let update = exp_avg_corrected.broadcast_div(&denom)?;
let weight_decay_term =
param.broadcast_mul(&Tensor::new(self.weight_decay as f32, &device)?)?;
let full_update = update
.broadcast_add(&weight_decay_term)?
.broadcast_mul(&step_size)?;
*param = param.broadcast_sub(&full_update)?;
self.state.page_to_cpu(name, &exp_avg, &exp_avg_sq)?;
Ok(())
}
#[must_use]
pub fn memory_stats(&self) -> (usize, usize) {
let cpu_bytes: usize = self
.state
.exp_avg
.values()
.chain(self.state.exp_avg_sq.values())
.map(|t| t.elem_count() * 4)
.sum();
(cpu_bytes, self.state.current_gpu_usage)
}
}
pub struct QLoraTrainer {
config: QLoraTrainingConfig,
state: AdapterTrainingState,
device: Device,
varmap: VarMap,
optimizer: Option<AdamW>,
paged_optimizer: Option<PagedAdamW>,
accumulation_step: usize,
}
impl QLoraTrainer {
#[must_use]
pub fn new(config: QLoraTrainingConfig, device: Device) -> Self {
let state = AdapterTrainingState::new(config.adapter_config.clone());
Self {
config,
state,
device,
varmap: VarMap::new(),
optimizer: None,
paged_optimizer: None,
accumulation_step: 0,
}
}
#[must_use]
pub fn var_builder(&self) -> VarBuilder<'_> {
VarBuilder::from_varmap(&self.varmap, DType::F32, &self.device)
}
pub fn init_optimizer(&mut self, layers: &[&QuantizedLinear]) -> Result<()> {
if self.config.use_paged_optimizer {
let mut paged = PagedAdamW::new(
self.config.adapter_config.learning_rate,
self.config.adapter_config.weight_decay,
self.config.page_size,
self.config.max_optimizer_memory,
);
let vars = self.varmap.all_vars();
if vars.is_empty() {
return Err(QLoraError::InvalidConfig(
"No trainable parameters found. Layers must be created using trainer.var_builder() \
so `LoRA` weights are registered in the `VarMap`.".into()
));
}
let params: Vec<(String, Tensor)> = self
.varmap
.data()
.lock()
.unwrap()
.iter()
.map(|(name, var)| (name.clone(), var.as_tensor().clone()))
.collect();
paged.init(¶ms)?;
self.paged_optimizer = Some(paged);
let _ = layers.len();
} else {
let vars = self.varmap.all_vars();
if vars.is_empty() {
return Err(QLoraError::InvalidConfig(
"No trainable parameters found. Layers must be created using trainer.var_builder() \
so `LoRA` weights are registered in the `VarMap`.".into()
));
}
let params = ParamsAdamW {
lr: self.config.adapter_config.learning_rate,
weight_decay: self.config.adapter_config.weight_decay,
beta1: 0.9,
beta2: 0.999,
eps: 1e-8,
};
self.optimizer = Some(AdamW::new(vars, params)?);
}
Ok(())
}
#[must_use]
pub fn state(&self) -> &AdapterTrainingState {
&self.state
}
#[must_use]
pub fn current_lr(&self) -> f64 {
self.state.current_lr()
}
#[must_use]
pub fn global_step(&self) -> usize {
self.state.global_step
}
#[must_use]
pub fn epoch(&self) -> usize {
self.state.epoch
}
#[allow(clippy::cast_precision_loss, clippy::excessive_nesting)]
pub fn training_step(
&mut self,
layers: &[&QuantizedLinear],
input: &Tensor,
targets: &Tensor,
) -> Result<f64> {
let mut output = input.clone();
for layer in layers {
output = layer.forward(&output)?;
}
let loss = output.sub(targets)?.sqr()?.mean_all()?;
let accum_steps = self.config.adapter_config.gradient_accumulation_steps;
let scaled_loss = if accum_steps > 1 {
let scale = Tensor::new(1.0 / accum_steps as f32, loss.device())?;
loss.broadcast_mul(&scale)?
} else {
loss.clone()
};
let loss_value = f64::from(loss.to_scalar::<f32>()?);
self.accumulation_step += 1;
if let Some(ref mut optimizer) = self.optimizer {
if self.accumulation_step >= accum_steps {
if let Some(max_norm) = self.config.adapter_config.max_grad_norm {
let _ = max_norm; }
optimizer.backward_step(&scaled_loss)?;
self.accumulation_step = 0;
} else {
let _ = scaled_loss.backward();
}
} else if let Some(ref mut paged_optimizer) = self.paged_optimizer {
if self.accumulation_step >= accum_steps {
let grads = scaled_loss.backward()?;
let mut varmap_data = self.varmap.data().lock().unwrap();
for (name, var) in varmap_data.iter_mut() {
if let Some(grad) = grads.get(var.as_tensor()) {
let mut param = var.as_tensor().clone();
paged_optimizer.step_param(name, &mut param, grad)?;
}
}
drop(varmap_data);
self.accumulation_step = 0;
} else {
let _ = scaled_loss.backward();
}
}
let should_log = self.state.step();
if should_log && self.state.global_step.is_multiple_of(self.config.log_every) {
#[cfg(feature = "logging")]
log::info!(
"Step {} | Loss: {:.4} | LR: {:.2e}",
self.state.global_step,
loss_value,
self.current_lr()
);
}
Ok(loss_value)
}
pub fn training_step_lm(
&mut self,
layers: &[&QuantizedLinear],
input: &Tensor,
target_ids: &Tensor,
) -> Result<f64> {
let mut logits = input.clone();
for layer in layers {
logits = layer.forward(&logits)?;
}
let loss = cross_entropy_loss(&logits, target_ids)?;
let loss_value = f64::from(loss.to_scalar::<f32>()?);
if let Some(ref mut optimizer) = self.optimizer {
optimizer.backward_step(&loss)?;
} else if let Some(ref mut paged_optimizer) = self.paged_optimizer {
let grads = loss.backward()?;
let mut varmap_data = self.varmap.data().lock().unwrap();
for (name, var) in varmap_data.iter_mut() {
if let Some(grad) = grads.get(var.as_tensor()) {
let mut param = var.as_tensor().clone();
paged_optimizer.step_param(name, &mut param, grad)?;
}
}
drop(varmap_data);
}
self.state.step();
Ok(loss_value)
}
pub fn start_epoch(&mut self) {
self.state.new_epoch();
self.accumulation_step = 0;
#[cfg(feature = "logging")]
log::info!("Starting epoch {}", self.state.epoch);
}
#[must_use]
pub fn should_continue(&self) -> bool {
self.state.epoch < self.config.num_epochs
}
pub fn update_lr(&mut self) {
let lr = self.current_lr();
if let Some(ref mut optimizer) = self.optimizer {
optimizer.set_learning_rate(lr);
}
if let Some(ref mut paged) = self.paged_optimizer {
paged.set_lr(lr);
}
}
#[must_use]
pub fn config(&self) -> &QLoraTrainingConfig {
&self.config
}
#[must_use]
pub fn optimizer_memory_stats(&self) -> Option<(usize, usize)> {
self.paged_optimizer.as_ref().map(PagedAdamW::memory_stats)
}
pub fn zero_grad(&mut self) {
self.accumulation_step = 0;
}
}
pub fn cross_entropy_loss(logits: &Tensor, targets: &Tensor) -> Result<Tensor> {
let (batch, seq_len, vocab_size) = logits.dims3()?;
let flat_logits = logits.reshape(&[batch * seq_len, vocab_size])?;
let flat_targets = targets.reshape(&[batch * seq_len])?;
let log_probs = candle_nn::ops::log_softmax(&flat_logits, 1)?;
let target_indices = flat_targets.unsqueeze(1)?;
let gathered = log_probs.gather(&target_indices, 1)?;
let loss = gathered.neg()?.mean_all()?;
Ok(loss)
}
#[derive(Debug, Clone, Default)]
pub struct TrainingMetrics {
pub total_loss: f64,
pub num_steps: usize,
pub best_loss: f64,
pub tokens_processed: usize,
}
impl TrainingMetrics {
#[must_use]
pub fn new() -> Self {
Self {
total_loss: 0.0,
num_steps: 0,
best_loss: f64::MAX,
tokens_processed: 0,
}
}
pub fn update(&mut self, loss: f64, num_tokens: usize) {
self.total_loss += loss;
self.num_steps += 1;
self.tokens_processed += num_tokens;
if loss < self.best_loss {
self.best_loss = loss;
}
}
#[must_use]
#[allow(clippy::cast_precision_loss)]
pub fn average_loss(&self) -> f64 {
if self.num_steps == 0 {
0.0
} else {
self.total_loss / self.num_steps as f64
}
}
pub fn reset(&mut self) {
self.total_loss = 0.0;
self.num_steps = 0;
self.tokens_processed = 0;
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::DType;
#[test]
fn test_training_config_default() {
let config = QLoraTrainingConfig::default();
assert_eq!(config.num_epochs, 3);
assert_eq!(config.batch_size, 4);
assert!((config.adapter_config.learning_rate - 2e-4).abs() < 1e-10);
}
#[test]
fn test_trainer_creation() {
let config = QLoraTrainingConfig::default();
let device = Device::Cpu;
let trainer = QLoraTrainer::new(config, device);
assert_eq!(trainer.global_step(), 0);
assert_eq!(trainer.epoch(), 0);
}
#[test]
fn test_training_metrics() {
let mut metrics = TrainingMetrics::new();
metrics.update(0.5, 128);
metrics.update(0.4, 128);
metrics.update(0.3, 128);
assert_eq!(metrics.num_steps, 3);
assert!((metrics.average_loss() - 0.4).abs() < 1e-10);
assert!((metrics.best_loss - 0.3).abs() < 1e-10);
}
#[test]
fn test_cross_entropy_loss_shape() {
let device = Device::Cpu;
let batch = 2;
let seq_len = 10;
let vocab_size = 100;
let logits = Tensor::zeros(&[batch, seq_len, vocab_size], DType::F32, &device).unwrap();
let targets = Tensor::zeros(&[batch, seq_len], DType::U32, &device).unwrap();
let loss = cross_entropy_loss(&logits, &targets).unwrap();
let dims: &[usize] = loss.dims();
assert!(dims.is_empty(), "Expected scalar loss, got dims: {dims:?}");
}
#[test]
fn test_trainer_epoch_progression() {
let config = QLoraTrainingConfig {
num_epochs: 2,
..Default::default()
};
let device = Device::Cpu;
let mut trainer = QLoraTrainer::new(config, device);
assert!(trainer.should_continue());
trainer.start_epoch();
assert_eq!(trainer.epoch(), 1);
assert!(trainer.should_continue());
trainer.start_epoch();
assert_eq!(trainer.epoch(), 2);
assert!(!trainer.should_continue());
}
}