1use crate::autograd::Tensor;
4use crate::cli::logging::log;
5use crate::cli::LogLevel;
6use crate::config::{MergeArgs, MergeMethod};
7use crate::merge::{
8 dare_merge, ensemble_merge, slerp_merge, ties_merge, DareConfig, EnsembleConfig, Model,
9 SlerpConfig, TiesConfig,
10};
11use safetensors::SafeTensors;
12use std::collections::HashMap;
13use std::path::Path;
14
15pub fn run_merge(args: MergeArgs, level: LogLevel) -> Result<(), String> {
16 if args.method == MergeMethod::LoraAdapter {
18 return run_lora_adapter_merge(&args, level);
19 }
20
21 log_merge_start(&args, level);
22 validate_model_count(&args)?;
23
24 let models = load_all_models(&args.models, level)?;
25 let merged = perform_merge(&models, &args)?;
26 export_merged_model(&merged, &args)?;
27
28 log_merge_complete(&merged, &args, level);
29 Ok(())
30}
31
32fn log_merge_start(args: &MergeArgs, level: LogLevel) {
34 log(
35 level,
36 LogLevel::Normal,
37 &format!("Merging {} models using {:?}", args.models.len(), args.method),
38 );
39
40 for (i, model) in args.models.iter().enumerate() {
41 log(level, LogLevel::Verbose, &format!(" Model {}: {}", i + 1, model.display()));
42 }
43 log(level, LogLevel::Verbose, &format!(" Output: {}", args.output.display()));
44}
45
46fn validate_model_count(args: &MergeArgs) -> Result<(), String> {
48 if args.models.len() < 2 {
49 return Err("Need at least 2 models to merge".to_string());
50 }
51 Ok(())
52}
53
54fn load_all_models(paths: &[std::path::PathBuf], level: LogLevel) -> Result<Vec<Model>, String> {
56 let mut models: Vec<Model> = Vec::new();
57 for path in paths {
58 let model = load_single_model(path)?;
59 let tensor_count = model.len();
60 models.push(model);
61
62 log(
63 level,
64 LogLevel::Verbose,
65 &format!(" Loaded {} tensors from {}", tensor_count, path.display()),
66 );
67 }
68 Ok(models)
69}
70
71fn load_single_model(path: &Path) -> Result<Model, String> {
73 let data =
74 std::fs::read(path).map_err(|e| format!("Failed to read {}: {e}", path.display()))?;
75
76 let tensors = SafeTensors::deserialize(&data)
77 .map_err(|e| format!("Failed to parse {}: {e}", path.display()))?;
78
79 let mut model: Model = HashMap::new();
80 for name in tensors.names() {
81 if let Some(tensor) = extract_f32_tensor(&tensors, name)? {
82 model.insert((*name).to_string(), tensor);
83 }
84 }
85 Ok(model)
86}
87
88fn extract_f32_tensor(tensors: &SafeTensors<'_>, name: &str) -> Result<Option<Tensor>, String> {
90 let tensor = tensors.tensor(name).map_err(|e| format!("Failed to get tensor {name}: {e}"))?;
91
92 if tensor.dtype() != safetensors::tensor::Dtype::F32 {
93 return Ok(None);
94 }
95
96 let bytes = tensor.data();
97 let values: Vec<f32> = bytes
98 .chunks_exact(4)
99 .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
100 .collect();
101
102 Ok(Some(Tensor::from_vec(values, false)))
103}
104
105fn perform_merge(models: &[Model], args: &MergeArgs) -> Result<Model, String> {
107 match args.method {
108 MergeMethod::Ties => perform_ties_merge(models, args),
109 MergeMethod::Dare => perform_dare_merge(models, args),
110 MergeMethod::Slerp => perform_slerp_merge(models, args),
111 MergeMethod::Average => perform_average_merge(models, args),
112 MergeMethod::LoraAdapter => {
113 Err("LoRA adapter merge uses dedicated path".to_string())
115 }
116 }
117}
118
119fn perform_ties_merge(models: &[Model], args: &MergeArgs) -> Result<Model, String> {
121 let config = TiesConfig { density: args.density.unwrap_or(0.2) };
122 let base = &models[0];
123 ties_merge(models.get(1..).unwrap_or_default(), base, &config)
124 .map_err(|e| format!("TIES merge failed: {e}"))
125}
126
127fn perform_dare_merge(models: &[Model], args: &MergeArgs) -> Result<Model, String> {
129 let config = DareConfig { drop_prob: 1.0 - args.density.unwrap_or(0.5), seed: None };
130 let base = &models[0];
131 dare_merge(models.get(1..).unwrap_or_default(), base, &config)
132 .map_err(|e| format!("DARE merge failed: {e}"))
133}
134
135fn perform_slerp_merge(models: &[Model], args: &MergeArgs) -> Result<Model, String> {
137 if models.len() != 2 {
138 return Err("SLERP requires exactly 2 models".to_string());
139 }
140 let config = SlerpConfig { t: args.weight.unwrap_or(0.5) };
141 slerp_merge(&models[0], &models[1], &config).map_err(|e| format!("SLERP merge failed: {e}"))
142}
143
144fn perform_average_merge(models: &[Model], args: &MergeArgs) -> Result<Model, String> {
146 let config = build_ensemble_config(args)?;
147 ensemble_merge(models, &config).map_err(|e| format!("Average merge failed: {e}"))
148}
149
150fn build_ensemble_config(args: &MergeArgs) -> Result<EnsembleConfig, String> {
152 if let Some(w_str) = &args.weights {
153 let weights: Vec<f32> = w_str
154 .split(',')
155 .map(|s| s.trim().parse::<f32>())
156 .collect::<Result<Vec<_>, _>>()
157 .map_err(|e| format!("Invalid weights: {e}"))?;
158 Ok(EnsembleConfig::weighted_average(weights))
159 } else {
160 Ok(EnsembleConfig::uniform_average())
161 }
162}
163
164fn export_merged_model(merged: &Model, args: &MergeArgs) -> Result<(), String> {
166 let output_ext = args.output.extension().and_then(|s| s.to_str()).unwrap_or("json");
167
168 if output_ext == "safetensors" {
169 export_safetensors(merged, args)
170 } else {
171 export_json(merged, args)
172 }
173}
174
175fn export_safetensors(merged: &Model, args: &MergeArgs) -> Result<(), String> {
177 use safetensors::tensor::{Dtype, TensorView};
178
179 let tensor_data: Vec<(String, Vec<u8>, Vec<usize>)> = merged
180 .iter()
181 .map(|(name, tensor)| {
182 let data = tensor.data();
183 let bytes: Vec<u8> = bytemuck::cast_slice(data.as_slice().unwrap_or(&[])).to_vec();
184 let shape = vec![tensor.len()];
185 (name.clone(), bytes, shape)
186 })
187 .collect();
188
189 let views: Vec<(&str, TensorView<'_>)> = tensor_data
190 .iter()
191 .filter_map(|(name, bytes, shape)| {
192 TensorView::new(Dtype::F32, shape.clone(), bytes).ok().map(|view| (name.as_str(), view))
193 })
194 .collect();
195
196 let metadata = build_safetensor_metadata(merged, args);
197 let safetensor_bytes = safetensors::serialize(views, Some(metadata))
198 .map_err(|e| format!("Failed to serialize SafeTensors: {e}"))?;
199
200 std::fs::write(&args.output, safetensor_bytes)
201 .map_err(|e| format!("Failed to write output: {e}"))
202}
203
204fn build_safetensor_metadata(merged: &Model, args: &MergeArgs) -> HashMap<String, String> {
206 let mut metadata = HashMap::new();
207 metadata.insert("name".to_string(), "merged-model".to_string());
208 metadata.insert("merge_method".to_string(), format!("{:?}", args.method));
209 metadata.insert("tensor_count".to_string(), merged.len().to_string());
210 metadata
211}
212
213fn export_json(merged: &Model, args: &MergeArgs) -> Result<(), String> {
215 let output_data: HashMap<String, Vec<f32>> =
216 merged.iter().map(|(name, tensor)| (name.clone(), tensor.data().to_vec())).collect();
217
218 let json_data =
219 serde_json::to_vec_pretty(&output_data).map_err(|e| format!("Failed to serialize: {e}"))?;
220
221 std::fs::write(&args.output, &json_data).map_err(|e| format!("Failed to write output: {e}"))
222}
223
224fn log_merge_complete(merged: &Model, args: &MergeArgs, level: LogLevel) {
226 log(
227 level,
228 LogLevel::Normal,
229 &format!("Merge complete: {} tensors written to {}", merged.len(), args.output.display()),
230 );
231}
232
233fn run_lora_adapter_merge(args: &MergeArgs, level: LogLevel) -> Result<(), String> {
238 let base_path = args.base.as_ref().ok_or("--base required for lora-adapter merge")?;
239 let adapter_dir = args.adapter.as_ref().ok_or("--adapter required for lora-adapter merge")?;
240
241 let config_path = adapter_dir.join("adapter_config.json");
242 let adapter_path = adapter_dir.join("adapter_model.safetensors");
243
244 if !base_path.exists() {
245 return Err(format!("Base model not found: {}", base_path.display()));
246 }
247 if !config_path.exists() {
248 return Err(format!("adapter_config.json not found in {}", adapter_dir.display()));
249 }
250 if !adapter_path.exists() {
251 return Err(format!("adapter_model.safetensors not found in {}", adapter_dir.display()));
252 }
253
254 log(level, LogLevel::Normal, "LoRA adapter merge:");
255 log(level, LogLevel::Normal, &format!(" Base: {}", base_path.display()));
256 log(level, LogLevel::Normal, &format!(" Adapter: {}", adapter_dir.display()));
257
258 let config_str =
260 std::fs::read_to_string(&config_path).map_err(|e| format!("Read adapter config: {e}"))?;
261 let config: serde_json::Value =
262 serde_json::from_str(&config_str).map_err(|e| format!("Parse adapter config: {e}"))?;
263
264 let rank = config.get("r").and_then(serde_json::Value::as_u64).unwrap_or(8) as usize;
265 let alpha =
266 config.get("lora_alpha").and_then(serde_json::Value::as_f64).unwrap_or(rank as f64 * 2.0);
267 let use_rslora = config.get("use_rslora").and_then(serde_json::Value::as_bool).unwrap_or(false);
271 let scaling = if use_rslora {
272 crate::lora::LoRAScaling::RsLoRA
273 } else {
274 crate::lora::LoRAScaling::Standard
275 };
276 let scale = scaling.compute(alpha as f32, rank);
277
278 log(
279 level,
280 LogLevel::Normal,
281 &format!(" Rank: {rank}, Alpha: {alpha}, rsLoRA: {use_rslora}, Scale: {scale:.4}"),
282 );
283
284 let base_data = std::fs::read(base_path).map_err(|e| format!("Read base model: {e}"))?;
286 let base_tensors =
287 SafeTensors::deserialize(&base_data).map_err(|e| format!("Parse base model: {e}"))?;
288
289 let adapter_data = std::fs::read(&adapter_path).map_err(|e| format!("Read adapter: {e}"))?;
291 let adapter_tensors =
292 SafeTensors::deserialize(&adapter_data).map_err(|e| format!("Parse adapter: {e}"))?;
293
294 let adapter_names: Vec<String> =
296 adapter_tensors.names().iter().map(|s| (*s).to_string()).collect();
297 let base_names: Vec<String> = base_tensors.names().iter().map(|s| (*s).to_string()).collect();
298
299 let lora_pairs = build_lora_pairs(&adapter_names, &adapter_tensors)?;
301 let mut merged_count = 0usize;
302
303 let mut output_tensors: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
305
306 for name in &base_names {
307 let base_t = base_tensors.tensor(name).map_err(|e| format!("Get tensor {name}: {e}"))?;
308 let shape: Vec<usize> = base_t.shape().to_vec();
309
310 if let Some(((a_data, a_shape, a_dtype), (b_data, b_shape, b_dtype))) =
312 lora_pairs.get(name.as_str())
313 {
314 let base_f32 = bytes_to_f32(base_t.data(), base_t.dtype());
318 let a_f32 = bytes_to_f32(a_data, *a_dtype);
319 let b_f32 = bytes_to_f32(b_data, *b_dtype);
320
321 let d_out = b_shape[0];
322 let r = b_shape[1];
323 let d_in = a_shape[1];
324
325 let mut ba = vec![0.0f32; d_out * d_in];
327 for i in 0..d_out {
328 for j in 0..d_in {
329 let mut sum = 0.0f32;
330 for k in 0..r {
331 sum += b_f32[i * r + k] * a_f32[k * d_in + j];
332 }
333 ba[i * d_in + j] = sum;
334 }
335 }
336
337 let mut merged: Vec<f32> = base_f32;
339 for (i, val) in merged.iter_mut().enumerate() {
340 *val += scale * ba[i];
341 }
342
343 let bytes: Vec<u8> = bytemuck::cast_slice(&merged).to_vec();
344 output_tensors.push((name.clone(), bytes, shape));
345 merged_count += 1;
346 } else {
347 output_tensors.push((name.clone(), base_t.data().to_vec(), shape));
349 }
350 }
351
352 let views: Vec<(&str, safetensors::tensor::TensorView<'_>)> = output_tensors
354 .iter()
355 .filter_map(|(name, bytes, shape)| {
356 safetensors::tensor::TensorView::new(
357 safetensors::tensor::Dtype::F32,
358 shape.clone(),
359 bytes,
360 )
361 .ok()
362 .map(|view| (name.as_str(), view))
363 })
364 .collect();
365
366 let mut metadata = HashMap::new();
367 metadata.insert("format".to_string(), "entrenar-merged-lora".to_string());
368 metadata.insert("lora_rank".to_string(), rank.to_string());
369 metadata.insert("lora_alpha".to_string(), format!("{alpha}"));
370
371 let safetensor_bytes = safetensors::serialize(views, Some(metadata))
372 .map_err(|e| format!("Serialize merged model: {e}"))?;
373
374 std::fs::write(&args.output, safetensor_bytes)
375 .map_err(|e| format!("Write merged model: {e}"))?;
376
377 let output_size = std::fs::metadata(&args.output).map(|m| m.len()).unwrap_or(0);
378 log(
379 level,
380 LogLevel::Normal,
381 &format!(" Merged {merged_count} adapter weights into base model"),
382 );
383 log(
384 level,
385 LogLevel::Normal,
386 &format!(" Output: {} ({:.2} MB)", args.output.display(), output_size as f64 / 1e6),
387 );
388
389 Ok(())
390}
391
392type LoraFactor = (Vec<u8>, Vec<usize>, safetensors::tensor::Dtype);
397
398type LoraPair = (LoraFactor, LoraFactor);
400
401fn build_lora_pairs<'a>(
403 names: &[String],
404 tensors: &'a SafeTensors<'a>,
405) -> Result<HashMap<&'a str, LoraPair>, String> {
406 let mut pairs: HashMap<String, (Option<LoraFactor>, Option<LoraFactor>)> = HashMap::new();
407
408 for name in names {
409 let (base_name, is_a) = if let Some(stripped) = name.strip_suffix(".lora_A.weight") {
411 (stripped.replace("base_model.model.", "") + ".weight", true)
412 } else if let Some(stripped) = name.strip_suffix(".lora_B.weight") {
413 (stripped.replace("base_model.model.", "") + ".weight", false)
414 } else {
415 continue;
416 };
417
418 let tensor = tensors.tensor(name).map_err(|e| format!("Get adapter tensor {name}: {e}"))?;
419 let data = tensor.data().to_vec();
420 let shape = tensor.shape().to_vec();
421 let dtype = tensor.dtype();
422
423 let entry = pairs.entry(base_name).or_insert((None, None));
424 if is_a {
425 entry.0 = Some((data, shape, dtype));
426 } else {
427 entry.1 = Some((data, shape, dtype));
428 }
429 }
430
431 let mut result = HashMap::new();
432 for (base_name, (a, b)) in &pairs {
433 if let (Some(a_factor), Some(b_factor)) = (a, b) {
434 let key: &str = Box::leak(base_name.clone().into_boxed_str());
437 result.insert(key, (a_factor.clone(), b_factor.clone()));
438 }
439 }
440 Ok(result)
441}
442
443fn bytes_to_f32(data: &[u8], dtype: safetensors::tensor::Dtype) -> Vec<f32> {
445 match dtype {
446 safetensors::tensor::Dtype::F32 => {
447 data.chunks_exact(4).map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]])).collect()
448 }
449 safetensors::tensor::Dtype::F16 => data
450 .chunks_exact(2)
451 .map(|c| {
452 let bits = u16::from_le_bytes([c[0], c[1]]);
453 half::f16::from_bits(bits).to_f32()
454 })
455 .collect(),
456 safetensors::tensor::Dtype::BF16 => data
457 .chunks_exact(2)
458 .map(|c| {
459 let bits = u16::from_le_bytes([c[0], c[1]]);
460 half::bf16::from_bits(bits).to_f32()
461 })
462 .collect(),
463 _ => {
464 data.chunks_exact(4).map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]])).collect()
466 }
467 }
468}
469
470#[cfg(test)]
471mod tests {
472 #![allow(clippy::unwrap_used)]
473 use super::*;
474 use std::path::PathBuf;
475
476 #[test]
477 fn test_validate_model_count_zero() {
478 let args = MergeArgs {
479 models: vec![],
480 output: PathBuf::from("o.json"),
481 method: MergeMethod::Ties,
482 weight: None,
483 density: None,
484 weights: None,
485 base: None,
486 adapter: None,
487 };
488 assert!(validate_model_count(&args).is_err());
489 }
490
491 #[test]
492 fn test_validate_model_count_two_ok() {
493 let args = MergeArgs {
494 models: vec![PathBuf::from("a"), PathBuf::from("b")],
495 output: PathBuf::from("o.json"),
496 method: MergeMethod::Ties,
497 weight: None,
498 density: None,
499 weights: None,
500 base: None,
501 adapter: None,
502 };
503 assert!(validate_model_count(&args).is_ok());
504 }
505
506 #[test]
507 fn test_build_ensemble_config_no_weights() {
508 let args = MergeArgs {
509 models: vec![],
510 output: PathBuf::from("o.json"),
511 method: MergeMethod::Average,
512 weight: None,
513 density: None,
514 weights: None,
515 base: None,
516 adapter: None,
517 };
518 assert!(build_ensemble_config(&args).is_ok());
519 }
520
521 #[test]
522 fn test_build_ensemble_config_with_weights() {
523 let args = MergeArgs {
524 models: vec![],
525 output: PathBuf::from("o.json"),
526 method: MergeMethod::Average,
527 weight: None,
528 density: None,
529 weights: Some("0.3, 0.7".into()),
530 base: None,
531 adapter: None,
532 };
533 assert!(build_ensemble_config(&args).is_ok());
534 }
535
536 #[test]
537 fn test_build_ensemble_config_invalid() {
538 let args = MergeArgs {
539 models: vec![],
540 output: PathBuf::from("o.json"),
541 method: MergeMethod::Average,
542 weight: None,
543 density: None,
544 weights: Some("abc".into()),
545 base: None,
546 adapter: None,
547 };
548 assert!(build_ensemble_config(&args).unwrap_err().contains("Invalid weights"));
549 }
550
551 fn mk(keys: &[(&str, &[f32])]) -> Model {
552 keys.iter().map(|(n, v)| (n.to_string(), Tensor::from_vec(v.to_vec(), false))).collect()
553 }
554
555 #[test]
556 fn test_slerp_wrong_count() {
557 let ms = vec![mk(&[("w", &[1.0])]), mk(&[("w", &[2.0])]), mk(&[("w", &[3.0])])];
558 let a = MergeArgs {
559 models: vec![],
560 output: PathBuf::from("o"),
561 method: MergeMethod::Slerp,
562 weight: None,
563 density: None,
564 weights: None,
565 base: None,
566 adapter: None,
567 };
568 assert!(perform_slerp_merge(&ms, &a).unwrap_err().contains("SLERP requires exactly 2"));
569 }
570
571 #[test]
572 fn test_merge_lora_err() {
573 let a = MergeArgs {
574 models: vec![],
575 output: PathBuf::from("o"),
576 method: MergeMethod::LoraAdapter,
577 weight: None,
578 density: None,
579 weights: None,
580 base: None,
581 adapter: None,
582 };
583 assert!(perform_merge(&[], &a).is_err());
584 }
585
586 #[test]
587 fn test_bytes_to_f32_f32() {
588 let v = vec![1.0f32, 2.5];
589 let b: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
590 let r = bytes_to_f32(&b, safetensors::tensor::Dtype::F32);
591 assert!((r[0] - 1.0).abs() < 1e-6);
592 }
593
594 #[test]
595 fn test_bytes_to_f32_f16() {
596 let b = half::f16::from_f32(1.0).to_le_bytes().to_vec();
597 assert!((bytes_to_f32(&b, safetensors::tensor::Dtype::F16)[0] - 1.0).abs() < 0.01);
598 }
599
600 #[test]
601 fn test_bytes_to_f32_bf16() {
602 let b = half::bf16::from_f32(2.0).to_le_bytes().to_vec();
603 assert!((bytes_to_f32(&b, safetensors::tensor::Dtype::BF16)[0] - 2.0).abs() < 0.1);
604 }
605
606 #[test]
607 fn test_bytes_to_f32_fallback() {
608 let b: Vec<u8> = 42.0f32.to_le_bytes().to_vec();
609 assert!((bytes_to_f32(&b, safetensors::tensor::Dtype::I8)[0] - 42.0).abs() < 1e-6);
610 }
611
612 #[test]
613 fn test_bytes_to_f32_empty() {
614 assert!(bytes_to_f32(&[], safetensors::tensor::Dtype::F32).is_empty());
615 }
616
617 #[test]
618 fn test_safetensor_metadata() {
619 let m = mk(&[("a", &[1.0]), ("b", &[2.0])]);
620 let a = MergeArgs {
621 models: vec![],
622 output: PathBuf::from("o.st"),
623 method: MergeMethod::Dare,
624 weight: None,
625 density: None,
626 weights: None,
627 base: None,
628 adapter: None,
629 };
630 let md = build_safetensor_metadata(&m, &a);
631 assert_eq!(md["name"], "merged-model");
632 assert_eq!(md["tensor_count"], "2");
633 }
634
635 #[test]
636 fn test_export_json() {
637 let m = mk(&[("w", &[1.0])]);
638 let t = std::env::temp_dir().join("ent_merge_j.json");
639 let a = MergeArgs {
640 models: vec![],
641 output: t.clone(),
642 method: MergeMethod::Average,
643 weight: None,
644 density: None,
645 weights: None,
646 base: None,
647 adapter: None,
648 };
649 assert!(export_merged_model(&m, &a).is_ok());
650 let _ = std::fs::remove_file(&t);
651 }
652
653 #[test]
654 fn test_export_safetensors() {
655 let m = mk(&[("w", &[1.0])]);
656 let t = std::env::temp_dir().join("ent_merge_s.safetensors");
657 let a = MergeArgs {
658 models: vec![],
659 output: t.clone(),
660 method: MergeMethod::Average,
661 weight: None,
662 density: None,
663 weights: None,
664 base: None,
665 adapter: None,
666 };
667 assert!(export_merged_model(&m, &a).is_ok());
668 let _ = std::fs::remove_file(&t);
669 }
670
671 #[test]
672 fn test_ties_merge_ok() {
673 let a = MergeArgs {
674 models: vec![],
675 output: PathBuf::from("o"),
676 method: MergeMethod::Ties,
677 weight: None,
678 density: None,
679 weights: None,
680 base: None,
681 adapter: None,
682 };
683 assert!(perform_ties_merge(
685 &[mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.1, 2.1])]), mk(&[("w", &[1.2, 2.2])]),],
686 &a
687 )
688 .is_ok());
689 }
690
691 #[test]
692 fn test_dare_merge_ok() {
693 let a = MergeArgs {
694 models: vec![],
695 output: PathBuf::from("o"),
696 method: MergeMethod::Dare,
697 weight: None,
698 density: None,
699 weights: None,
700 base: None,
701 adapter: None,
702 };
703 assert!(
704 perform_dare_merge(&[mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.1, 2.1])])], &a).is_ok()
705 );
706 }
707
708 #[test]
709 fn test_average_merge() {
710 let a = MergeArgs {
711 models: vec![],
712 output: PathBuf::from("o"),
713 method: MergeMethod::Average,
714 weight: None,
715 density: None,
716 weights: None,
717 base: None,
718 adapter: None,
719 };
720 let r = perform_average_merge(&[mk(&[("w", &[2.0, 4.0])]), mk(&[("w", &[6.0, 8.0])])], &a)
721 .unwrap();
722 let s = r["w"].data().as_slice().unwrap().to_vec();
723 assert!((s[0] - 4.0).abs() < 1e-6);
724 }
725
726 #[test]
727 fn test_log_merge_no_panic() {
728 let a = MergeArgs {
729 models: vec![PathBuf::from("a"), PathBuf::from("b")],
730 output: PathBuf::from("o"),
731 method: MergeMethod::Ties,
732 weight: None,
733 density: None,
734 weights: None,
735 base: None,
736 adapter: None,
737 };
738 log_merge_start(&a, LogLevel::Quiet);
739 log_merge_start(&a, LogLevel::Verbose);
740 log_merge_complete(&mk(&[("w", &[1.0])]), &a, LogLevel::Normal);
741 }
742
743 #[test]
744 fn test_lora_missing_base() {
745 let a = MergeArgs {
746 models: vec![],
747 output: PathBuf::from("o"),
748 method: MergeMethod::LoraAdapter,
749 weight: None,
750 density: None,
751 weights: None,
752 base: None,
753 adapter: Some(PathBuf::from("/tmp")),
754 };
755 assert!(run_lora_adapter_merge(&a, LogLevel::Quiet)
756 .unwrap_err()
757 .contains("--base required"));
758 }
759
760 #[test]
761 fn test_lora_missing_adapter() {
762 let a = MergeArgs {
763 models: vec![],
764 output: PathBuf::from("o"),
765 method: MergeMethod::LoraAdapter,
766 weight: None,
767 density: None,
768 weights: None,
769 base: Some(PathBuf::from("/tmp/x")),
770 adapter: None,
771 };
772 assert!(run_lora_adapter_merge(&a, LogLevel::Quiet)
773 .unwrap_err()
774 .contains("--adapter required"));
775 }
776
777 #[test]
778 fn test_lora_base_not_found() {
779 let a = MergeArgs {
780 models: vec![],
781 output: PathBuf::from("o"),
782 method: MergeMethod::LoraAdapter,
783 weight: None,
784 density: None,
785 weights: None,
786 base: Some(PathBuf::from("/no/base")),
787 adapter: Some(PathBuf::from("/tmp")),
788 };
789 assert!(run_lora_adapter_merge(&a, LogLevel::Quiet)
790 .unwrap_err()
791 .contains("Base model not found"));
792 }
793
794 #[test]
795 fn test_load_nonexistent() {
796 assert!(load_single_model(std::path::Path::new("/no/m"))
797 .unwrap_err()
798 .contains("Failed to read"));
799 }
800
801 #[test]
802 fn test_run_merge_too_few() {
803 let a = MergeArgs {
804 models: vec![PathBuf::from("a")],
805 output: PathBuf::from("o"),
806 method: MergeMethod::Ties,
807 weight: None,
808 density: None,
809 weights: None,
810 base: None,
811 adapter: None,
812 };
813 assert!(run_merge(a, LogLevel::Quiet).unwrap_err().contains("Need at least 2"));
814 }
815
816 #[test]
817 fn test_run_merge_lora_routes() {
818 let a = MergeArgs {
819 models: vec![],
820 output: PathBuf::from("o"),
821 method: MergeMethod::LoraAdapter,
822 weight: None,
823 density: None,
824 weights: None,
825 base: None,
826 adapter: None,
827 };
828 assert!(run_merge(a, LogLevel::Quiet).unwrap_err().contains("--base required"));
829 }
830
831 #[test]
834 fn test_perform_merge_ties_route() {
835 let models =
836 vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.1, 2.1])]), mk(&[("w", &[1.2, 2.2])])];
837 let a = MergeArgs {
838 models: vec![],
839 output: PathBuf::from("o"),
840 method: MergeMethod::Ties,
841 weight: None,
842 density: Some(0.5),
843 weights: None,
844 base: None,
845 adapter: None,
846 };
847 assert!(perform_merge(&models, &a).is_ok());
848 }
849
850 #[test]
851 fn test_perform_merge_dare_route() {
852 let models = vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])])];
853 let a = MergeArgs {
854 models: vec![],
855 output: PathBuf::from("o"),
856 method: MergeMethod::Dare,
857 weight: None,
858 density: Some(0.3),
859 weights: None,
860 base: None,
861 adapter: None,
862 };
863 assert!(perform_merge(&models, &a).is_ok());
864 }
865
866 #[test]
867 fn test_perform_merge_slerp_route() {
868 let models = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
869 let a = MergeArgs {
870 models: vec![],
871 output: PathBuf::from("o"),
872 method: MergeMethod::Slerp,
873 weight: Some(0.5),
874 density: None,
875 weights: None,
876 base: None,
877 adapter: None,
878 };
879 assert!(perform_merge(&models, &a).is_ok());
880 }
881
882 #[test]
883 fn test_perform_merge_average_route() {
884 let models = vec![mk(&[("w", &[2.0])]), mk(&[("w", &[4.0])])];
885 let a = MergeArgs {
886 models: vec![],
887 output: PathBuf::from("o"),
888 method: MergeMethod::Average,
889 weight: None,
890 density: None,
891 weights: None,
892 base: None,
893 adapter: None,
894 };
895 let result = perform_merge(&models, &a).unwrap();
896 let vals = result["w"].data().as_slice().unwrap().to_vec();
897 assert!((vals[0] - 3.0).abs() < 1e-6);
898 }
899
900 #[test]
903 fn test_slerp_merge_two_models_ok() {
904 let ms = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
905 let a = MergeArgs {
906 models: vec![],
907 output: PathBuf::from("o"),
908 method: MergeMethod::Slerp,
909 weight: Some(0.3),
910 density: None,
911 weights: None,
912 base: None,
913 adapter: None,
914 };
915 assert!(perform_slerp_merge(&ms, &a).is_ok());
916 }
917
918 #[test]
919 fn test_slerp_merge_default_weight() {
920 let ms = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
921 let a = MergeArgs {
922 models: vec![],
923 output: PathBuf::from("o"),
924 method: MergeMethod::Slerp,
925 weight: None, density: None,
927 weights: None,
928 base: None,
929 adapter: None,
930 };
931 assert!(perform_slerp_merge(&ms, &a).is_ok());
932 }
933
934 #[test]
937 fn test_ties_merge_with_density() {
938 let a = MergeArgs {
939 models: vec![],
940 output: PathBuf::from("o"),
941 method: MergeMethod::Ties,
942 weight: None,
943 density: Some(0.8),
944 weights: None,
945 base: None,
946 adapter: None,
947 };
948 let models =
949 vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])]), mk(&[("w", &[1.2, 2.2])])];
950 let result = perform_ties_merge(&models, &a);
951 assert!(result.is_ok());
952 }
953
954 #[test]
957 fn test_dare_merge_with_density() {
958 let a = MergeArgs {
959 models: vec![],
960 output: PathBuf::from("o"),
961 method: MergeMethod::Dare,
962 weight: None,
963 density: Some(0.9),
964 weights: None,
965 base: None,
966 adapter: None,
967 };
968 let models = vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])])];
969 assert!(perform_dare_merge(&models, &a).is_ok());
970 }
971
972 #[test]
975 fn test_average_merge_with_weights() {
976 let a = MergeArgs {
977 models: vec![],
978 output: PathBuf::from("o"),
979 method: MergeMethod::Average,
980 weight: None,
981 density: None,
982 weights: Some("0.8,0.2".to_string()),
983 base: None,
984 adapter: None,
985 };
986 let models = vec![mk(&[("w", &[10.0])]), mk(&[("w", &[0.0])])];
987 let result = perform_average_merge(&models, &a).unwrap();
988 let vals = result["w"].data().as_slice().unwrap().to_vec();
989 assert!((vals[0] - 8.0).abs() < 1e-4);
991 }
992
993 #[test]
996 fn test_build_ensemble_config_single_weight() {
997 let a = MergeArgs {
998 models: vec![],
999 output: PathBuf::from("o.json"),
1000 method: MergeMethod::Average,
1001 weight: None,
1002 density: None,
1003 weights: Some("1.0".to_string()),
1004 base: None,
1005 adapter: None,
1006 };
1007 let config = build_ensemble_config(&a);
1008 assert!(config.is_ok());
1009 }
1010
1011 #[test]
1012 fn test_build_ensemble_config_three_weights() {
1013 let a = MergeArgs {
1014 models: vec![],
1015 output: PathBuf::from("o.json"),
1016 method: MergeMethod::Average,
1017 weight: None,
1018 density: None,
1019 weights: Some("0.2, 0.3, 0.5".to_string()),
1020 base: None,
1021 adapter: None,
1022 };
1023 let config = build_ensemble_config(&a);
1024 assert!(config.is_ok());
1025 }
1026
1027 #[test]
1028 fn test_build_ensemble_config_empty_weights_string() {
1029 let a = MergeArgs {
1030 models: vec![],
1031 output: PathBuf::from("o.json"),
1032 method: MergeMethod::Average,
1033 weight: None,
1034 density: None,
1035 weights: Some(String::new()),
1036 base: None,
1037 adapter: None,
1038 };
1039 assert!(build_ensemble_config(&a).is_err());
1041 }
1042
1043 #[test]
1046 fn test_validate_model_count_one() {
1047 let a = MergeArgs {
1048 models: vec![PathBuf::from("a")],
1049 output: PathBuf::from("o"),
1050 method: MergeMethod::Ties,
1051 weight: None,
1052 density: None,
1053 weights: None,
1054 base: None,
1055 adapter: None,
1056 };
1057 assert!(validate_model_count(&a).is_err());
1058 }
1059
1060 #[test]
1061 fn test_validate_model_count_three() {
1062 let a = MergeArgs {
1063 models: vec![PathBuf::from("a"), PathBuf::from("b"), PathBuf::from("c")],
1064 output: PathBuf::from("o"),
1065 method: MergeMethod::Ties,
1066 weight: None,
1067 density: None,
1068 weights: None,
1069 base: None,
1070 adapter: None,
1071 };
1072 assert!(validate_model_count(&a).is_ok());
1073 }
1074
1075 #[test]
1078 fn test_export_merged_model_no_extension() {
1079 let m = mk(&[("w", &[1.0])]);
1080 let t = std::env::temp_dir().join("ent_merge_noext");
1081 let a = MergeArgs {
1082 models: vec![],
1083 output: t.clone(),
1084 method: MergeMethod::Average,
1085 weight: None,
1086 density: None,
1087 weights: None,
1088 base: None,
1089 adapter: None,
1090 };
1091 assert!(export_merged_model(&m, &a).is_ok());
1093 let _ = std::fs::remove_file(&t);
1094 }
1095
1096 #[test]
1099 fn test_bytes_to_f32_f32_multiple() {
1100 let vals = vec![1.0f32, 2.0, 3.5, -1.0];
1101 let bytes: Vec<u8> = vals.iter().flat_map(|x| x.to_le_bytes()).collect();
1102 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F32);
1103 assert_eq!(result.len(), 4);
1104 assert!((result[0] - 1.0).abs() < 1e-6);
1105 assert!((result[1] - 2.0).abs() < 1e-6);
1106 assert!((result[2] - 3.5).abs() < 1e-6);
1107 assert!((result[3] - (-1.0)).abs() < 1e-6);
1108 }
1109
1110 #[test]
1111 fn test_bytes_to_f32_f16_multiple() {
1112 let vals = vec![half::f16::from_f32(0.5), half::f16::from_f32(1.5)];
1113 let bytes: Vec<u8> = vals.iter().flat_map(|x| x.to_le_bytes()).collect();
1114 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F16);
1115 assert_eq!(result.len(), 2);
1116 assert!((result[0] - 0.5).abs() < 0.01);
1117 assert!((result[1] - 1.5).abs() < 0.01);
1118 }
1119
1120 #[test]
1121 fn test_bytes_to_f32_bf16_multiple() {
1122 let vals = vec![half::bf16::from_f32(3.0), half::bf16::from_f32(-1.0)];
1123 let bytes: Vec<u8> = vals.iter().flat_map(|x| x.to_le_bytes()).collect();
1124 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::BF16);
1125 assert_eq!(result.len(), 2);
1126 assert!((result[0] - 3.0).abs() < 0.1);
1127 assert!((result[1] - (-1.0)).abs() < 0.1);
1128 }
1129
1130 #[test]
1133 fn test_safetensor_metadata_ties() {
1134 let m = mk(&[("a", &[1.0]), ("b", &[2.0]), ("c", &[3.0])]);
1135 let a = MergeArgs {
1136 models: vec![],
1137 output: PathBuf::from("o.st"),
1138 method: MergeMethod::Ties,
1139 weight: None,
1140 density: None,
1141 weights: None,
1142 base: None,
1143 adapter: None,
1144 };
1145 let md = build_safetensor_metadata(&m, &a);
1146 assert_eq!(md["name"], "merged-model");
1147 assert_eq!(md["tensor_count"], "3");
1148 assert!(md["merge_method"].contains("Ties"));
1149 }
1150
1151 #[test]
1152 fn test_safetensor_metadata_slerp() {
1153 let m = mk(&[("x", &[1.0])]);
1154 let a = MergeArgs {
1155 models: vec![],
1156 output: PathBuf::from("o.st"),
1157 method: MergeMethod::Slerp,
1158 weight: None,
1159 density: None,
1160 weights: None,
1161 base: None,
1162 adapter: None,
1163 };
1164 let md = build_safetensor_metadata(&m, &a);
1165 assert!(md["merge_method"].contains("Slerp"));
1166 }
1167
1168 #[test]
1171 fn test_log_merge_start_normal() {
1172 let a = MergeArgs {
1173 models: vec![PathBuf::from("m1"), PathBuf::from("m2")],
1174 output: PathBuf::from("out"),
1175 method: MergeMethod::Average,
1176 weight: None,
1177 density: None,
1178 weights: None,
1179 base: None,
1180 adapter: None,
1181 };
1182 log_merge_start(&a, LogLevel::Normal);
1183 }
1184
1185 #[test]
1186 fn test_log_merge_complete_verbose() {
1187 let m = mk(&[("a", &[1.0, 2.0])]);
1188 let a = MergeArgs {
1189 models: vec![],
1190 output: PathBuf::from("merged.json"),
1191 method: MergeMethod::Dare,
1192 weight: None,
1193 density: None,
1194 weights: None,
1195 base: None,
1196 adapter: None,
1197 };
1198 log_merge_complete(&m, &a, LogLevel::Verbose);
1199 }
1200
1201 #[test]
1204 fn test_lora_adapter_config_not_found() {
1205 let dir = tempfile::tempdir().unwrap();
1207 let base_file = dir.path().join("base.safetensors");
1209 std::fs::write(&base_file, b"fake").unwrap();
1210 let a = MergeArgs {
1211 models: vec![],
1212 output: PathBuf::from("o"),
1213 method: MergeMethod::LoraAdapter,
1214 weight: None,
1215 density: None,
1216 weights: None,
1217 base: Some(base_file),
1218 adapter: Some(dir.path().to_path_buf()),
1219 };
1220 let err = run_lora_adapter_merge(&a, LogLevel::Quiet).unwrap_err();
1221 assert!(err.contains("adapter_config.json"), "Error: {err}");
1222 }
1223
1224 #[test]
1225 fn test_lora_adapter_model_not_found() {
1226 let dir = tempfile::tempdir().unwrap();
1227 let base_file = dir.path().join("base.safetensors");
1228 std::fs::write(&base_file, b"fake").unwrap();
1229 std::fs::write(dir.path().join("adapter_config.json"), r#"{"r": 8, "lora_alpha": 16}"#)
1231 .unwrap();
1232 let a = MergeArgs {
1233 models: vec![],
1234 output: PathBuf::from("o"),
1235 method: MergeMethod::LoraAdapter,
1236 weight: None,
1237 density: None,
1238 weights: None,
1239 base: Some(base_file),
1240 adapter: Some(dir.path().to_path_buf()),
1241 };
1242 let err = run_lora_adapter_merge(&a, LogLevel::Quiet).unwrap_err();
1243 assert!(err.contains("adapter_model.safetensors"), "Error: {err}");
1244 }
1245
1246 #[test]
1249 fn test_run_merge_nonexistent_models() {
1250 let a = MergeArgs {
1251 models: vec![PathBuf::from("/no/m1"), PathBuf::from("/no/m2")],
1252 output: PathBuf::from("o"),
1253 method: MergeMethod::Ties,
1254 weight: None,
1255 density: None,
1256 weights: None,
1257 base: None,
1258 adapter: None,
1259 };
1260 let err = run_merge(a, LogLevel::Quiet).unwrap_err();
1261 assert!(err.contains("Failed to read"), "Error: {err}");
1262 }
1263
1264 #[test]
1267 fn test_mk_helper_creates_model() {
1268 let model = mk(&[("a", &[1.0, 2.0, 3.0]), ("b", &[4.0])]);
1269 assert_eq!(model.len(), 2);
1270 assert!(model.contains_key("a"));
1271 assert!(model.contains_key("b"));
1272 assert_eq!(model["a"].len(), 3);
1273 assert_eq!(model["b"].len(), 1);
1274 }
1275
1276 #[test]
1279 fn test_export_safetensors_multiple_tensors() {
1280 let m = mk(&[("w1", &[1.0, 2.0]), ("w2", &[3.0, 4.0, 5.0])]);
1281 let t = std::env::temp_dir().join("ent_merge_multi.safetensors");
1282 let a = MergeArgs {
1283 models: vec![],
1284 output: t.clone(),
1285 method: MergeMethod::Average,
1286 weight: None,
1287 density: None,
1288 weights: None,
1289 base: None,
1290 adapter: None,
1291 };
1292 assert!(export_merged_model(&m, &a).is_ok());
1293 assert!(t.exists());
1295 let _ = std::fs::remove_file(&t);
1296 }
1297
1298 #[test]
1301 fn test_export_json_roundtrip() {
1302 let m = mk(&[("w1", &[1.0, 2.0]), ("w2", &[3.0])]);
1303 let t = std::env::temp_dir().join("ent_merge_roundtrip.json");
1304 let a = MergeArgs {
1305 models: vec![],
1306 output: t.clone(),
1307 method: MergeMethod::Average,
1308 weight: None,
1309 density: None,
1310 weights: None,
1311 base: None,
1312 adapter: None,
1313 };
1314 assert!(export_merged_model(&m, &a).is_ok());
1315 let content = std::fs::read_to_string(&t).unwrap();
1316 let parsed: HashMap<String, Vec<f32>> = serde_json::from_str(&content).unwrap();
1317 assert!(parsed.contains_key("w1"));
1318 assert!(parsed.contains_key("w2"));
1319 assert_eq!(parsed["w1"].len(), 2);
1320 let _ = std::fs::remove_file(&t);
1321 }
1322
1323 fn mk_args(method: MergeMethod) -> MergeArgs {
1329 MergeArgs {
1330 models: vec![],
1331 output: PathBuf::from("out.json"),
1332 method,
1333 weight: None,
1334 density: None,
1335 weights: None,
1336 base: None,
1337 adapter: None,
1338 }
1339 }
1340
1341 #[test]
1344 fn test_cov2_bytes_to_f32_f32_zeros() {
1345 let zeros = vec![0.0f32; 10];
1346 let bytes: Vec<u8> = zeros.iter().flat_map(|x| x.to_le_bytes()).collect();
1347 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F32);
1348 assert_eq!(result.len(), 10);
1349 assert!(result.iter().all(|&v| v == 0.0));
1350 }
1351
1352 #[test]
1353 fn test_cov2_bytes_to_f32_f32_negative() {
1354 let vals = vec![-1.0f32, -100.0, -0.001];
1355 let bytes: Vec<u8> = vals.iter().flat_map(|x| x.to_le_bytes()).collect();
1356 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F32);
1357 assert_eq!(result.len(), 3);
1358 assert!((result[0] - (-1.0)).abs() < 1e-6);
1359 assert!((result[1] - (-100.0)).abs() < 1e-6);
1360 assert!((result[2] - (-0.001)).abs() < 1e-6);
1361 }
1362
1363 #[test]
1364 fn test_cov2_bytes_to_f32_f32_large() {
1365 let vals = vec![1e30f32, -1e30];
1366 let bytes: Vec<u8> = vals.iter().flat_map(|x| x.to_le_bytes()).collect();
1367 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F32);
1368 assert_eq!(result.len(), 2);
1369 assert!((result[0] - 1e30).abs() / 1e30 < 1e-6);
1370 }
1371
1372 #[test]
1373 fn test_cov2_bytes_to_f32_f16_zero() {
1374 let zero = half::f16::from_f32(0.0);
1375 let bytes = zero.to_le_bytes().to_vec();
1376 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F16);
1377 assert_eq!(result.len(), 1);
1378 assert!((result[0]).abs() < 1e-6);
1379 }
1380
1381 #[test]
1382 fn test_cov2_bytes_to_f32_bf16_zero() {
1383 let zero = half::bf16::from_f32(0.0);
1384 let bytes = zero.to_le_bytes().to_vec();
1385 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::BF16);
1386 assert_eq!(result.len(), 1);
1387 assert!((result[0]).abs() < 1e-6);
1388 }
1389
1390 #[test]
1391 fn test_cov2_bytes_to_f32_f16_negative() {
1392 let neg = half::f16::from_f32(-3.14);
1393 let bytes = neg.to_le_bytes().to_vec();
1394 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F16);
1395 assert_eq!(result.len(), 1);
1396 assert!((result[0] - (-3.14)).abs() < 0.01);
1397 }
1398
1399 #[test]
1400 fn test_cov2_bytes_to_f32_bf16_negative() {
1401 let neg = half::bf16::from_f32(-5.0);
1402 let bytes = neg.to_le_bytes().to_vec();
1403 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::BF16);
1404 assert_eq!(result.len(), 1);
1405 assert!((result[0] - (-5.0)).abs() < 0.5);
1406 }
1407
1408 #[test]
1411 fn test_cov2_bytes_to_f32_f32_truncated() {
1412 let bytes: Vec<u8> = vec![0, 0, 128, 63, 99];
1414 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F32);
1415 assert_eq!(result.len(), 1);
1416 assert!((result[0] - 1.0).abs() < 1e-6);
1417 }
1418
1419 #[test]
1420 fn test_cov2_bytes_to_f32_f16_truncated() {
1421 let val = half::f16::from_f32(2.0);
1423 let mut bytes = val.to_le_bytes().to_vec();
1424 bytes.push(0xFF);
1425 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::F16);
1426 assert_eq!(result.len(), 1);
1427 assert!((result[0] - 2.0).abs() < 0.01);
1428 }
1429
1430 #[test]
1433 fn test_cov2_bytes_to_f32_i64_fallback() {
1434 let v = 3.14f32;
1435 let bytes = v.to_le_bytes().to_vec();
1436 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::I64);
1437 assert_eq!(result.len(), 1);
1438 assert!((result[0] - 3.14).abs() < 1e-6);
1439 }
1440
1441 #[test]
1442 fn test_cov2_bytes_to_f32_u8_fallback() {
1443 let v = 7.0f32;
1444 let bytes = v.to_le_bytes().to_vec();
1445 let result = bytes_to_f32(&bytes, safetensors::tensor::Dtype::U8);
1446 assert_eq!(result.len(), 1);
1447 assert!((result[0] - 7.0).abs() < 1e-6);
1448 }
1449
1450 #[test]
1453 fn test_cov2_build_ensemble_config_whitespace_weights() {
1454 let a = MergeArgs {
1455 weights: Some(" 0.5 , 0.3 , 0.2 ".to_string()),
1456 ..mk_args(MergeMethod::Average)
1457 };
1458 let config = build_ensemble_config(&a);
1459 assert!(config.is_ok());
1460 }
1461
1462 #[test]
1465 fn test_cov2_build_ensemble_config_negative_weights() {
1466 let a =
1467 MergeArgs { weights: Some("-0.5, 1.5".to_string()), ..mk_args(MergeMethod::Average) };
1468 let config = build_ensemble_config(&a);
1469 assert!(config.is_ok());
1471 }
1472
1473 #[test]
1476 fn test_cov2_build_ensemble_config_many_weights() {
1477 let w_str = (0..10).map(|_| "0.1").collect::<Vec<_>>().join(",");
1478 let a = MergeArgs { weights: Some(w_str), ..mk_args(MergeMethod::Average) };
1479 let config = build_ensemble_config(&a);
1480 assert!(config.is_ok());
1481 }
1482
1483 #[test]
1486 fn test_cov2_safetensor_metadata_average() {
1487 let m = mk(&[("w", &[1.0])]);
1488 let a = mk_args(MergeMethod::Average);
1489 let md = build_safetensor_metadata(&m, &a);
1490 assert!(md["merge_method"].contains("Average"));
1491 assert_eq!(md["tensor_count"], "1");
1492 }
1493
1494 #[test]
1495 fn test_cov2_safetensor_metadata_dare() {
1496 let m = mk(&[("a", &[1.0]), ("b", &[2.0])]);
1497 let a = mk_args(MergeMethod::Dare);
1498 let md = build_safetensor_metadata(&m, &a);
1499 assert!(md["merge_method"].contains("Dare"));
1500 assert_eq!(md["tensor_count"], "2");
1501 }
1502
1503 #[test]
1504 fn test_cov2_safetensor_metadata_lora() {
1505 let m = mk(&[("w", &[1.0])]);
1506 let a = mk_args(MergeMethod::LoraAdapter);
1507 let md = build_safetensor_metadata(&m, &a);
1508 assert!(md["merge_method"].contains("LoraAdapter"));
1509 }
1510
1511 #[test]
1512 fn test_cov2_safetensor_metadata_empty_model() {
1513 let m: Model = HashMap::new();
1514 let a = mk_args(MergeMethod::Ties);
1515 let md = build_safetensor_metadata(&m, &a);
1516 assert_eq!(md["tensor_count"], "0");
1517 }
1518
1519 #[test]
1522 fn test_cov2_validate_model_count_exactly_2() {
1523 let a = MergeArgs {
1524 models: vec![PathBuf::from("a"), PathBuf::from("b")],
1525 ..mk_args(MergeMethod::Average)
1526 };
1527 assert!(validate_model_count(&a).is_ok());
1528 }
1529
1530 #[test]
1531 fn test_cov2_validate_model_count_large() {
1532 let models: Vec<PathBuf> = (0..100).map(|i| PathBuf::from(format!("m{i}"))).collect();
1533 let a = MergeArgs { models, ..mk_args(MergeMethod::Average) };
1534 assert!(validate_model_count(&a).is_ok());
1535 }
1536
1537 #[test]
1540 fn test_cov2_perform_merge_lora_error_msg() {
1541 let a = mk_args(MergeMethod::LoraAdapter);
1542 let err = perform_merge(&[], &a).unwrap_err();
1543 assert_eq!(err, "LoRA adapter merge uses dedicated path");
1544 }
1545
1546 #[test]
1549 fn test_cov2_export_json_bad_path() {
1550 let m = mk(&[("w", &[1.0])]);
1551 let a = MergeArgs {
1552 output: PathBuf::from("/nonexistent_dir_xxxx/output.json"),
1553 ..mk_args(MergeMethod::Average)
1554 };
1555 let result = export_merged_model(&m, &a);
1556 assert!(result.is_err());
1557 assert!(result.unwrap_err().contains("Failed to write"));
1558 }
1559
1560 #[test]
1561 fn test_cov2_export_safetensors_bad_path() {
1562 let m = mk(&[("w", &[1.0])]);
1563 let a = MergeArgs {
1564 output: PathBuf::from("/nonexistent_dir_xxxx/output.safetensors"),
1565 ..mk_args(MergeMethod::Average)
1566 };
1567 let result = export_merged_model(&m, &a);
1568 assert!(result.is_err());
1569 assert!(result.unwrap_err().contains("Failed to write"));
1570 }
1571
1572 #[test]
1575 fn test_cov2_export_safetensors_roundtrip() {
1576 let m = mk(&[("layer1", &[1.0, 2.0, 3.0]), ("layer2", &[4.0, 5.0])]);
1577 let t = std::env::temp_dir().join("ent_merge_cov2_rt.safetensors");
1578 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Ties) };
1579 assert!(export_merged_model(&m, &a).is_ok());
1580 let data = std::fs::read(&t).unwrap();
1582 let tensors = SafeTensors::deserialize(&data).unwrap();
1583 let names: Vec<&str> = tensors.names().clone();
1584 assert!(names.contains(&"layer1"));
1585 assert!(names.contains(&"layer2"));
1586 let _ = std::fs::remove_file(&t);
1587 }
1588
1589 #[test]
1592 fn test_cov2_export_json_empty_model() {
1593 let m: Model = HashMap::new();
1594 let t = std::env::temp_dir().join("ent_merge_cov2_empty.json");
1595 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Average) };
1596 assert!(export_merged_model(&m, &a).is_ok());
1597 let content = std::fs::read_to_string(&t).unwrap();
1598 let parsed: HashMap<String, Vec<f32>> = serde_json::from_str(&content).unwrap();
1599 assert!(parsed.is_empty());
1600 let _ = std::fs::remove_file(&t);
1601 }
1602
1603 #[test]
1606 fn test_cov2_export_safetensors_empty_model() {
1607 let m: Model = HashMap::new();
1608 let t = std::env::temp_dir().join("ent_merge_cov2_empty.safetensors");
1609 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Average) };
1610 assert!(export_merged_model(&m, &a).is_ok());
1611 let _ = std::fs::remove_file(&t);
1612 }
1613
1614 #[test]
1617 fn test_cov2_ties_merge_default_density() {
1618 let a = MergeArgs { density: None, ..mk_args(MergeMethod::Ties) };
1619 let models =
1621 vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])]), mk(&[("w", &[1.2, 2.2])])];
1622 assert!(perform_ties_merge(&models, &a).is_ok());
1623 }
1624
1625 #[test]
1628 fn test_cov2_dare_merge_default_density() {
1629 let a = MergeArgs { density: None, ..mk_args(MergeMethod::Dare) };
1630 let models = vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])])];
1632 assert!(perform_dare_merge(&models, &a).is_ok());
1633 }
1634
1635 #[test]
1638 fn test_cov2_slerp_merge_default_weight() {
1639 let a = MergeArgs { weight: None, ..mk_args(MergeMethod::Slerp) };
1640 let models = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
1641 let result = perform_slerp_merge(&models, &a);
1642 assert!(result.is_ok());
1643 }
1644
1645 #[test]
1648 fn test_cov2_slerp_single_model() {
1649 let a = mk_args(MergeMethod::Slerp);
1650 let models = vec![mk(&[("w", &[1.0])])];
1651 let err = perform_slerp_merge(&models, &a).unwrap_err();
1652 assert!(err.contains("SLERP requires exactly 2"));
1653 }
1654
1655 #[test]
1658 fn test_cov2_run_merge_lora_missing_both() {
1659 let a = MergeArgs {
1660 method: MergeMethod::LoraAdapter,
1661 base: None,
1662 adapter: None,
1663 ..mk_args(MergeMethod::LoraAdapter)
1664 };
1665 let err = run_merge(a, LogLevel::Quiet).unwrap_err();
1666 assert!(err.contains("--base required"));
1667 }
1668
1669 #[test]
1670 fn test_cov2_run_merge_lora_has_base_no_adapter() {
1671 let a = MergeArgs {
1672 method: MergeMethod::LoraAdapter,
1673 base: Some(PathBuf::from("/tmp/some_base")),
1674 adapter: None,
1675 ..mk_args(MergeMethod::LoraAdapter)
1676 };
1677 let err = run_merge(a, LogLevel::Quiet).unwrap_err();
1678 assert!(err.contains("--adapter required"));
1679 }
1680
1681 #[test]
1684 fn test_cov2_load_single_model_empty_file() {
1685 let dir = tempfile::tempdir().unwrap();
1686 let path = dir.path().join("empty.safetensors");
1687 std::fs::write(&path, b"").unwrap();
1688 let err = load_single_model(&path).unwrap_err();
1689 assert!(err.contains("Failed to parse"));
1690 }
1691
1692 #[test]
1695 fn test_cov2_load_single_model_garbage() {
1696 let dir = tempfile::tempdir().unwrap();
1697 let path = dir.path().join("garbage.safetensors");
1698 std::fs::write(&path, b"this is not a safetensors file at all").unwrap();
1699 let err = load_single_model(&path).unwrap_err();
1700 assert!(err.contains("Failed to parse"));
1701 }
1702
1703 #[test]
1706 fn test_cov2_run_merge_first_model_missing() {
1707 let dir = tempfile::tempdir().unwrap();
1708 let a = MergeArgs {
1709 models: vec![
1710 dir.path().join("no_exist_1.safetensors"),
1711 dir.path().join("no_exist_2.safetensors"),
1712 ],
1713 output: dir.path().join("out.json"),
1714 method: MergeMethod::Average,
1715 ..mk_args(MergeMethod::Average)
1716 };
1717 let err = run_merge(a, LogLevel::Quiet).unwrap_err();
1718 assert!(err.contains("Failed to read"));
1719 }
1720
1721 #[test]
1724 fn test_cov2_log_merge_start_quiet() {
1725 let a = MergeArgs {
1726 models: vec![PathBuf::from("m1"), PathBuf::from("m2"), PathBuf::from("m3")],
1727 output: PathBuf::from("out.safetensors"),
1728 ..mk_args(MergeMethod::Dare)
1729 };
1730 log_merge_start(&a, LogLevel::Quiet);
1731 }
1732
1733 #[test]
1734 fn test_cov2_log_merge_complete_quiet() {
1735 let m = mk(&[("w", &[1.0, 2.0, 3.0])]);
1736 let a = MergeArgs {
1737 output: PathBuf::from("merged.safetensors"),
1738 ..mk_args(MergeMethod::Average)
1739 };
1740 log_merge_complete(&m, &a, LogLevel::Quiet);
1741 }
1742
1743 #[test]
1746 fn test_cov2_average_merge_multi_tensor() {
1747 let a = mk_args(MergeMethod::Average);
1748 let m1 = mk(&[("a", &[1.0, 2.0]), ("b", &[3.0])]);
1749 let m2 = mk(&[("a", &[3.0, 4.0]), ("b", &[5.0])]);
1750 let result = perform_average_merge(&[m1, m2], &a).unwrap();
1751 let a_vals = result["a"].data().as_slice().unwrap().to_vec();
1752 let b_vals = result["b"].data().as_slice().unwrap().to_vec();
1753 assert!((a_vals[0] - 2.0).abs() < 1e-6);
1754 assert!((a_vals[1] - 3.0).abs() < 1e-6);
1755 assert!((b_vals[0] - 4.0).abs() < 1e-6);
1756 }
1757
1758 #[test]
1761 fn test_cov2_slerp_merge_weight_0() {
1762 let a = MergeArgs { weight: Some(0.0), ..mk_args(MergeMethod::Slerp) };
1763 let models = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
1764 let result = perform_slerp_merge(&models, &a).unwrap();
1765 let vals = result["w"].data().as_slice().unwrap().to_vec();
1766 assert!((vals[0] - 1.0).abs() < 0.1);
1768 }
1769
1770 #[test]
1771 fn test_cov2_slerp_merge_weight_1() {
1772 let a = MergeArgs { weight: Some(1.0), ..mk_args(MergeMethod::Slerp) };
1773 let models = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
1774 let result = perform_slerp_merge(&models, &a).unwrap();
1775 let vals = result["w"].data().as_slice().unwrap().to_vec();
1776 assert!((vals[1] - 1.0).abs() < 0.1);
1778 }
1779
1780 #[test]
1783 fn test_cov2_ties_merge_high_density() {
1784 let a = MergeArgs { density: Some(0.99), ..mk_args(MergeMethod::Ties) };
1785 let models = vec![
1786 mk(&[("w", &[1.0, 2.0, 3.0])]),
1787 mk(&[("w", &[1.1, 2.1, 3.1])]),
1788 mk(&[("w", &[1.2, 2.2, 3.2])]),
1789 ];
1790 assert!(perform_ties_merge(&models, &a).is_ok());
1791 }
1792
1793 #[test]
1796 fn test_cov2_dare_merge_low_density() {
1797 let a = MergeArgs { density: Some(0.01), ..mk_args(MergeMethod::Dare) };
1798 let models = vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])])];
1799 assert!(perform_dare_merge(&models, &a).is_ok());
1800 }
1801
1802 #[test]
1805 fn test_cov2_lora_adapter_config_exists_no_model() {
1806 let dir = tempfile::tempdir().unwrap();
1807 let base = dir.path().join("base.safetensors");
1808 std::fs::write(&base, b"fake").unwrap();
1809 let adapter_dir = dir.path().join("adapter");
1810 std::fs::create_dir_all(&adapter_dir).unwrap();
1811 std::fs::write(adapter_dir.join("adapter_config.json"), r#"{"r":8}"#).unwrap();
1812 let a = MergeArgs {
1814 base: Some(base),
1815 adapter: Some(adapter_dir),
1816 ..mk_args(MergeMethod::LoraAdapter)
1817 };
1818 let err = run_lora_adapter_merge(&a, LogLevel::Quiet).unwrap_err();
1819 assert!(err.contains("adapter_model.safetensors"));
1820 }
1821
1822 #[test]
1825 fn test_cov2_lora_base_path_not_found() {
1826 let dir = tempfile::tempdir().unwrap();
1827 let adapter_dir = dir.path().join("adapter");
1828 std::fs::create_dir_all(&adapter_dir).unwrap();
1829 let a = MergeArgs {
1830 base: Some(PathBuf::from("/definitely/not/exist/base.st")),
1831 adapter: Some(adapter_dir),
1832 ..mk_args(MergeMethod::LoraAdapter)
1833 };
1834 let err = run_lora_adapter_merge(&a, LogLevel::Quiet).unwrap_err();
1835 assert!(err.contains("Base model not found"));
1836 }
1837
1838 #[test]
1841 fn test_cov2_export_extension_safetensors() {
1842 let m = mk(&[("w", &[1.0])]);
1843 let t = std::env::temp_dir().join("ent_merge_cov2_ext.safetensors");
1844 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Average) };
1845 assert!(export_merged_model(&m, &a).is_ok());
1846 let data = std::fs::read(&t).unwrap();
1848 assert!(SafeTensors::deserialize(&data).is_ok());
1849 let _ = std::fs::remove_file(&t);
1850 }
1851
1852 #[test]
1853 fn test_cov2_export_extension_json() {
1854 let m = mk(&[("w", &[1.0])]);
1855 let t = std::env::temp_dir().join("ent_merge_cov2_ext.json");
1856 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Average) };
1857 assert!(export_merged_model(&m, &a).is_ok());
1858 let content = std::fs::read_to_string(&t).unwrap();
1859 assert!(serde_json::from_str::<HashMap<String, Vec<f32>>>(&content).is_ok());
1860 let _ = std::fs::remove_file(&t);
1861 }
1862
1863 #[test]
1864 fn test_cov2_export_extension_unknown() {
1865 let m = mk(&[("w", &[1.0])]);
1866 let t = std::env::temp_dir().join("ent_merge_cov2_ext.bin");
1867 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Average) };
1868 assert!(export_merged_model(&m, &a).is_ok());
1870 let content = std::fs::read_to_string(&t).unwrap();
1871 assert!(serde_json::from_str::<HashMap<String, Vec<f32>>>(&content).is_ok());
1872 let _ = std::fs::remove_file(&t);
1873 }
1874
1875 #[test]
1878 fn test_cov2_mk_empty_model() {
1879 let model = mk(&[]);
1880 assert!(model.is_empty());
1881 }
1882
1883 #[test]
1884 fn test_cov2_mk_single_empty_tensor() {
1885 let model = mk(&[("empty", &[])]);
1886 assert_eq!(model.len(), 1);
1887 assert_eq!(model["empty"].len(), 0);
1888 }
1889
1890 #[test]
1893 fn test_cov2_perform_merge_all_methods() {
1894 let models3 =
1896 vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.1, 2.1])]), mk(&[("w", &[1.2, 2.2])])];
1897 assert!(perform_merge(&models3, &mk_args(MergeMethod::Ties)).is_ok());
1898
1899 let models2 = vec![mk(&[("w", &[1.0, 2.0])]), mk(&[("w", &[1.5, 2.5])])];
1901 assert!(perform_merge(&models2, &mk_args(MergeMethod::Dare)).is_ok());
1902
1903 let models_s = vec![mk(&[("w", &[1.0, 0.0])]), mk(&[("w", &[0.0, 1.0])])];
1905 assert!(perform_merge(
1906 &models_s,
1907 &MergeArgs { weight: Some(0.5), ..mk_args(MergeMethod::Slerp) }
1908 )
1909 .is_ok());
1910
1911 let models_a = vec![mk(&[("w", &[2.0])]), mk(&[("w", &[4.0])])];
1913 assert!(perform_merge(&models_a, &mk_args(MergeMethod::Average)).is_ok());
1914
1915 assert!(perform_merge(&[], &mk_args(MergeMethod::LoraAdapter)).is_err());
1917 }
1918
1919 #[test]
1922 fn test_cov2_large_tensor_roundtrip() {
1923 let large_data: Vec<f32> = (0..1000).map(|i| i as f32 * 0.001).collect();
1924 let m = mk(&[("big", large_data.as_slice())]);
1925 let t = std::env::temp_dir().join("ent_merge_cov2_large.json");
1926 let a = MergeArgs { output: t.clone(), ..mk_args(MergeMethod::Average) };
1927 assert!(export_merged_model(&m, &a).is_ok());
1928 let content = std::fs::read_to_string(&t).unwrap();
1929 let parsed: HashMap<String, Vec<f32>> = serde_json::from_str(&content).unwrap();
1930 assert_eq!(parsed["big"].len(), 1000);
1931 assert!((parsed["big"][500] - 0.5).abs() < 1e-3);
1932 let _ = std::fs::remove_file(&t);
1933 }
1934
1935 fn write_st(path: &Path, tensors: &[(&str, safetensors::tensor::Dtype, Vec<usize>, Vec<u8>)]) {
1943 let views: Vec<(&str, safetensors::tensor::TensorView<'_>)> = tensors
1944 .iter()
1945 .map(|(name, dt, shape, bytes)| {
1946 (*name, safetensors::tensor::TensorView::new(*dt, shape.clone(), bytes).unwrap())
1947 })
1948 .collect();
1949 let bytes = safetensors::serialize(views, None).unwrap();
1950 std::fs::write(path, bytes).unwrap();
1951 }
1952
1953 fn f32_bytes(vals: &[f32]) -> Vec<u8> {
1954 vals.iter().flat_map(|x| x.to_le_bytes()).collect()
1955 }
1956
1957 fn bf16_bytes(vals: &[f32]) -> Vec<u8> {
1958 vals.iter().flat_map(|x| half::bf16::from_f32(*x).to_le_bytes()).collect()
1959 }
1960
1961 #[test]
1967 fn test_lora_merge_rslora_scaling_pmat897() {
1968 let dir = tempfile::tempdir().unwrap();
1969 let adapter_dir = dir.path();
1970
1971 std::fs::write(
1973 adapter_dir.join("adapter_config.json"),
1974 r#"{"r": 16, "lora_alpha": 16, "use_rslora": true}"#,
1975 )
1976 .unwrap();
1977
1978 let base_path = adapter_dir.join("base.safetensors");
1980 write_st(
1981 &base_path,
1982 &[("w.weight", safetensors::tensor::Dtype::F32, vec![1, 1], f32_bytes(&[0.0]))],
1983 );
1984
1985 let mut a_vals = vec![0.0f32; 16];
1988 a_vals[0] = 1.0;
1989 let mut b_vals = vec![0.0f32; 16];
1990 b_vals[0] = 1.0;
1991 write_st(
1992 &adapter_dir.join("adapter_model.safetensors"),
1993 &[
1994 (
1995 "base_model.model.w.lora_A.weight",
1996 safetensors::tensor::Dtype::F32,
1997 vec![16, 1],
1998 f32_bytes(&a_vals),
1999 ),
2000 (
2001 "base_model.model.w.lora_B.weight",
2002 safetensors::tensor::Dtype::F32,
2003 vec![1, 16],
2004 f32_bytes(&b_vals),
2005 ),
2006 ],
2007 );
2008
2009 let out = dir.path().join("merged.safetensors");
2010 let args = MergeArgs {
2011 output: out.clone(),
2012 base: Some(base_path),
2013 adapter: Some(adapter_dir.to_path_buf()),
2014 ..mk_args(MergeMethod::LoraAdapter)
2015 };
2016 run_lora_adapter_merge(&args, LogLevel::Quiet).unwrap();
2017
2018 let data = std::fs::read(&out).unwrap();
2020 let st = SafeTensors::deserialize(&data).unwrap();
2021 let merged =
2022 bytes_to_f32(st.tensor("w.weight").unwrap().data(), safetensors::tensor::Dtype::F32);
2023 assert!(
2025 (merged[0] - 4.0).abs() < 1e-4,
2026 "rsLoRA scale must be alpha/sqrt(rank)=4.0; got merged={} (main: 1.0 = alpha/rank)",
2027 merged[0]
2028 );
2029 }
2030
2031 #[test]
2038 fn test_lora_merge_bf16_adapter_dtype_pmat897() {
2039 let dir = tempfile::tempdir().unwrap();
2040 let adapter_dir = dir.path();
2041
2042 std::fs::write(
2043 adapter_dir.join("adapter_config.json"),
2044 r#"{"r": 4, "lora_alpha": 4, "use_rslora": false}"#,
2045 )
2046 .unwrap();
2047
2048 let base_path = adapter_dir.join("base.safetensors");
2050 write_st(
2051 &base_path,
2052 &[("w.weight", safetensors::tensor::Dtype::F32, vec![1, 1], f32_bytes(&[0.5]))],
2053 );
2054
2055 let a_vals = [2.0f32, 0.0, 0.0, 0.0];
2057 let b_vals = [3.0f32, 0.0, 0.0, 0.0];
2058 write_st(
2059 &adapter_dir.join("adapter_model.safetensors"),
2060 &[
2061 (
2062 "base_model.model.w.lora_A.weight",
2063 safetensors::tensor::Dtype::BF16,
2064 vec![4, 1],
2065 bf16_bytes(&a_vals),
2066 ),
2067 (
2068 "base_model.model.w.lora_B.weight",
2069 safetensors::tensor::Dtype::BF16,
2070 vec![1, 4],
2071 bf16_bytes(&b_vals),
2072 ),
2073 ],
2074 );
2075
2076 let out = dir.path().join("merged.safetensors");
2077 let args = MergeArgs {
2078 output: out.clone(),
2079 base: Some(base_path),
2080 adapter: Some(adapter_dir.to_path_buf()),
2081 ..mk_args(MergeMethod::LoraAdapter)
2082 };
2083 run_lora_adapter_merge(&args, LogLevel::Quiet).unwrap();
2084
2085 let data = std::fs::read(&out).unwrap();
2086 let st = SafeTensors::deserialize(&data).unwrap();
2087 let merged =
2088 bytes_to_f32(st.tensor("w.weight").unwrap().data(), safetensors::tensor::Dtype::F32);
2089 assert_eq!(merged.len(), 1, "merged w.weight must be 1 element");
2091 assert!(
2092 (merged[0] - 6.5).abs() < 1e-2,
2093 "BF16 adapter must decode to base+scale*(B@A)=6.5; got {} (main: BF16 read as f32 → garbage)",
2094 merged[0]
2095 );
2096 }
2097}