use crate::error::{Error, Result};
use crate::progress::{ProgressEvent, ProgressFn};
use candle_core::{Device, Tensor};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fs;
use std::path::Path;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LoRAConfig {
pub r: usize,
pub lora_alpha: f64,
pub lora_dropout: Option<f64>,
pub target_modules: Vec<String>,
pub base_model_name_or_path: Option<String>,
pub task_type: Option<String>,
pub peft_type: Option<String>,
pub use_rslora: Option<bool>,
pub fan_in_fan_out: Option<bool>,
pub bias: Option<String>,
pub modules_to_save: Option<Vec<String>>,
#[serde(default)]
pub custom_config: HashMap<String, serde_json::Value>,
}
impl LoRAConfig {
pub fn new(r: usize, lora_alpha: f64) -> Self {
Self {
r,
lora_alpha,
lora_dropout: Some(0.1),
target_modules: vec!["q_proj".to_string(), "v_proj".to_string()],
base_model_name_or_path: None,
task_type: Some("CAUSAL_LM".to_string()),
peft_type: Some("LORA".to_string()),
use_rslora: Some(false),
fan_in_fan_out: Some(false),
bias: Some("none".to_string()),
modules_to_save: None,
custom_config: HashMap::new(),
}
}
pub fn with_target_modules(mut self, modules: Vec<String>) -> Self {
self.target_modules = modules;
self
}
pub fn with_base_model<S: Into<String>>(mut self, path: S) -> Self {
self.base_model_name_or_path = Some(path.into());
self
}
pub fn with_task_type<S: Into<String>>(mut self, task_type: S) -> Self {
self.task_type = Some(task_type.into());
self
}
pub fn with_custom<K: Into<String>, V: Into<serde_json::Value>>(
mut self,
key: K,
value: V,
) -> Self {
self.custom_config.insert(key.into(), value.into());
self
}
pub fn scaling_factor(&self) -> f64 {
if self.r == 0 {
1.0
} else {
self.lora_alpha / self.r as f64
}
}
pub fn is_target_module(&self, module_name: &str) -> bool {
self.target_modules
.iter()
.any(|target| module_name.contains(target) || module_name.ends_with(target))
}
}
#[derive(Debug, Clone)]
pub struct LoRAWeights {
pub lora_a: Tensor,
pub lora_b: Tensor,
pub lora_bias: Option<Tensor>,
pub scaling: f64,
}
impl LoRAWeights {
pub fn new(lora_a: Tensor, lora_b: Tensor, scaling: f64) -> Self {
Self {
lora_a,
lora_b,
lora_bias: None,
scaling,
}
}
pub fn with_bias(mut self, bias: Tensor) -> Self {
self.lora_bias = Some(bias);
self
}
pub fn compute_update(&self) -> Result<Tensor> {
let update = self.lora_b.matmul(&self.lora_a)?;
let scaled_update = (update * self.scaling)?;
Ok(scaled_update)
}
pub fn target_shape(&self) -> Result<(usize, usize)> {
let a_shape = self.lora_a.shape();
let b_shape = self.lora_b.shape();
if a_shape.rank() != 2 || b_shape.rank() != 2 {
return Err(Error::invalid_format(
"LoRA matrices must be 2-dimensional".to_string(),
));
}
let input_dim = a_shape.dims()[1];
let output_dim = b_shape.dims()[0];
Ok((output_dim, input_dim))
}
}
#[derive(Debug, Clone)]
pub struct LoRAAdapter {
pub config: LoRAConfig,
pub weights: HashMap<String, LoRAWeights>,
pub metadata: HashMap<String, String>,
}
impl LoRAAdapter {
pub fn new(config: LoRAConfig) -> Self {
Self {
config,
weights: HashMap::new(),
metadata: HashMap::new(),
}
}
pub fn add_module(&mut self, module_name: String, weights: LoRAWeights) -> Result<()> {
if !self.config.is_target_module(&module_name) {
return Err(Error::invalid_config(format!(
"Module '{}' is not in target_modules list",
module_name
)));
}
self.weights.insert(module_name, weights);
Ok(())
}
pub fn get_module(&self, module_name: &str) -> Option<&LoRAWeights> {
self.weights.get(module_name)
}
pub fn modules(&self) -> Vec<&String> {
self.weights.keys().collect()
}
pub fn num_modules(&self) -> usize {
self.weights.len()
}
pub fn add_metadata<K: Into<String>, V: Into<String>>(&mut self, key: K, value: V) {
self.metadata.insert(key.into(), value.into());
}
pub fn merge_adapters(adapters: &[(Self, f64)]) -> Result<Self> {
if adapters.is_empty() {
return Err(Error::invalid_config(
"Cannot merge empty list of adapters".to_string(),
));
}
let base_config = adapters[0].0.config.clone();
let mut merged = Self::new(base_config);
let mut all_modules = std::collections::HashSet::new();
for (adapter, _) in adapters {
for module in adapter.modules() {
all_modules.insert(module.clone());
}
}
for module_name in all_modules {
let mut merged_lora_a: Option<Tensor> = None;
let mut merged_lora_b: Option<Tensor> = None;
let mut total_weight = 0.0;
for (adapter, weight) in adapters {
if let Some(module_weights) = adapter.get_module(&module_name) {
let weighted_a = (&module_weights.lora_a * *weight)?;
let weighted_b = (&module_weights.lora_b * *weight)?;
if let Some(ref mut acc_a) = merged_lora_a {
*acc_a = (acc_a.clone() + weighted_a)?;
} else {
merged_lora_a = Some(weighted_a);
}
if let Some(ref mut acc_b) = merged_lora_b {
*acc_b = (acc_b.clone() + weighted_b)?;
} else {
merged_lora_b = Some(weighted_b);
}
total_weight += weight;
}
}
if let (Some(lora_a), Some(lora_b)) = (merged_lora_a, merged_lora_b) {
let normalized_a = (lora_a / total_weight)?;
let normalized_b = (lora_b / total_weight)?;
let merged_weights =
LoRAWeights::new(normalized_a, normalized_b, merged.config.scaling_factor());
merged.weights.insert(module_name, merged_weights);
}
}
Ok(merged)
}
}
pub struct LoRAModel {
pub base_tensors: HashMap<String, Tensor>,
pub adapter: LoRAAdapter,
pub is_merged: bool,
}
impl LoRAModel {
pub fn new(base_tensors: HashMap<String, Tensor>, adapter: LoRAAdapter) -> Self {
Self {
base_tensors,
adapter,
is_merged: false,
}
}
pub fn merge(&mut self) -> Result<()> {
if self.is_merged {
return Ok(()); }
for (module_name, lora_weights) in &self.adapter.weights {
if let Some(base_weight) = self.base_tensors.get_mut(module_name) {
let update = lora_weights.compute_update()?;
*base_weight = (base_weight.clone() + update)?;
}
}
self.is_merged = true;
Ok(())
}
pub fn unmerge(&mut self) -> Result<()> {
if !self.is_merged {
return Ok(()); }
for (module_name, lora_weights) in &self.adapter.weights {
if let Some(base_weight) = self.base_tensors.get_mut(module_name) {
let update = lora_weights.compute_update()?;
*base_weight = (base_weight.clone() - update)?;
}
}
self.is_merged = false;
Ok(())
}
pub fn get_weight(&self, module_name: &str) -> Result<Tensor> {
let base_weight = self
.base_tensors
.get(module_name)
.ok_or_else(|| Error::tensor_name_mapping(module_name.to_string()))?;
if self.is_merged {
Ok(base_weight.clone())
} else if let Some(lora_weights) = self.adapter.get_module(module_name) {
let update = lora_weights.compute_update()?;
Ok((base_weight.clone() + update)?)
} else {
Ok(base_weight.clone())
}
}
}
pub mod lora {
use super::*;
pub fn is_lora_adapter<P: AsRef<Path>>(path: P) -> bool {
let path = path.as_ref();
let config_file = path.join("adapter_config.json");
if !config_file.exists() {
return false;
}
if let Ok(config_data) = fs::read_to_string(&config_file) {
if let Ok(config) = serde_json::from_str::<serde_json::Value>(&config_data) {
if let Some(peft_type) = config.get("peft_type").and_then(|v| v.as_str()) {
return peft_type == "LORA";
}
}
}
false
}
pub fn load_config<P: AsRef<Path>>(path: P) -> Result<LoRAConfig> {
let config_file = path.as_ref().join("adapter_config.json");
let config_data = fs::read_to_string(&config_file)
.map_err(|e| Error::model_loading(format!("Failed to read adapter config: {}", e)))?;
let config: LoRAConfig = serde_json::from_str(&config_data)
.map_err(|e| Error::model_loading(format!("Failed to parse adapter config: {}", e)))?;
Ok(config)
}
pub fn load_adapter<P: AsRef<Path>>(
path: P,
device: &Device,
progress_callback: Option<ProgressFn>,
) -> Result<LoRAAdapter> {
let path = path.as_ref();
if let Some(ref progress) = progress_callback {
progress(ProgressEvent::Status {
message: format!("Loading LoRA adapter from {}", path.display()),
});
}
let config = load_config(path)?;
let mut adapter = LoRAAdapter::new(config.clone());
let model_file = if path.join("adapter_model.safetensors").exists() {
path.join("adapter_model.safetensors")
} else if path.join("adapter_model.bin").exists() {
return Err(Error::unsupported_format(
"PyTorch .bin files not yet supported for LoRA adapters".to_string(),
));
} else {
return Err(Error::model_loading(
"No adapter model file found (adapter_model.safetensors)".to_string(),
));
};
let all_tensors = candle_core::safetensors::load(&model_file, device)
.map_err(|e| Error::model_loading(format!("Failed to load SafeTensors: {}", e)))?;
let mut lora_modules: HashMap<String, (Option<Tensor>, Option<Tensor>)> = HashMap::new();
for (tensor_name, tensor) in all_tensors {
if let Some(lora_info) = parse_lora_tensor_name(&tensor_name) {
let module_entry = lora_modules
.entry(lora_info.module_name.clone())
.or_default();
match lora_info.matrix_type.as_str() {
"lora_A" => module_entry.0 = Some(tensor),
"lora_B" => module_entry.1 = Some(tensor),
_ => continue, }
}
}
for (module_name, (lora_a_opt, lora_b_opt)) in lora_modules {
if let (Some(lora_a), Some(lora_b)) = (lora_a_opt, lora_b_opt) {
let weights = LoRAWeights::new(lora_a, lora_b, config.scaling_factor());
adapter.add_module(module_name, weights)?;
}
}
if let Some(ref progress) = progress_callback {
progress(ProgressEvent::Status {
message: format!("Loaded LoRA adapter with {} modules", adapter.num_modules()),
});
}
Ok(adapter)
}
pub fn save_adapter<P: AsRef<Path>>(
adapter: &LoRAAdapter,
path: P,
progress_callback: Option<ProgressFn>,
) -> Result<()> {
let path = path.as_ref();
if !path.exists() {
fs::create_dir_all(path).map_err(|e| {
Error::io_error(format!("Failed to create adapter directory: {}", e))
})?;
}
if let Some(ref progress) = progress_callback {
progress(ProgressEvent::Status {
message: "Saving LoRA adapter configuration...".to_string(),
});
}
let config_file = path.join("adapter_config.json");
let config_data = serde_json::to_string_pretty(&adapter.config).map_err(|e| {
Error::model_saving(format!("Failed to serialize adapter config: {}", e))
})?;
fs::write(&config_file, config_data)
.map_err(|e| Error::io_error(format!("Failed to write adapter config: {}", e)))?;
if let Some(ref progress) = progress_callback {
progress(ProgressEvent::Status {
message: "Saving LoRA adapter weights...".to_string(),
});
}
let mut tensors_to_save = HashMap::new();
for (module_name, weights) in &adapter.weights {
let lora_a_name = format!("base_model.{}.lora_A.weight", module_name);
tensors_to_save.insert(lora_a_name, weights.lora_a.clone());
let lora_b_name = format!("base_model.{}.lora_B.weight", module_name);
tensors_to_save.insert(lora_b_name, weights.lora_b.clone());
if let Some(ref bias) = weights.lora_bias {
let bias_name = format!("base_model.{}.lora_bias", module_name);
tensors_to_save.insert(bias_name, bias.clone());
}
}
let mut metadata = HashMap::new();
metadata.insert("format".to_string(), "pt".to_string());
metadata.insert("peft_type".to_string(), "LORA".to_string());
for (key, value) in &adapter.metadata {
metadata.insert(key.clone(), value.clone());
}
let model_file = path.join("adapter_model.safetensors");
crate::formats::safetensors_export::save_safetensors_with_metadata(
&model_file,
&tensors_to_save,
&metadata,
)?;
if let Some(ref progress) = progress_callback {
progress(ProgressEvent::Status {
message: "LoRA adapter saved successfully".to_string(),
});
}
Ok(())
}
pub fn load_model_with_adapter<P1: AsRef<Path>, P2: AsRef<Path>>(
_base_model_path: P1,
adapter_path: P2,
device: &Device,
progress_callback: Option<ProgressFn>,
) -> Result<LoRAModel> {
if let Some(ref progress) = progress_callback {
progress(ProgressEvent::Status {
message: "Loading base model...".to_string(),
});
}
let base_tensors = HashMap::new();
let adapter = load_adapter(adapter_path, device, progress_callback)?;
Ok(LoRAModel::new(base_tensors, adapter))
}
#[derive(Debug, Clone)]
pub struct LoRATensorInfo {
pub module_name: String,
pub matrix_type: String,
}
pub fn parse_lora_tensor_name(tensor_name: &str) -> Option<LoRATensorInfo> {
if !tensor_name.contains("lora_") {
return None;
}
let parts: Vec<&str> = tensor_name.split('.').collect();
if parts.len() < 3 {
return None;
}
let mut lora_idx = None;
let mut matrix_type = None;
for (i, part) in parts.iter().enumerate() {
if part.starts_with("lora_") && (part == &"lora_A" || part == &"lora_B") {
lora_idx = Some(i);
matrix_type = Some(part[5..].to_string()); break;
}
}
if let (Some(idx), Some(mat_type)) = (lora_idx, matrix_type) {
let module_parts = &parts[..idx];
let module_name = module_parts.join(".");
let cleaned_module = if module_name.starts_with("base_model.") {
module_name
.strip_prefix("base_model.")
.unwrap_or(&module_name)
} else {
&module_name
};
Some(LoRATensorInfo {
module_name: cleaned_module.to_string(),
matrix_type: format!("lora_{}", mat_type),
})
} else {
None
}
}
pub mod advanced {
use super::*;
#[derive(Debug, Clone)]
pub struct CompositionOptions {
pub strategy: CompositionStrategy,
pub task_weights: HashMap<String, f64>,
pub normalize_weights: bool,
}
#[derive(Debug, Clone)]
pub enum CompositionStrategy {
WeightedSum,
Concatenation,
LearnedComposition {
gate_weights: HashMap<String, Tensor>,
},
}
impl Default for CompositionOptions {
fn default() -> Self {
Self {
strategy: CompositionStrategy::WeightedSum,
task_weights: HashMap::new(),
normalize_weights: true,
}
}
}
pub fn compose_adapters(
adapters: &[(String, LoRAAdapter, f64)], options: CompositionOptions,
) -> Result<LoRAAdapter> {
if adapters.is_empty() {
return Err(Error::invalid_config(
"Cannot compose empty list of adapters".to_string(),
));
}
let base_config = adapters[0].1.config.clone();
let mut composed = LoRAAdapter::new(base_config);
let mut all_modules = std::collections::HashSet::new();
for (_, adapter, _) in adapters {
for module in adapter.modules() {
all_modules.insert(module.clone());
}
}
match options.strategy {
CompositionStrategy::WeightedSum => {
compose_weighted_sum(adapters, &mut composed, all_modules, &options)?;
}
CompositionStrategy::Concatenation => {
compose_concatenation(adapters, &mut composed, all_modules)?;
}
CompositionStrategy::LearnedComposition { gate_weights: _ } => {
return Err(Error::unsupported_format(
"Learned composition not yet implemented".to_string(),
));
}
}
let adapter_names: Vec<String> =
adapters.iter().map(|(name, _, _)| name.clone()).collect();
composed.add_metadata("composed_from", adapter_names.join(","));
composed.add_metadata("composition_strategy", format!("{:?}", options.strategy));
Ok(composed)
}
fn compose_weighted_sum(
adapters: &[(String, LoRAAdapter, f64)],
composed: &mut LoRAAdapter,
all_modules: std::collections::HashSet<String>,
options: &CompositionOptions,
) -> Result<()> {
for module_name in all_modules {
let mut merged_lora_a: Option<Tensor> = None;
let mut merged_lora_b: Option<Tensor> = None;
let mut _total_weight = 0.0;
for (adapter_name, adapter, base_weight) in adapters {
if let Some(module_weights) = adapter.get_module(&module_name) {
let task_weight = options
.task_weights
.get(adapter_name)
.unwrap_or(base_weight);
let effective_weight = if options.normalize_weights {
*task_weight / adapters.len() as f64
} else {
*task_weight
};
let weighted_a = (&module_weights.lora_a * effective_weight)?;
let weighted_b = (&module_weights.lora_b * effective_weight)?;
merged_lora_a = Some(if let Some(existing) = merged_lora_a {
(&existing + weighted_a)?
} else {
weighted_a
});
merged_lora_b = Some(if let Some(existing) = merged_lora_b {
(&existing + weighted_b)?
} else {
weighted_b
});
_total_weight += effective_weight;
}
}
if let (Some(lora_a), Some(lora_b)) = (merged_lora_a, merged_lora_b) {
let weights =
LoRAWeights::new(lora_a, lora_b, composed.config.scaling_factor());
composed.weights.insert(module_name, weights);
}
}
Ok(())
}
fn compose_concatenation(
adapters: &[(String, LoRAAdapter, f64)],
composed: &mut LoRAAdapter,
all_modules: std::collections::HashSet<String>,
) -> Result<()> {
for module_name in all_modules {
let mut lora_a_tensors = Vec::new();
let mut lora_b_tensors = Vec::new();
for (_, adapter, _) in adapters {
if let Some(module_weights) = adapter.get_module(&module_name) {
lora_a_tensors.push(module_weights.lora_a.clone());
lora_b_tensors.push(module_weights.lora_b.clone());
}
}
if !lora_a_tensors.is_empty() {
let concatenated_a = Tensor::cat(&lora_a_tensors, 0)?; let concatenated_b = Tensor::cat(&lora_b_tensors, 1)?;
let weights = LoRAWeights::new(
concatenated_a,
concatenated_b,
composed.config.scaling_factor(),
);
composed.weights.insert(module_name, weights);
}
}
Ok(())
}
pub fn progressive_rank_expansion(
base_adapter: &LoRAAdapter,
target_rank: usize,
device: &Device,
) -> Result<LoRAAdapter> {
if target_rank <= base_adapter.config.r {
return Err(Error::invalid_config(
"Target rank must be larger than current rank".to_string(),
));
}
let mut expanded_config = base_adapter.config.clone();
expanded_config.r = target_rank;
let mut expanded = LoRAAdapter::new(expanded_config);
for (module_name, weights) in &base_adapter.weights {
let current_rank = base_adapter.config.r;
let rank_diff = target_rank - current_rank;
let (output_dim, input_dim) = weights.target_shape()?;
let zeros_a =
Tensor::zeros((rank_diff, input_dim), weights.lora_a.dtype(), device)?;
let expanded_a = Tensor::cat(&[weights.lora_a.clone(), zeros_a], 0)?;
let zeros_b =
Tensor::zeros((output_dim, rank_diff), weights.lora_b.dtype(), device)?;
let expanded_b = Tensor::cat(&[weights.lora_b.clone(), zeros_b], 1)?;
let expanded_weights = LoRAWeights::new(expanded_a, expanded_b, weights.scaling);
expanded
.weights
.insert(module_name.clone(), expanded_weights);
}
expanded.metadata = base_adapter.metadata.clone();
expanded.add_metadata("expanded_from_rank", base_adapter.config.r.to_string());
expanded.add_metadata("expanded_to_rank", target_rank.to_string());
Ok(expanded)
}
pub fn quantize_adapter(
adapter: &LoRAAdapter,
quantization_bits: u8,
) -> Result<LoRAAdapter> {
if quantization_bits != 8 && quantization_bits != 4 {
return Err(Error::unsupported_format(
"Only 4-bit and 8-bit quantization supported".to_string(),
));
}
let mut quantized = LoRAAdapter::new(adapter.config.clone());
for (module_name, weights) in &adapter.weights {
let quantized_a = quantize_tensor(&weights.lora_a, quantization_bits)?;
let quantized_b = quantize_tensor(&weights.lora_b, quantization_bits)?;
let quantized_weights = LoRAWeights::new(quantized_a, quantized_b, weights.scaling);
quantized
.weights
.insert(module_name.clone(), quantized_weights);
}
quantized.metadata = adapter.metadata.clone();
quantized.add_metadata("quantized", "true");
quantized.add_metadata("quantization_bits", quantization_bits.to_string());
Ok(quantized)
}
fn quantize_tensor(tensor: &Tensor, bits: u8) -> Result<Tensor> {
let max_val = tensor.abs()?.max_keepdim(0)?.max_keepdim(1)?;
let scale = (max_val / ((1 << (bits - 1)) - 1) as f64)?;
let divided = (tensor / &scale)?;
let quantized_int = divided.round()?;
let quantized = (&quantized_int * &scale)?;
Ok(quantized)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
#[test]
fn test_lora_config_creation() {
let config = LoRAConfig::new(16, 32.0)
.with_target_modules(vec![
"q_proj".to_string(),
"k_proj".to_string(),
"v_proj".to_string(),
])
.with_task_type("CAUSAL_LM");
assert_eq!(config.r, 16);
assert_eq!(config.lora_alpha, 32.0);
assert_eq!(config.scaling_factor(), 2.0); assert!(config.is_target_module("self_attn.q_proj"));
assert!(!config.is_target_module("layer_norm"));
}
#[test]
fn test_lora_tensor_name_parsing() {
let tensor_name = "base_model.model.layers.0.self_attn.q_proj.lora_A.weight";
let info = lora::parse_lora_tensor_name(tensor_name);
assert!(info.is_some());
let info = info.unwrap();
assert_eq!(info.module_name, "model.layers.0.self_attn.q_proj");
assert_eq!(info.matrix_type, "lora_A");
}
#[test]
fn test_lora_adapter_creation() {
let config = LoRAConfig::new(8, 16.0);
let mut adapter = LoRAAdapter::new(config);
assert_eq!(adapter.num_modules(), 0);
adapter.add_metadata("created_by", "test");
assert_eq!(
adapter.metadata.get("created_by"),
Some(&"test".to_string())
);
}
}