Skip to main content

entrenar/lora/adapter/
lora_adapter.rs

1//! LoRA adapter serialization and deserialization
2//!
3//! Contains the main LoRAAdapter struct for saving and loading adapters.
4
5use super::error::AdapterError;
6use super::metadata::AdapterMetadata;
7use crate::lora::LoRALayer;
8use crate::Tensor;
9use serde::{Deserialize, Serialize};
10use std::fs::File;
11use std::io::{BufReader, BufWriter};
12use std::path::Path;
13
14/// Serializable LoRA adapter format
15///
16/// Contains all information needed to reconstruct a LoRA adapter
17/// (excluding the base weight, which remains frozen and separate)
18#[derive(Serialize, Deserialize, Debug, Clone)]
19pub struct LoRAAdapter {
20    /// Format version for future compatibility
21    version: String,
22    /// LoRA rank
23    rank: usize,
24    /// LoRA alpha parameter
25    alpha: f32,
26    /// Output dimension
27    d_out: usize,
28    /// Input dimension
29    d_in: usize,
30    /// Computed scale factor (alpha/rank)
31    scale: f32,
32    /// LoRA A matrix weights [rank * d_in]
33    lora_a: Vec<f32>,
34    /// LoRA B matrix weights [d_out * rank]
35    lora_b: Vec<f32>,
36}
37
38impl LoRAAdapter {
39    /// Current adapter format version
40    const VERSION: &'static str = "1.0";
41
42    /// Create adapter from LoRALayer
43    ///
44    /// # Arguments
45    /// * `layer` - LoRALayer to extract adapter from
46    /// * `rank` - LoRA rank
47    /// * `alpha` - LoRA alpha parameter
48    pub fn from_layer(layer: &LoRALayer, rank: usize, alpha: f32) -> Self {
49        Self {
50            version: Self::VERSION.to_string(),
51            rank,
52            alpha,
53            d_out: layer.d_out(),
54            d_in: layer.d_in(),
55            scale: layer.scale(),
56            lora_a: layer.lora_a().data().to_vec(),
57            lora_b: layer.lora_b().data().to_vec(),
58        }
59    }
60
61    /// Load adapter and apply to base weight
62    ///
63    /// # Arguments
64    /// * `base_weight` - Frozen base weight tensor [d_out * d_in]
65    ///
66    /// # Returns
67    /// LoRALayer with loaded adapter weights
68    pub fn to_layer(&self, base_weight: Tensor) -> Result<LoRALayer, AdapterError> {
69        // Validate dimensions
70        if base_weight.len() != self.d_out * self.d_in {
71            return Err(AdapterError::DimensionMismatch {
72                expected: format!("{}x{} = {}", self.d_out, self.d_in, self.d_out * self.d_in),
73                actual: base_weight.len().to_string(),
74            });
75        }
76
77        if self.lora_a.len() != self.rank * self.d_in {
78            return Err(AdapterError::Validation(format!(
79                "LoRA A size mismatch: expected {} (rank {} * d_in {}), got {}",
80                self.rank * self.d_in,
81                self.rank,
82                self.d_in,
83                self.lora_a.len()
84            )));
85        }
86
87        if self.lora_b.len() != self.d_out * self.rank {
88            return Err(AdapterError::Validation(format!(
89                "LoRA B size mismatch: expected {} (d_out {} * rank {}), got {}",
90                self.d_out * self.rank,
91                self.d_out,
92                self.rank,
93                self.lora_b.len()
94            )));
95        }
96
97        // Create layer with loaded weights. `new` recomputes scale = alpha/rank
98        // (Standard mode); restore the SERIALIZED scale so a non-Standard adapter
99        // (e.g. rsLoRA, scale = alpha/sqrt(rank)) round-trips losslessly instead of
100        // being silently re-scaled to alpha/rank.
101        let mut layer = LoRALayer::new(base_weight, self.d_out, self.d_in, self.rank, self.alpha)
102            .with_scale(self.scale);
103
104        // Replace LoRA weights with loaded values
105        *layer.lora_a_mut().data_mut() = ndarray::arr1(&self.lora_a);
106        *layer.lora_b_mut().data_mut() = ndarray::arr1(&self.lora_b);
107
108        Ok(layer)
109    }
110
111    /// Save adapter to JSON file
112    ///
113    /// # Arguments
114    /// * `path` - File path to save to
115    pub fn save<P: AsRef<Path>>(&self, path: P) -> Result<(), AdapterError> {
116        let file = File::create(path)?;
117        let writer = BufWriter::new(file);
118        serde_json::to_writer_pretty(writer, self)?;
119        Ok(())
120    }
121
122    /// Load adapter from JSON file
123    ///
124    /// # Arguments
125    /// * `path` - File path to load from
126    pub fn load<P: AsRef<Path>>(path: P) -> Result<Self, AdapterError> {
127        let file = File::open(path)?;
128        let reader = BufReader::new(file);
129        let adapter: LoRAAdapter = serde_json::from_reader(reader)?;
130
131        // Validate version
132        if adapter.version != Self::VERSION {
133            return Err(AdapterError::Validation(format!(
134                "Unsupported adapter version: {} (expected {})",
135                adapter.version,
136                Self::VERSION
137            )));
138        }
139
140        Ok(adapter)
141    }
142
143    /// Get adapter metadata
144    pub fn metadata(&self) -> AdapterMetadata {
145        AdapterMetadata {
146            version: self.version.clone(),
147            rank: self.rank,
148            alpha: self.alpha,
149            d_out: self.d_out,
150            d_in: self.d_in,
151            scale: self.scale,
152            num_params: self.lora_a.len() + self.lora_b.len(),
153        }
154    }
155}
156
157#[cfg(test)]
158mod tests {
159    use super::*;
160    use tempfile::NamedTempFile;
161
162    fn make_test_adapter() -> LoRAAdapter {
163        LoRAAdapter {
164            version: "1.0".to_string(),
165            rank: 4,
166            alpha: 8.0,
167            d_out: 8,
168            d_in: 16,
169            scale: 2.0,
170            lora_a: vec![0.1; 4 * 16], // rank * d_in
171            lora_b: vec![0.2; 8 * 4],  // d_out * rank
172        }
173    }
174
175    #[test]
176    fn test_adapter_from_layer() {
177        let base_weight = Tensor::zeros(8 * 16, false);
178        let layer = LoRALayer::new(base_weight, 8, 16, 4, 8.0);
179        let adapter = LoRAAdapter::from_layer(&layer, 4, 8.0);
180        assert_eq!(adapter.rank, 4);
181        assert_eq!(adapter.alpha, 8.0);
182        assert_eq!(adapter.d_out, 8);
183        assert_eq!(adapter.d_in, 16);
184    }
185
186    #[test]
187    fn test_adapter_to_layer_valid() {
188        let adapter = make_test_adapter();
189        let base_weight = Tensor::zeros(8 * 16, false);
190        let layer = adapter.to_layer(base_weight).expect("operation should succeed");
191        assert_eq!(layer.d_out(), 8);
192        assert_eq!(layer.d_in(), 16);
193    }
194
195    /// FALSIFY-LORA-ADAPTER-SCALE-001: an rsLoRA adapter's scale (= alpha/sqrt(rank))
196    /// must survive a from_layer -> to_layer round-trip. `to_layer` previously rebuilt
197    /// the layer via `LoRALayer::new`, which recomputes Standard scale = alpha/rank and
198    /// discarded the serialized scale — silently re-scaling rsLoRA by a factor of
199    /// sqrt(rank) (e.g. 4x for rank=16) with no error, corrupting the reloaded model.
200    /// Every prior test used scale == alpha/rank (the Standard case the bug left intact).
201    #[test]
202    fn test_rslora_scale_survives_roundtrip() {
203        use crate::lora::LoRAScaling;
204        let (d_out, d_in, rank, alpha) = (4usize, 4usize, 16usize, 16.0f32);
205        let base = Tensor::zeros(d_out * d_in, false);
206        // rsLoRA scale = alpha / sqrt(rank) = 16 / 4 = 4.0  (Standard would be 16/16 = 1.0)
207        let layer = LoRALayer::new_with_scaling(
208            base.clone(),
209            d_out,
210            d_in,
211            rank,
212            alpha,
213            LoRAScaling::RsLoRA,
214        );
215        assert!(
216            (layer.scale() - 4.0).abs() < 1e-6,
217            "precondition: rsLoRA scale = {} (expected 4.0)",
218            layer.scale()
219        );
220
221        let adapter = LoRAAdapter::from_layer(&layer, rank, alpha);
222        assert!((adapter.metadata().scale - 4.0).abs() < 1e-6, "from_layer must capture 4.0");
223
224        let reloaded = adapter.to_layer(base).expect("to_layer should succeed");
225        // RED pre-fix: reloaded.scale() == alpha/rank == 1.0. GREEN post-fix: 4.0.
226        assert!(
227            (reloaded.scale() - 4.0).abs() < 1e-6,
228            "rsLoRA scale dropped on load: reloaded scale = {} (expected 4.0 = alpha/sqrt(rank))",
229            reloaded.scale()
230        );
231    }
232
233    #[test]
234    fn test_adapter_to_layer_dimension_mismatch() {
235        let adapter = make_test_adapter();
236        let base_weight = Tensor::zeros(100, false); // Wrong size
237        let result = adapter.to_layer(base_weight);
238        assert!(result.is_err());
239        match result {
240            Err(AdapterError::DimensionMismatch { .. }) => {}
241            _ => panic!("Expected DimensionMismatch error"),
242        }
243    }
244
245    #[test]
246    fn test_adapter_to_layer_lora_a_mismatch() {
247        let mut adapter = make_test_adapter();
248        adapter.lora_a = vec![0.1; 10]; // Wrong size
249        let base_weight = Tensor::zeros(8 * 16, false);
250        let result = adapter.to_layer(base_weight);
251        assert!(result.is_err());
252        match result {
253            Err(AdapterError::Validation(msg)) => {
254                assert!(msg.contains("LoRA A size mismatch"));
255            }
256            _ => panic!("Expected Validation error"),
257        }
258    }
259
260    #[test]
261    fn test_adapter_to_layer_lora_b_mismatch() {
262        let mut adapter = make_test_adapter();
263        adapter.lora_b = vec![0.2; 10]; // Wrong size
264        let base_weight = Tensor::zeros(8 * 16, false);
265        let result = adapter.to_layer(base_weight);
266        assert!(result.is_err());
267        match result {
268            Err(AdapterError::Validation(msg)) => {
269                assert!(msg.contains("LoRA B size mismatch"));
270            }
271            _ => panic!("Expected Validation error"),
272        }
273    }
274
275    #[test]
276    fn test_adapter_save_load_roundtrip() {
277        let adapter = make_test_adapter();
278        let file = NamedTempFile::new().expect("temp file creation should succeed");
279
280        adapter.save(file.path()).expect("save should succeed");
281        let loaded = LoRAAdapter::load(file.path()).expect("load should succeed");
282
283        assert_eq!(adapter.rank, loaded.rank);
284        assert_eq!(adapter.alpha, loaded.alpha);
285        assert_eq!(adapter.d_out, loaded.d_out);
286        assert_eq!(adapter.d_in, loaded.d_in);
287        assert_eq!(adapter.lora_a.len(), loaded.lora_a.len());
288        assert_eq!(adapter.lora_b.len(), loaded.lora_b.len());
289    }
290
291    #[test]
292    fn test_adapter_load_invalid_version() {
293        let mut adapter = make_test_adapter();
294        adapter.version = "0.0".to_string();
295        let file = NamedTempFile::new().expect("temp file creation should succeed");
296        adapter.save(file.path()).expect("save should succeed");
297
298        let result = LoRAAdapter::load(file.path());
299        assert!(result.is_err());
300        match result {
301            Err(AdapterError::Validation(msg)) => {
302                assert!(msg.contains("Unsupported adapter version"));
303            }
304            _ => panic!("Expected Validation error"),
305        }
306    }
307
308    #[test]
309    fn test_adapter_load_nonexistent_file() {
310        let result = LoRAAdapter::load("/nonexistent/path/adapter.json");
311        assert!(result.is_err());
312    }
313
314    #[test]
315    fn test_adapter_save_invalid_path() {
316        let adapter = make_test_adapter();
317        let result = adapter.save("/nonexistent/dir/adapter.json");
318        assert!(result.is_err());
319    }
320
321    #[test]
322    fn test_adapter_metadata() {
323        let adapter = make_test_adapter();
324        let meta = adapter.metadata();
325        assert_eq!(meta.rank, 4);
326        assert_eq!(meta.alpha, 8.0);
327        assert_eq!(meta.d_out, 8);
328        assert_eq!(meta.d_in, 16);
329        assert_eq!(meta.num_params, 4 * 16 + 8 * 4);
330    }
331
332    #[test]
333    fn test_adapter_clone() {
334        let adapter = make_test_adapter();
335        let cloned = adapter.clone();
336        assert_eq!(adapter.rank, cloned.rank);
337        assert_eq!(adapter.lora_a.len(), cloned.lora_a.len());
338    }
339
340    #[test]
341    fn test_adapter_debug() {
342        let adapter = make_test_adapter();
343        let debug = format!("{adapter:?}");
344        assert!(debug.contains("LoRAAdapter"));
345    }
346}