Skip to main content

entrenar/cli/commands/
merge.rs

1//! Merge command implementation
2
3use 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    // ENT-LoRA-017: LoRA adapter merge path
17    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
32/// Log merge operation start
33fn 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
46/// Validate we have enough models
47fn 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
54/// Load all models from paths
55fn 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
71/// Load a single model from a SafeTensors file
72fn 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
88/// Extract a tensor as f32 values (returns None for non-F32 tensors)
89fn 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
105/// Perform the merge based on the specified method
106fn 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            // Handled by early return in run_merge; shouldn't reach here
114            Err("LoRA adapter merge uses dedicated path".to_string())
115        }
116    }
117}
118
119/// TIES merge: first model is base, rest are task-specific
120fn 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
127/// DARE merge with dropout
128fn 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
135/// SLERP merge (requires exactly 2 models)
136fn 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
144/// Average/ensemble merge with optional weights
145fn 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
150/// Build ensemble config from args
151fn 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
164/// Export merged model to file
165fn 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
175/// Export to SafeTensors format
176fn 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
204/// Build SafeTensors metadata
205fn 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
213/// Export to JSON format
214fn 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
224/// Log merge completion
225fn 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
233/// Merge LoRA adapter into base model (ENT-LoRA-017)
234///
235/// Computes W_merged = W_base + scale * B @ A for each adapted module,
236/// producing a standard safetensors model with no LoRA tensors.
237fn 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    // Read adapter config
259    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    // F-LORA-MERGE-RSLORA-001: honor PEFT/Unsloth `use_rslora`.
268    // rsLoRA scale = alpha/sqrt(rank); Standard = alpha/rank. Use the in-tree
269    // LoRAScaling enum so the merge matches LoRALayer's training-time scaling.
270    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    // Load base model
285    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    // Load adapter
290    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    // Merge: copy all base tensors, apply LoRA delta where adapters exist
295    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    // Build map of adapter A/B pairs grouped by module path
300    let lora_pairs = build_lora_pairs(&adapter_names, &adapter_tensors)?;
301    let mut merged_count = 0usize;
302
303    // Prepare output tensors
304    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        // Check if this weight has a LoRA adapter
311        if let Some(((a_data, a_shape, a_dtype), (b_data, b_shape, b_dtype))) =
312            lora_pairs.get(name.as_str())
313        {
314            // W_merged = W_base + scale * B @ A
315            // F-LORA-MERGE-ADAPTER-DTYPE-001: decode each factor with its OWN dtype
316            // (PEFT/Unsloth adapters are typically BF16/FP16), not a hardcoded f32.
317            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            // Compute B @ A: [d_out, r] @ [r, d_in] -> [d_out, d_in]
326            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            // W_merged = W_base + scale * BA
338            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            // Pass through base tensor unchanged
348            output_tensors.push((name.clone(), base_t.data().to_vec(), shape));
349        }
350    }
351
352    // Serialize to safetensors
353    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
392/// Per-tensor adapter factor: raw bytes, shape, and the ACTUAL safetensors dtype.
393///
394/// F-LORA-MERGE-ADAPTER-DTYPE-001: the dtype MUST be preserved (PEFT/Unsloth
395/// adapters are commonly BF16/FP16); decoding everything as f32 produces garbage.
396type LoraFactor = (Vec<u8>, Vec<usize>, safetensors::tensor::Dtype);
397
398/// A matched A/B adapter pair: (A_factor, B_factor).
399type LoraPair = (LoraFactor, LoraFactor);
400
401/// Build a map of base weight name -> (A_factor, B_factor), each carrying its dtype.
402fn 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        // PEFT naming: base_model.model.{path}.lora_A.weight / lora_B.weight
410        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            // Leak the base_name string to get a &'a str — safe in this context
435            // as the result lives only for the merge duration
436            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
443/// Convert tensor bytes to f32 based on dtype
444fn 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            // For other dtypes, treat as f32 (best effort)
465            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        // ties_merge needs base + at least 2 delta models (3 total)
684        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    // ── perform_merge routing tests ─────────────────────────────────────
832
833    #[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    // ── slerp merge with exactly 2 models ───────────────────────────────
901
902    #[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, // defaults to 0.5
926            density: None,
927            weights: None,
928            base: None,
929            adapter: None,
930        };
931        assert!(perform_slerp_merge(&ms, &a).is_ok());
932    }
933
934    // ── ties merge with density ─────────────────────────────────────────
935
936    #[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    // ── dare merge with density ─────────────────────────────────────────
955
956    #[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    // ── average merge with explicit weights ─────────────────────────────
973
974    #[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        // 0.8 * 10.0 + 0.2 * 0.0 = 8.0
990        assert!((vals[0] - 8.0).abs() < 1e-4);
991    }
992
993    // ── build_ensemble_config edge cases ────────────────────────────────
994
995    #[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        // Empty string should fail to parse as f32
1040        assert!(build_ensemble_config(&a).is_err());
1041    }
1042
1043    // ── validate_model_count edge cases ─────────────────────────────────
1044
1045    #[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    // ── export_merged_model extension detection ─────────────────────────
1076
1077    #[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        // Should fall through to JSON export (default)
1092        assert!(export_merged_model(&m, &a).is_ok());
1093        let _ = std::fs::remove_file(&t);
1094    }
1095
1096    // ── bytes_to_f32 additional edge cases ──────────────────────────────
1097
1098    #[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    // ── build_safetensor_metadata tests ─────────────────────────────────
1131
1132    #[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    // ── log_merge_start and log_merge_complete with different levels ────
1169
1170    #[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    // ── LoRA merge error paths ──────────────────────────────────────────
1202
1203    #[test]
1204    fn test_lora_adapter_config_not_found() {
1205        // adapter dir exists but no adapter_config.json inside
1206        let dir = tempfile::tempdir().unwrap();
1207        // Create a fake base file
1208        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        // Create adapter_config.json but not adapter_model.safetensors
1230        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    // ── run_merge with nonexistent model files ──────────────────────────
1247
1248    #[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    // ── mk helper verify ────────────────────────────────────────────────
1265
1266    #[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    // ── export safetensors with multiple tensors ────────────────────────
1277
1278    #[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        // Verify file was created and has content
1294        assert!(t.exists());
1295        let _ = std::fs::remove_file(&t);
1296    }
1297
1298    // ── export json roundtrip ───────────────────────────────────────────
1299
1300    #[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    // =========================================================================
1324    // test_cov2_* — Additional coverage tests
1325    // =========================================================================
1326
1327    /// Helper to build MergeArgs easily
1328    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    // ── bytes_to_f32 with zero-value data ────────────────────────────────
1342
1343    #[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    // ── bytes_to_f32 with truncated data (not aligned) ──────────────────
1409
1410    #[test]
1411    fn test_cov2_bytes_to_f32_f32_truncated() {
1412        // 5 bytes → only 1 full f32 chunk (4 bytes), remainder ignored
1413        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        // 3 bytes → only 1 full f16 chunk (2 bytes), remainder ignored
1422        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    // ── bytes_to_f32 with I64 fallback (other dtype) ────────────────────
1431
1432    #[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    // ── build_ensemble_config with whitespace-padded weights ────────────
1451
1452    #[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    // ── build_ensemble_config with negative weights ─────────────────────
1463
1464    #[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        // Parsing should succeed (negative floats are valid f32)
1470        assert!(config.is_ok());
1471    }
1472
1473    // ── build_ensemble_config with large number of weights ──────────────
1474
1475    #[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    // ── build_safetensor_metadata for each method ───────────────────────
1484
1485    #[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    // ── validate_model_count edge: exactly 2 ────────────────────────────
1520
1521    #[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    // ── perform_merge LoRA early error message ──────────────────────────
1538
1539    #[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    // ── export_merged_model to bad path ─────────────────────────────────
1547
1548    #[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    // ── export safetensors roundtrip ────────────────────────────────────
1573
1574    #[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        // Read back and verify
1581        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    // ── export json with empty model ────────────────────────────────────
1590
1591    #[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    // ── export safetensors with empty model ─────────────────────────────
1604
1605    #[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    // ── perform_ties_merge with default density ─────────────────────────
1615
1616    #[test]
1617    fn test_cov2_ties_merge_default_density() {
1618        let a = MergeArgs { density: None, ..mk_args(MergeMethod::Ties) };
1619        // density defaults to 0.2
1620        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    // ── perform_dare_merge with default density ─────────────────────────
1626
1627    #[test]
1628    fn test_cov2_dare_merge_default_density() {
1629        let a = MergeArgs { density: None, ..mk_args(MergeMethod::Dare) };
1630        // density defaults to 0.5 → drop_prob = 0.5
1631        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    // ── perform_slerp_merge with default weight ─────────────────────────
1636
1637    #[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    // ── slerp with single model → error ─────────────────────────────────
1646
1647    #[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    // ── run_merge with LoRA routes to lora function ─────────────────────
1656
1657    #[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    // ── load_single_model with empty file ───────────────────────────────
1682
1683    #[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    // ── load_single_model with garbage data ─────────────────────────────
1693
1694    #[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    // ── run_merge models don't exist → load error ───────────────────────
1704
1705    #[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    // ── log functions with all log levels ────────────────────────────────
1722
1723    #[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    // ── average merge with multiple tensors per model ───────────────────
1744
1745    #[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    // ── slerp merge with custom weight ──────────────────────────────────
1759
1760    #[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        // t=0 should give model 0's values
1767        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        // t=1 should give model 1's values
1777        assert!((vals[1] - 1.0).abs() < 0.1);
1778    }
1779
1780    // ── ties merge with explicit density close to 1.0 ───────────────────
1781
1782    #[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    // ── dare merge with density close to 0 ──────────────────────────────
1794
1795    #[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    // ── LoRA merge: adapter dir exists, has config, but no model.safetensors ─
1803
1804    #[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        // No adapter_model.safetensors
1813        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    // ── LoRA merge: base path doesn't exist ─────────────────────────────
1823
1824    #[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    // ── export extension detection ──────────────────────────────────────
1839
1840    #[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        // Verify file exists and is valid safetensors
1847        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        // Unknown extension → falls through to JSON
1869        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    // ── mk helper edge cases ────────────────────────────────────────────
1876
1877    #[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    // ── perform_merge dispatch coverage ─────────────────────────────────
1891
1892    #[test]
1893    fn test_cov2_perform_merge_all_methods() {
1894        // Ties
1895        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        // Dare
1900        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        // Slerp
1904        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        // Average
1912        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        // LoraAdapter → error
1916        assert!(perform_merge(&[], &mk_args(MergeMethod::LoraAdapter)).is_err());
1917    }
1918
1919    // ── large tensor export/import roundtrip ────────────────────────────
1920
1921    #[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    // ─────────────────────────────────────────────────────────────────────
1936    // PMAT-897: LoRA-adapter merge correctness falsifiers
1937    //   F-LORA-MERGE-RSLORA-001       — use_rslora scaling honored
1938    //   F-LORA-MERGE-ADAPTER-DTYPE-001 — adapter tensors decoded by real dtype
1939    // ─────────────────────────────────────────────────────────────────────
1940
1941    /// Serialize a single-tensor safetensors file with the given dtype/bytes.
1942    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    /// F-LORA-MERGE-RSLORA-001: `use_rslora: true` must select scale = alpha/sqrt(rank).
1962    ///
1963    /// rank=16, alpha=16 ⇒ rsLoRA scale = 16/sqrt(16) = 4.0 (NOT alpha/rank = 1.0).
1964    /// With a base of zeros and B@A == 1.0 in one cell, the merged weight there must
1965    /// equal scale (4.0). RED on main: scale hardcoded to alpha/rank = 1.0 → 4× too small.
1966    #[test]
1967    fn test_lora_merge_rslora_scaling_pmat897() {
1968        let dir = tempfile::tempdir().unwrap();
1969        let adapter_dir = dir.path();
1970
1971        // adapter_config.json: rank=16, alpha=16, use_rslora=true
1972        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        // Base weight "w.weight": shape [d_out=1, d_in=1], value 0.0
1979        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        // Adapter: lora_A [rank=16, d_in=1] all 0 except A[0]=1 ; lora_B [d_out=1, rank=16] all 0 except B[0]=1
1986        // ⇒ (B@A)[0,0] = sum_k B[0,k]*A[k,0] = 1*1 = 1.0
1987        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        // Read back merged "w.weight"
2019        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        // rsLoRA scale = alpha/sqrt(rank) = 16/4 = 4.0; base 0.0 + 4.0*(B@A=1.0) = 4.0
2024        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    /// F-LORA-MERGE-ADAPTER-DTYPE-001: BF16 adapter tensors must be decoded as BF16,
2032    /// not reinterpreted as f32.
2033    ///
2034    /// Standard scaling (rank=4, alpha=4 ⇒ scale=1.0). Adapter lora_A/lora_B are BF16.
2035    /// merged = base + 1.0*(B@A). RED on main: BF16 bytes read via hardcoded Dtype::F32
2036    /// → garbage (and len mismatch / wrong magnitude).
2037    #[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        // Base [d_out=1, d_in=1] = 0.5
2049        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        // BF16 adapter: A [4,1] = [2,0,0,0], B [1,4] = [3,0,0,0] ⇒ (B@A)[0,0] = 6.0
2056        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        // scale=1.0, base 0.5 + 1.0*(B@A=6.0) = 6.5
2090        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}