1use 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#[derive(Serialize, Deserialize, Debug, Clone)]
19pub struct LoRAAdapter {
20 version: String,
22 rank: usize,
24 alpha: f32,
26 d_out: usize,
28 d_in: usize,
30 scale: f32,
32 lora_a: Vec<f32>,
34 lora_b: Vec<f32>,
36}
37
38impl LoRAAdapter {
39 const VERSION: &'static str = "1.0";
41
42 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 pub fn to_layer(&self, base_weight: Tensor) -> Result<LoRALayer, AdapterError> {
69 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 let mut layer = LoRALayer::new(base_weight, self.d_out, self.d_in, self.rank, self.alpha)
102 .with_scale(self.scale);
103
104 *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 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 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 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 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], lora_b: vec![0.2; 8 * 4], }
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 #[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 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 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); 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]; 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]; 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}