//! Tests for LoRA (Low-Rank Adaptation) functionality //! //! TDD: Define expected behavior for LoRA compression use anyhow::Result; use approx::assert_abs_diff_eq; use rtx_compress::lora::{LoRACompressor, LoRAConfig, LoRALayer}; use rtx_tensor::{Device, Tensor}; #[test] fn test_lora_config_creation() { let config = LoRAConfig { rank: 16, alpha: 32.0, dropout: 0.1, target_modules: vec![ "q_proj".to_string(), "v_proj".to_string(), "k_proj".to_string(), "o_proj".to_string(), ], merge_weights: false, }; assert_eq!(config.rank, 16); assert_eq!(config.alpha, 32.0); assert_eq!(config.scaling(), 32.0 / 16.0); // alpha / rank } #[test] fn test_lora_layer_creation() { let device = Device::cpu(); // Original weight dimensions let in_features = 768; let out_features = 768; let config = LoRAConfig { rank: 8, alpha: 16.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }; let lora_layer = LoRALayer::new(in_features, out_features, config, &device); assert!(lora_layer.is_ok()); let lora = lora_layer.unwrap(); // Check dimensions of LoRA matrices assert_eq!(lora.lora_a_shape(), vec![8, 768]); // [rank, in_features] assert_eq!(lora.lora_b_shape(), vec![768, 8]); // [out_features, rank] // Parameter reduction let original_params = in_features * out_features; // 768 * 768 = 589,824 let lora_params = lora.num_parameters(); // 8*768 + 768*8 = 12,288 assert_eq!(lora_params, 8 * 768 * 2); assert!(lora_params < original_params / 10); // Should be much smaller } #[test] fn test_lora_forward_pass() { let device = Device::cpu(); // Create original weight matrix let weight = Tensor::randn(&[256, 512], &device).unwrap(); // Create LoRA decomposition let config = LoRAConfig { rank: 4, alpha: 8.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }; let lora = LoRALayer::new(512, 256, config, &device).unwrap(); // Input let input = Tensor::randn(&[2, 512], &device).unwrap(); // Forward pass through original weight let output_original = input.matmul(&weight.transpose(0, 1).unwrap()).unwrap(); // Forward pass with LoRA let output_lora = lora.forward(&input, &weight).unwrap(); // Shapes should match assert_eq!(output_lora.shape().dims(), output_original.shape().dims()); assert_eq!(output_lora.shape().dims(), &[2, 256]); } #[test] fn test_lora_weight_merging() { let device = Device::cpu(); let weight = Tensor::randn(&[128, 128], &device).unwrap(); let config = LoRAConfig { rank: 8, alpha: 16.0, dropout: 0.0, target_modules: vec![], merge_weights: true, // Enable merging }; let mut lora = LoRALayer::new(128, 128, config, &device).unwrap(); // Initialize LoRA matrices with specific values lora.initialize_lora_weights().unwrap(); // Merge LoRA weights into original weight let merged = lora.merge_weights(&weight).unwrap(); // Check dimensions assert_eq!(merged.shape().dims(), weight.shape().dims()); // Unmerge should recover original weight let unmerged = lora.unmerge_weights(&merged).unwrap(); let weight_data = weight.to_vec().unwrap(); let unmerged_data = unmerged.to_vec().unwrap(); for i in 0..weight_data.len() { assert_abs_diff_eq!(weight_data[i], unmerged_data[i], epsilon = 1e-4); } } #[test] fn test_lora_compression_ratio() { let device = Device::cpu(); // Test compression for different matrix sizes and ranks let test_cases = vec![ (768, 768, 8), // Transformer attention (768, 3072, 16), // Transformer FFN (1024, 1024, 4), // Large model, very low rank ]; for (in_dim, out_dim, rank) in test_cases { let config = LoRAConfig { rank, alpha: rank as f32 * 2.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }; let lora = LoRALayer::new(in_dim, out_dim, config, &device).unwrap(); let original_params = in_dim * out_dim; let lora_params = lora.num_parameters(); let compression_ratio = lora.compression_ratio(); assert_eq!( compression_ratio, original_params as f32 / lora_params as f32 ); assert!(compression_ratio > 1.0); // Verify parameter count assert_eq!(lora_params, rank * (in_dim + out_dim)); } } #[test] fn test_lora_adapter_training() { let device = Device::cpu(); // Simulate fine-tuning scenario let base_weight = Tensor::randn(&[256, 256], &device).unwrap(); let config = LoRAConfig { rank: 8, alpha: 16.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }; let mut lora = LoRALayer::new(256, 256, config, &device).unwrap(); // Freeze base weight, only train LoRA parameters let trainable_params = lora.get_trainable_parameters(); assert_eq!(trainable_params.len(), 2); // lora_A and lora_B // Simulate gradient update let grad_a = Tensor::randn(&[8, 256], &device).unwrap(); let grad_b = Tensor::randn(&[256, 8], &device).unwrap(); lora.update_parameters(&grad_a, &grad_b, 0.001).unwrap(); // Parameters should have changed let new_params = lora.get_trainable_parameters(); assert_eq!(new_params.len(), 2); } #[test] fn test_multi_layer_lora_compression() { let device = Device::cpu(); // Compress multiple layers with LoRA let compressor = LoRACompressor::new( LoRAConfig { rank: 16, alpha: 32.0, dropout: 0.0, target_modules: vec![ "attention.q".to_string(), "attention.k".to_string(), "attention.v".to_string(), ], merge_weights: false, }, &device, ) .unwrap(); // Original model weights let weights = vec![ ( "attention.q".to_string(), Tensor::randn(&[768, 768], &device).unwrap(), ), ( "attention.k".to_string(), Tensor::randn(&[768, 768], &device).unwrap(), ), ( "attention.v".to_string(), Tensor::randn(&[768, 768], &device).unwrap(), ), ( "attention.o".to_string(), Tensor::randn(&[768, 768], &device).unwrap(), ), ( "ffn.w1".to_string(), Tensor::randn(&[3072, 768], &device).unwrap(), ), ( "ffn.w2".to_string(), Tensor::randn(&[768, 3072], &device).unwrap(), ), ]; // Apply LoRA compression let compressed = compressor.compress_model(&weights).unwrap(); // Check that only target modules are compressed assert_eq!(compressed.num_lora_layers(), 3); assert!(compressed.has_lora("attention.q")); assert!(compressed.has_lora("attention.k")); assert!(compressed.has_lora("attention.v")); assert!(!compressed.has_lora("attention.o")); // Calculate total compression let original_size = weights .iter() .filter(|(name, _)| compressor.is_target_module(name)) .map(|(_, w)| w.numel()) .sum::(); let compressed_size = compressed.total_lora_parameters(); let compression_ratio = original_size as f32 / compressed_size as f32; assert!(compression_ratio > 10.0); // Should achieve significant compression } #[test] fn test_lora_with_different_ranks() { let device = Device::cpu(); // Test that higher rank preserves more information let weight = Tensor::randn(&[512, 512], &device).unwrap(); let input = Tensor::randn(&[4, 512], &device).unwrap(); // Original output let output_original = input.matmul(&weight.transpose(0, 1).unwrap()).unwrap(); // Low rank LoRA let lora_low = LoRALayer::new( 512, 512, LoRAConfig { rank: 4, alpha: 8.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }, &device, ) .unwrap(); // High rank LoRA let lora_high = LoRALayer::new( 512, 512, LoRAConfig { rank: 32, alpha: 64.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }, &device, ) .unwrap(); let output_low = lora_low.forward(&input, &weight).unwrap(); let output_high = lora_high.forward(&input, &weight).unwrap(); // Both should have correct shape assert_eq!(output_low.shape().dims(), output_original.shape().dims()); assert_eq!(output_high.shape().dims(), output_original.shape().dims()); // Higher rank should have more parameters assert!(lora_high.num_parameters() > lora_low.num_parameters()); } #[test] fn test_lora_initialization_methods() { let device = Device::cpu(); let config = LoRAConfig { rank: 8, alpha: 16.0, dropout: 0.0, target_modules: vec![], merge_weights: false, }; // Test different initialization methods let mut lora = LoRALayer::new(256, 256, config, &device).unwrap(); // Kaiming initialization (default) lora.initialize_kaiming().unwrap(); let params_kaiming = lora.get_lora_a().to_vec().unwrap(); // Xavier initialization lora.initialize_xavier().unwrap(); let params_xavier = lora.get_lora_a().to_vec().unwrap(); // Zero initialization for lora_B lora.initialize_zero_b().unwrap(); let params_b = lora.get_lora_b().to_vec().unwrap(); // Check that B is zeros for val in params_b { assert_abs_diff_eq!(val, 0.0, epsilon = 1e-6); } // A matrices should be different between methods let mut differences = 0; for i in 0..params_kaiming.len() { if (params_kaiming[i] - params_xavier[i]).abs() > 1e-6 { differences += 1; } } assert!(differences > params_kaiming.len() / 2); // Most values should differ }