use entrenar_common::{EntrenarError, Result};
use std::path::Path;
#[derive(Debug)]
pub struct MergeEngine {
scale: f32,
}
impl Default for MergeEngine {
fn default() -> Self {
Self::new()
}
}
impl MergeEngine {
pub fn new() -> Self {
Self { scale: 1.0 }
}
pub fn with_scale(mut self, scale: f32) -> Self {
self.scale = scale;
self
}
pub fn merge(
&self,
base_weights: &[f32],
lora_a: &[f32],
lora_b: &[f32],
alpha: f32,
rank: u32,
) -> Vec<f32> {
let scale_factor = self.scale * alpha / rank as f32;
base_weights
.iter()
.enumerate()
.map(|(i, &w)| {
let a_val = lora_a.get(i % lora_a.len()).copied().unwrap_or(0.0);
let b_val = lora_b.get(i % lora_b.len()).copied().unwrap_or(0.0);
w + scale_factor * a_val * b_val
})
.collect()
}
pub fn merge_multiple(&self, base_weights: &[f32], adapters: &[AdapterWeights]) -> Vec<f32> {
let mut result = base_weights.to_vec();
for adapter in adapters {
let scale_factor = adapter.scale * adapter.alpha / adapter.rank as f32;
for (i, w) in result.iter_mut().enumerate() {
let a_val = adapter.lora_a.get(i % adapter.lora_a.len()).copied().unwrap_or(0.0);
let b_val = adapter.lora_b.get(i % adapter.lora_b.len()).copied().unwrap_or(0.0);
*w += scale_factor * a_val * b_val;
}
}
result
}
pub fn merge_from_file(
&self,
base_path: &Path,
adapter_path: &Path,
output_path: &Path,
) -> Result<MergeResult> {
if !base_path.exists() {
return Err(EntrenarError::ModelNotFound { path: base_path.to_path_buf() });
}
if !adapter_path.exists() {
return Err(EntrenarError::ModelNotFound { path: adapter_path.to_path_buf() });
}
Ok(MergeResult {
output_path: output_path.to_path_buf(),
merged_params: 0,
base_size_bytes: 0,
output_size_bytes: 0,
})
}
}
#[derive(Debug, Clone)]
pub struct AdapterWeights {
pub lora_a: Vec<f32>,
pub lora_b: Vec<f32>,
pub alpha: f32,
pub rank: u32,
pub scale: f32,
}
impl AdapterWeights {
pub fn new(lora_a: Vec<f32>, lora_b: Vec<f32>, alpha: f32, rank: u32) -> Self {
Self { lora_a, lora_b, alpha, rank, scale: 1.0 }
}
pub fn with_scale(mut self, scale: f32) -> Self {
self.scale = scale;
self
}
}
#[derive(Debug, Clone)]
pub struct MergeResult {
pub output_path: std::path::PathBuf,
pub merged_params: u64,
pub base_size_bytes: u64,
pub output_size_bytes: u64,
}
impl MergeResult {
pub fn size_increase_percent(&self) -> f64 {
if self.base_size_bytes == 0 {
return 0.0;
}
((self.output_size_bytes as f64 - self.base_size_bytes as f64)
/ self.base_size_bytes as f64)
* 100.0
}
}
#[derive(Debug, Clone)]
pub struct AdapterAnalysis {
pub rank: u32,
pub alpha: f32,
pub scale: f32,
pub effective_rank: f32,
pub rank_utilization: f64,
pub sparsity: f64,
pub frobenius_norm: f64,
}
pub fn analyze_adapter(lora_a: &[f32], lora_b: &[f32], alpha: f32, rank: u32) -> AdapterAnalysis {
let sparsity = calculate_sparsity(lora_a) * 0.5 + calculate_sparsity(lora_b) * 0.5;
let effective_rank = (rank as f32) * (1.0 - sparsity as f32);
let frobenius_norm = f64::from(
(lora_a.iter().map(|x| x * x).sum::<f32>() + lora_b.iter().map(|x| x * x).sum::<f32>())
.sqrt(),
);
AdapterAnalysis {
rank,
alpha,
scale: alpha / rank as f32,
effective_rank,
rank_utilization: f64::from(effective_rank / rank as f32) * 100.0,
sparsity: sparsity * 100.0,
frobenius_norm,
}
}
fn calculate_sparsity(values: &[f32]) -> f64 {
if values.is_empty() {
return 0.0;
}
let near_zero = values.iter().filter(|&&x| x.abs() < 1e-6).count();
near_zero as f64 / values.len() as f64
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_merge_adds_adapter_contribution() {
let engine = MergeEngine::new();
let base = vec![1.0, 2.0, 3.0, 4.0];
let lora_a = vec![0.1, 0.2];
let lora_b = vec![0.5, 0.5];
let merged = engine.merge(&base, &lora_a, &lora_b, 16.0, 64);
assert!(merged.iter().zip(&base).any(|(m, b)| (m - b).abs() > 1e-6));
}
#[test]
fn test_merge_scale_affects_result() {
let base = vec![1.0, 2.0, 3.0, 4.0];
let lora_a = vec![0.1, 0.2];
let lora_b = vec![0.5, 0.5];
let merged_1 = MergeEngine::new().with_scale(1.0).merge(&base, &lora_a, &lora_b, 16.0, 64);
let merged_2 = MergeEngine::new().with_scale(2.0).merge(&base, &lora_a, &lora_b, 16.0, 64);
let diff_1: f32 = merged_1.iter().zip(&base).map(|(m, b)| (m - b).abs()).sum();
let diff_2: f32 = merged_2.iter().zip(&base).map(|(m, b)| (m - b).abs()).sum();
assert!(diff_2 > diff_1);
}
#[test]
fn test_merge_multiple_adapters() {
let engine = MergeEngine::new();
let base = vec![1.0, 2.0, 3.0, 4.0];
let adapters = vec![
AdapterWeights::new(vec![0.1, 0.1], vec![0.5, 0.5], 16.0, 64),
AdapterWeights::new(vec![0.2, 0.2], vec![0.3, 0.3], 8.0, 32).with_scale(0.5),
];
let merged = engine.merge_multiple(&base, &adapters);
assert!(merged.iter().zip(&base).any(|(m, b)| (m - b).abs() > 1e-6));
}
#[test]
fn test_adapter_analysis() {
let lora_a = vec![0.1, 0.2, 0.3, 0.0, 0.0];
let lora_b = vec![0.5, 0.5, 0.0, 0.0, 0.5];
let analysis = analyze_adapter(&lora_a, &lora_b, 16.0, 64);
assert_eq!(analysis.rank, 64);
assert_eq!(analysis.alpha, 16.0);
assert!(analysis.sparsity > 0.0); assert!(analysis.frobenius_norm > 0.0);
}
#[test]
fn test_sparsity_calculation() {
let sparse = vec![0.0, 0.0, 0.0, 1.0];
assert!((calculate_sparsity(&sparse) - 0.75).abs() < 0.01);
let dense = vec![1.0, 2.0, 3.0, 4.0];
assert!((calculate_sparsity(&dense)).abs() < 0.01);
}
#[test]
fn test_merge_result_size_increase() {
let result = MergeResult {
output_path: std::path::PathBuf::from("/tmp/out"),
merged_params: 1000,
base_size_bytes: 1000,
output_size_bytes: 1100,
};
assert!((result.size_increase_percent() - 10.0).abs() < 0.01);
}
#[test]
fn test_merge_result_zero_base() {
let result = MergeResult {
output_path: std::path::PathBuf::from("/tmp/out"),
merged_params: 1000,
base_size_bytes: 0,
output_size_bytes: 1100,
};
assert_eq!(result.size_increase_percent(), 0.0);
}
#[test]
fn test_merge_engine_default() {
let engine = MergeEngine::default();
let base = vec![1.0, 2.0];
let lora_a = vec![1.0];
let lora_b = vec![1.0];
let merged = engine.merge(&base, &lora_a, &lora_b, 16.0, 16);
assert!((merged[0] - 2.0).abs() < 0.01);
}
#[test]
fn test_adapter_weights_with_scale() {
let adapter = AdapterWeights::new(vec![0.1], vec![0.2], 8.0, 32).with_scale(0.5);
assert_eq!(adapter.scale, 0.5);
assert_eq!(adapter.alpha, 8.0);
assert_eq!(adapter.rank, 32);
}
#[test]
fn test_merge_with_empty_adapters() {
let engine = MergeEngine::new();
let base = vec![1.0, 2.0, 3.0];
let adapters: Vec<AdapterWeights> = vec![];
let merged = engine.merge_multiple(&base, &adapters);
assert_eq!(merged, base);
}
#[test]
fn test_merge_from_file_missing_base() {
let engine = MergeEngine::new();
let result = engine.merge_from_file(
Path::new("/nonexistent/base.safetensors"),
Path::new("/nonexistent/adapter.safetensors"),
Path::new("/tmp/output.safetensors"),
);
assert!(result.is_err());
if let Err(EntrenarError::ModelNotFound { path }) = result {
assert!(path.to_string_lossy().contains("base.safetensors"));
}
}
#[test]
fn test_sparsity_empty_input() {
assert_eq!(calculate_sparsity(&[]), 0.0);
}
#[test]
fn test_sparsity_all_zeros() {
let zeros = vec![0.0, 0.0, 0.0, 0.0];
assert!((calculate_sparsity(&zeros) - 1.0).abs() < 0.01);
}
#[test]
fn test_adapter_analysis_effective_rank() {
let lora_a = vec![0.1, 0.2, 0.3, 0.4];
let lora_b = vec![0.5, 0.6, 0.7, 0.8];
let analysis = analyze_adapter(&lora_a, &lora_b, 16.0, 64);
assert!(analysis.effective_rank > 60.0);
assert!(analysis.rank_utilization > 90.0);
}
#[test]
fn test_adapter_analysis_scale_calculation() {
let lora_a = vec![0.1];
let lora_b = vec![0.1];
let analysis = analyze_adapter(&lora_a, &lora_b, 32.0, 64);
assert!((analysis.scale - 0.5).abs() < 0.01);
}
#[test]
fn test_merge_engine_with_scale_builder() {
let engine = MergeEngine::new().with_scale(0.75);
let base = vec![1.0, 1.0];
let lora_a = vec![1.0];
let lora_b = vec![1.0];
let merged = engine.merge(&base, &lora_a, &lora_b, 8.0, 8);
assert!((merged[0] - 1.75).abs() < 0.01);
}
}