//! LoRA voice fine-tuning for CSM-1B. //! //! ## What ships //! //! - [`LoraConfig`] — the canonical hyperparameter bundle (rank, alpha, //! target_modules, dropout) with sensible defaults derived from the //! 2024–2025 LoRA-on-TTS literature (StyleSpeech, UtterTune, Koel-TTS). //! - [`LoraAdapter`] — a single rank-`r` adapter pair `(A, B)` for one //! target weight. `forward(&base_out, &xs)` adds the low-rank update. //! - [`LoraSet`] — collection of adapters keyed by target safetensors //! name (e.g., `backbone.layers.5.attn.q_proj.weight`). //! - [`merge_into_safetensors`] — **offline merge**: read base weights, //! read LoRA adapters, write out `W' = W + (B @ A) * (alpha / r)`. //! This is the production path — once merged, inference uses the //! existing un-modified `csm_fork::Model` with no per-call overhead. //! //! ## What's deferred //! //! The **training loop** itself. That requires: //! 1. A paired-data pipeline: list of `(text, audio_24khz)` for the target //! speaker, ~10–30 minutes total. //! 2. Forward pass that exposes per-codebook logits at training time //! (the current `generate_frame` samples internally; we'd need an //! `forward_loss` variant). //! 3. A loss function: cross-entropy on Mimi codes for c0..c31. //! 4. AdamW optimizer with `parking_lot::RwLock` parameter //! handles for the rank-r matrices. //! 5. Mixed precision (bf16 forward, f32 master weights for stability). //! 6. Maybe gradient checkpointing if backbone activations spill. //! //! Estimated training-loop effort: 1 week. Once that lands, an end-to-end //! voice-clone pipeline is `extract_audio → train_lora → merge → generate`. //! //! ## Recipe defaults (per literature) //! - rank: 8 //! - alpha: 16 (alpha/rank = 2 — modest update strength) //! - target: `q_proj` + `v_proj` on backbone ONLY (not decoder, not FFN) //! - dropout: 0.05 //! - learning rate: 1e-4, AdamW, cosine schedule //! - epochs: 3–5 on ~30 min of audio (~150-300 utterances) //! //! ## References //! - LoRA: Hu et al. arXiv:2106.09685 //! - StyleSpeech (TTS LoRA recipe): arXiv:2408.14713 //! - UtterTune: arXiv:2508.09767 use crate::error::{CsmError, Result}; use candle_core::{Module, Tensor}; use candle_nn::{Linear, VarBuilder, VarMap}; use std::collections::HashMap; use std::path::{Path, PathBuf}; // Hand-impl Debug since Tensor doesn't pretty-print. impl std::fmt::Debug for LoraDelta { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("LoraDelta") .field("a_shape", &self.a.shape()) .field("b_shape", &self.b.shape()) .field("scale", &self.scale) .finish() } } /// Additive-only LoRA delta: `delta(xs) = scale * (B @ A @ xs)`. /// /// Unlike [`LoraLinear`] this does NOT own a base layer — it's the bare /// adapter, intended to be added to an existing layer's output. This is the /// right shape for slotting LoRA into an existing model with private layer /// types: we leave the base layer untouched and just inject an additive path. /// /// Init follows LoRA paper: A ~ Randn(std=1/r), B = 0 (so initial delta is 0 /// and the model behaves identically to the un-adapted base until training). #[derive(Clone)] pub struct LoraDelta { pub a: Tensor, // (rank, in_dim) — trainable pub b: Tensor, // (out_dim, rank) — trainable pub scale: f64, } impl LoraDelta { #[allow(clippy::too_many_arguments)] pub fn new( rank: usize, alpha: f64, in_dim: usize, out_dim: usize, prefix: &str, vm: &VarMap, device: &candle_core::Device, _dtype: candle_core::DType, ) -> Result { // LoRA params live in F32 regardless of the surrounding model dtype: // candle's autograd is most reliable in F32, and the rank-r adapters // are tiny (~1 MB total) so the precision cost is negligible. We cast // to the input dtype at forward time. let dtype = candle_core::DType::F32; let scale = alpha / rank as f64; let std = 1.0 / (rank as f64); let init_a = candle_nn::Init::Randn { mean: 0.0, stdev: std, }; let init_b = candle_nn::Init::Const(0.0); let a = vm .get( (rank, in_dim), &format!("{prefix}.lora_a"), init_a, dtype, device, ) .map_err(|e| CsmError::Other(anyhow::anyhow!("vm lora_a: {e}")))?; let b = vm .get( (out_dim, rank), &format!("{prefix}.lora_b"), init_b, dtype, device, ) .map_err(|e| CsmError::Other(anyhow::anyhow!("vm lora_b: {e}")))?; Ok(Self { a, b, scale }) } /// Refresh `a` and `b` from a VarMap (post-optimizer-step). The underlying /// `Var` storage was updated by the optimizer; this struct holds plain /// Tensor handles, so we re-snapshot to see the new values on next forward. pub fn refresh_from(&mut self, vm: &VarMap, prefix: &str) -> Result<()> { let vars = vm.data().lock().unwrap(); if let Some(a) = vars.get(&format!("{prefix}.lora_a")) { self.a = a.as_tensor().clone(); } if let Some(b) = vars.get(&format!("{prefix}.lora_b")) { self.b = b.as_tensor().clone(); } Ok(()) } pub fn forward(&self, xs: &Tensor) -> candle_core::Result { // LoRA params live in F32 for stable autograd; activations may be // F16 (Metal) or F32 (CPU). Run the adapter math in F32 and cast // back to xs's dtype at the end so the addition with the base output // doesn't trip on a dtype mismatch. let target_dtype = xs.dtype(); let xs_f32 = xs.to_dtype(candle_core::DType::F32)?; // Broadcast A and B to match input rank. let a_t = match *xs_f32.dims() { [b1, b2, _, _] => self.a.t()?.broadcast_left((b1, b2))?, [bsize, _, _] => self.a.t()?.broadcast_left(bsize)?, _ => self.a.t()?, }; let b_t = match *xs_f32.dims() { [b1, b2, _, _] => self.b.t()?.broadcast_left((b1, b2))?, [bsize, _, _] => self.b.t()?.broadcast_left(bsize)?, _ => self.b.t()?, }; let xs_a = xs_f32.matmul(&a_t)?; let xs_ab = xs_a.matmul(&b_t)?; let scaled = (xs_ab * self.scale)?; scaled.to_dtype(target_dtype) } } /// LoRA-augmented linear layer: `out = base(xs) + scale * (B @ A @ xs)` where /// `base` is the frozen pre-trained weight and `A: (rank, in)`, `B: (out, rank)` /// are the small trainable matrices. /// /// At forward time we compose the contributions; at training time only `a` and /// `b` accumulate gradients (the base Linear is constructed with non-Var /// tensors so candle's autograd treats it as constant). /// /// Init convention follows the original LoRA paper: A is normal-sampled with /// std = 1/rank, B is zero — so the adapter starts as a no-op and the model /// behaves identically to the base until training begins. #[derive(Clone)] pub struct LoraLinear { pub base: Linear, pub a: Tensor, // (rank, in_dim) — trainable pub b: Tensor, // (out_dim, rank) — trainable pub scale: f64, pub rank: usize, } impl LoraLinear { /// Wrap an existing frozen `Linear` with a trainable rank-r LoRA adapter. /// The A/B params get registered into `vm` under `.lora_a` / /// `.lora_b` so AdamW (or any optimizer) can find and update them. #[allow(clippy::too_many_arguments)] pub fn wrap( base: Linear, rank: usize, alpha: f64, in_dim: usize, out_dim: usize, prefix: &str, vm: &VarMap, device: &candle_core::Device, dtype: candle_core::DType, ) -> Result { let scale = alpha / rank as f64; let std = 1.0 / (rank as f64); let init_a = candle_nn::Init::Randn { mean: 0.0, stdev: std, }; let init_b = candle_nn::Init::Const(0.0); let a = vm .get( (rank, in_dim), &format!("{prefix}.lora_a"), init_a, dtype, device, ) .map_err(|e| CsmError::Other(anyhow::anyhow!("vm lora_a: {e}")))?; let b = vm .get( (out_dim, rank), &format!("{prefix}.lora_b"), init_b, dtype, device, ) .map_err(|e| CsmError::Other(anyhow::anyhow!("vm lora_b: {e}")))?; Ok(Self { base, a, b, scale, rank, }) } } impl Module for LoraLinear { fn forward(&self, xs: &Tensor) -> candle_core::Result { let base_out = self.base.forward(xs)?; // xs (.., in) @ A.t (in, rank) → (.., rank) let xs_a = xs.matmul(&self.a.t()?)?; // (.., rank) @ B.t (rank, out) → (.., out) let xs_ab = xs_a.matmul(&self.b.t()?)?; let scaled = (xs_ab * self.scale)?; base_out + scaled } } /// Convenience: scan a VarBuilder path and `wrap` the named base linears with /// LoRA according to a `LoraConfig`. Used during model construction in the /// LoRA-aware training fork. Returned Vec is keyed by safetensors-name so /// callers can plug them back in as Module replacements. #[allow(dead_code)] pub fn build_lora_set( cfg: &LoraConfig, bases: &HashMap, vm: &VarMap, device: &candle_core::Device, dtype: candle_core::DType, ) -> Result> { let mut out = HashMap::new(); for (name, (base, in_dim, out_dim)) in bases.iter() { if !cfg.matches(name) { continue; } let lora = LoraLinear::wrap( base.clone(), cfg.rank, cfg.alpha as f64, *in_dim, *out_dim, name, vm, device, dtype, )?; out.insert(name.clone(), lora); } Ok(out) } // Suppress unused warning when the trainer path isn't compiled. #[allow(dead_code)] fn _vb_unused(_vb: &VarBuilder) {} #[derive(Debug, Clone)] pub struct LoraConfig { pub rank: usize, pub alpha: f32, /// Substring patterns matched against safetensors keys. A key containing /// any pattern is targeted. Default: `["q_proj", "v_proj"]` against /// `backbone.*` only (the Llama-3.2 1B backbone projections). pub target_modules: Vec, /// Patterns that EXCLUDE a target even if it matches `target_modules`. /// Default: `["decoder.", "audio_head", "codebook0_head"]`. pub exclude_patterns: Vec, /// Training-only dropout; not used at merge time. pub dropout: f32, } impl Default for LoraConfig { fn default() -> Self { Self { rank: 8, alpha: 16.0, target_modules: vec!["q_proj".into(), "v_proj".into()], exclude_patterns: vec![ "decoder.".into(), "audio_head".into(), "codebook0_head".into(), "audio_embeddings".into(), "text_embeddings".into(), "projection".into(), ], dropout: 0.05, } } } impl LoraConfig { pub fn scale(&self) -> f32 { self.alpha / self.rank as f32 } pub fn matches(&self, name: &str) -> bool { if self.exclude_patterns.iter().any(|p| name.contains(p)) { return false; } self.target_modules.iter().any(|p| name.contains(p)) } /// Extended coverage (Phase 12.1): all four attention projections plus the /// SwiGLU MLP. Use this when fine-tuning on hours of audio for stronger /// prosody adaptation; the q+v default is a safer floor when data is /// scarce (~30 min). Adapter param count for rank 8 / 16 backbone layers /// goes from ~1 MB (q+v) to ~6 MB (full) — still negligible vs the 1 B /// base model. Substring patterns match the safetensors keys produced by /// CSM's torchtune-style backbone (`q_proj`, `k_proj`, `v_proj`, /// `output_proj`, `mlp.w1`, `mlp.w2`, `mlp.w3`). pub fn extended() -> Self { Self { target_modules: vec![ "q_proj".into(), "k_proj".into(), "v_proj".into(), "output_proj".into(), "mlp.w1".into(), "mlp.w2".into(), "mlp.w3".into(), ], ..Self::default() } } } /// One LoRA adapter for one weight. `A` is rank × in_features, `B` is /// out_features × rank. The update is `B @ A * scale` added to the base weight. #[derive(Debug, Clone)] pub struct LoraAdapter { /// rank × in pub a: Vec, /// out × rank pub b: Vec, pub rank: usize, pub in_features: usize, pub out_features: usize, } impl LoraAdapter { pub fn new_zero(rank: usize, in_features: usize, out_features: usize) -> Self { Self { a: vec![0.0; rank * in_features], b: vec![0.0; out_features * rank], rank, in_features, out_features, } } /// Compute `B @ A` as a flat `out × in` matrix. Used by the offline merger. pub fn delta_w(&self, scale: f32) -> Vec { let mut out = vec![0.0f32; self.out_features * self.in_features]; for o in 0..self.out_features { for i in 0..self.in_features { let mut acc = 0.0f32; for r in 0..self.rank { acc += self.b[o * self.rank + r] * self.a[r * self.in_features + i]; } out[o * self.in_features + i] = acc * scale; } } out } } #[derive(Debug, Default, Clone)] pub struct LoraSet { pub config: LoraConfig, /// keyed by base safetensors weight name (e.g. `backbone.layers.5.attn.q_proj.weight`) pub adapters: HashMap, } impl LoraSet { pub fn new(config: LoraConfig) -> Self { Self { config, adapters: HashMap::new(), } } /// Load adapters from a directory containing `.a.f32` and /// `.b.f32` raw little-endian f32 dumps. Useful when training /// is done in Python and adapters are exported as plain bytes — avoids /// any safetensors/PyTorch coupling for v1. pub fn load_from_dir>(_dir: P, _config: LoraConfig) -> Result { // Stub. Once a training loop exists, fill this in to walk the dir, // pair `.a` / `.b` files, infer shapes from filename or sidecar JSON. Err(CsmError::Config( "LoraSet::load_from_dir: not yet implemented (training loop not yet shipped)".into(), )) } pub fn insert(&mut self, name: String, adapter: LoraAdapter) { self.adapters.insert(name, adapter); } } /// Read a trained LoRA adapter file (produced by /// `training::save_lora_adapter[_with_metadata]`) and reconstruct the /// `LoraSet` keyed by base-weight names. The adapter file's tensors are /// named `.lora_a` / `.lora_b`; we strip the suffix and /// pair them up. pub fn load_lora_set_from_safetensors>( path: P, config: LoraConfig, ) -> Result { let device = candle_core::Device::Cpu; let tensors = candle_core::safetensors::load(path.as_ref(), &device).map_err(|e| { CsmError::Other(anyhow::anyhow!( "load adapter {}: {e}", path.as_ref().display() )) })?; let mut a_map: HashMap = HashMap::new(); let mut b_map: HashMap = HashMap::new(); for (name, t) in tensors.into_iter() { if let Some(prefix) = name.strip_suffix(".lora_a") { a_map.insert(prefix.to_string(), t); } else if let Some(prefix) = name.strip_suffix(".lora_b") { b_map.insert(prefix.to_string(), t); } } let mut set = LoraSet::new(config); for (prefix, a_t) in a_map.into_iter() { let Some(b_t) = b_map.remove(&prefix) else { tracing::warn!("LoRA adapter `{prefix}.lora_a` has no matching `.lora_b` — skipping"); continue; }; // (rank, in_features) for A, (out_features, rank) for B let a_dims = a_t.dims(); let b_dims = b_t.dims(); if a_dims.len() != 2 || b_dims.len() != 2 || a_dims[0] != b_dims[1] { tracing::warn!( "LoRA adapter `{prefix}` shape mismatch: a={a_dims:?} b={b_dims:?} — skipping" ); continue; } let (rank, in_features) = (a_dims[0], a_dims[1]); let out_features = b_dims[0]; let a_vec: Vec = a_t .flatten_all() .and_then(|t| t.to_dtype(candle_core::DType::F32)) .and_then(|t| t.to_vec1()) .map_err(|e| CsmError::Other(anyhow::anyhow!("a→f32 flatten: {e}")))?; let b_vec: Vec = b_t .flatten_all() .and_then(|t| t.to_dtype(candle_core::DType::F32)) .and_then(|t| t.to_vec1()) .map_err(|e| CsmError::Other(anyhow::anyhow!("b→f32 flatten: {e}")))?; // Base-weight name is `.weight`. let weight_name = format!("{prefix}.weight"); set.insert( weight_name, LoraAdapter { a: a_vec, b: b_vec, rank, in_features, out_features, }, ); } if !b_map.is_empty() { tracing::warn!( "LoRA adapter has dangling .lora_b without .lora_a: {:?}", b_map.keys() ); } Ok(set) } /// **Offline merger**: read CSM safetensors, fold LoRA deltas into the /// targeted weights, write out a new safetensors file. Once merged, the /// existing `csm_fork::Model` and `Generator::load_csm_1b_from_path` /// paths use the merged checkpoint with zero per-inference overhead. /// /// `scale` is the LoRA scale (typically `alpha / rank`); pass the value /// the adapter was trained with — it's stored in the adapter's metadata /// when saved via `save_lora_adapter_with_metadata`. pub fn merge_into_safetensors>( base_safetensors: P, lora: &LoraSet, scale: f32, output_safetensors: P, ) -> Result { let device = candle_core::Device::Cpu; let base = candle_core::safetensors::load(base_safetensors.as_ref(), &device).map_err(|e| { CsmError::Other(anyhow::anyhow!( "load base {}: {e}", base_safetensors.as_ref().display() )) })?; let mut merged: HashMap = HashMap::new(); let mut report = MergeReport::default(); for (name, base_t) in base.into_iter() { if let Some(adapter) = lora.adapters.get(&name) { let dtype = base_t.dtype(); let shape = base_t.shape().clone(); // base → F32 for the additive merge let base_f32 = base_t .to_dtype(candle_core::DType::F32) .map_err(|e| CsmError::Other(anyhow::anyhow!("upcast {name}: {e}")))?; let dims = base_f32.dims(); if dims.len() != 2 || dims[0] != adapter.out_features || dims[1] != adapter.in_features { tracing::warn!( "LoRA shape mismatch for `{name}`: base={dims:?} adapter=({},{}) — passing through unchanged", adapter.out_features, adapter.in_features ); merged.insert(name, base_t); report.skipped += 1; continue; } let delta = adapter.delta_w(scale); let delta_t = candle_core::Tensor::from_vec( delta, (adapter.out_features, adapter.in_features), &device, ) .map_err(|e| CsmError::Other(anyhow::anyhow!("delta tensor: {e}")))?; let summed = base_f32 .add(&delta_t) .map_err(|e| CsmError::Other(anyhow::anyhow!("add delta: {e}")))?; // Cast back to the original dtype (F16/BF16) for storage parity. let out_t = summed .to_dtype(dtype) .and_then(|t| t.reshape(shape)) .map_err(|e| CsmError::Other(anyhow::anyhow!("downcast {name}: {e}")))?; merged.insert(name, out_t); report.merged += 1; } else { merged.insert(name, base_t); report.passthrough += 1; } } safetensors::serialize_to_file(&merged, None, output_safetensors.as_ref()) .map_err(|e| CsmError::Other(anyhow::anyhow!("save merged: {e}")))?; tracing::info!( "lora merge: {} merged, {} skipped, {} passthrough → {}", report.merged, report.skipped, report.passthrough, output_safetensors.as_ref().display() ); Ok(report) } #[derive(Debug, Default)] pub struct MergeReport { /// Number of base tensors that had LoRA deltas applied. pub merged: usize, /// Number of base tensors that LoRA targeted but had shape mismatches /// (passed through unchanged with a warn). pub skipped: usize, /// Number of base tensors with no LoRA target (byte-copied). pub passthrough: usize, /// Where the merged checkpoint was written. pub output_path: PathBuf, } #[cfg(test)] mod tests { use super::*; use candle_core::{DType, Device}; #[test] fn lora_linear_starts_as_noop_when_b_is_zero() { let dev = Device::Cpu; // Build a base Linear with a known weight: 4x3. let w = Tensor::from_slice( &[ 1.0f32, 2.0, 3.0, 0.5, -0.5, 1.0, 0.0, 0.0, 1.0, 2.0, 1.0, 0.5, ], (4, 3), &dev, ) .unwrap(); let base = Linear::new(w, None); let xs = Tensor::from_slice(&[1.0f32, 2.0, 3.0], (1, 3), &dev).unwrap(); let base_out = base.forward(&xs).unwrap().to_vec2::().unwrap(); // Wrap with LoraLinear; B init is zero so the adapter contribution is 0. let vm = VarMap::new(); let lora = LoraLinear::wrap(base, 2, 4.0, 3, 4, "test", &vm, &dev, DType::F32).unwrap(); let lora_out = lora.forward(&xs).unwrap().to_vec2::().unwrap(); for (a, b) in base_out[0].iter().zip(&lora_out[0]) { assert!((a - b).abs() < 1e-5, "no-op violated: {a} vs {b}"); } } #[test] fn lora_linear_diverges_after_perturbing_b() { let dev = Device::Cpu; let w = Tensor::zeros((4, 3), DType::F32, &dev).unwrap(); let base = Linear::new(w, None); let xs = Tensor::from_slice(&[1.0f32, 2.0, 3.0], (1, 3), &dev).unwrap(); let vm = VarMap::new(); let mut lora = LoraLinear::wrap(base, 2, 4.0, 3, 4, "test", &vm, &dev, DType::F32).unwrap(); // Manually overwrite B with non-zero values. lora.b = Tensor::from_slice(&[1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0, 0.0, 0.0], (4, 2), &dev).unwrap(); // And A with something concrete. lora.a = Tensor::from_slice(&[0.1f32, 0.2, 0.3, 0.0, 0.5, 0.0], (2, 3), &dev).unwrap(); let out = lora.forward(&xs).unwrap().to_vec2::().unwrap(); // Base is zero so output = scale * (B @ A @ xs) // A @ xs: [0.1+0.4+0.9, 0+1+0] = [1.4, 1.0] // B @ (A @ xs): [1.4, 0+1, 1.4+1.0, 0] = [1.4, 1.0, 2.4, 0] // scale = alpha/rank = 4/2 = 2.0 // result: [2.8, 2.0, 4.8, 0.0] let expected = [2.8_f32, 2.0, 4.8, 0.0]; for (got, exp) in out[0].iter().zip(&expected) { assert!((got - exp).abs() < 1e-4, "got {got} expected {exp}"); } } #[test] fn config_default_targets_q_v_only() { let c = LoraConfig::default(); assert!(c.matches("backbone.layers.0.attn.q_proj.weight")); assert!(c.matches("backbone.layers.0.attn.v_proj.weight")); assert!(!c.matches("backbone.layers.0.attn.k_proj.weight")); assert!(!c.matches("backbone.layers.0.attn.o_proj.weight")); // Decoder excluded assert!(!c.matches("decoder.layers.0.attn.q_proj.weight")); // Heads excluded assert!(!c.matches("codebook0_head.weight")); assert!(!c.matches("audio_embeddings.weight")); } #[test] fn config_extended_targets_full_attn_and_mlp() { let c = LoraConfig::extended(); // All four attention projections (note: real safetensors key is // `output_proj`, not `o_proj`). assert!(c.matches("backbone.layers.0.attn.q_proj.weight")); assert!(c.matches("backbone.layers.0.attn.k_proj.weight")); assert!(c.matches("backbone.layers.0.attn.v_proj.weight")); assert!(c.matches("backbone.layers.0.attn.output_proj.weight")); // SwiGLU MLP (Llama-style w1/w2/w3 naming). Patterns are scoped to // `mlp.wN` so they don't accidentally pick up unrelated `wN`-bearing // strings. assert!(c.matches("backbone.layers.0.mlp.w1.weight")); assert!(c.matches("backbone.layers.0.mlp.w2.weight")); assert!(c.matches("backbone.layers.0.mlp.w3.weight")); // Decoder + heads still excluded. assert!(!c.matches("decoder.layers.0.attn.q_proj.weight")); assert!(!c.matches("decoder.layers.0.mlp.w1.weight")); assert!(!c.matches("codebook0_head.weight")); assert!(!c.matches("projection.weight")); assert!(!c.matches("audio_embeddings.weight")); } #[test] fn scale_is_alpha_over_rank() { let c = LoraConfig::default(); assert!((c.scale() - 2.0).abs() < 1e-6); } #[test] fn adapter_zero_init_has_zero_delta() { let a = LoraAdapter::new_zero(8, 2048, 2048); let d = a.delta_w(2.0); assert!(d.iter().all(|&x| x == 0.0)); } #[test] fn adapter_nonzero_delta_shape() { let mut a = LoraAdapter::new_zero(2, 3, 4); // a (2x3) [[1,0,0],[0,1,0]] b (4x2) [[1,0],[0,1],[1,1],[2,0]] a.a = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0]; a.b = vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 2.0, 0.0]; let d = a.delta_w(1.0); // BA = b * a: // row0: [1,0]*[a] = [1,0,0] // row1: [0,1]*[a] = [0,1,0] // row2: [1,1]*[a] = [1,1,0] // row3: [2,0]*[a] = [2,0,0] assert_eq!( d, vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0, 2.0, 0.0, 0.0] ); } #[test] fn merge_errors_on_missing_base_file() { let lora = LoraSet::new(LoraConfig::default()); let r = merge_into_safetensors("/nonexistent", &lora, 1.0, "/nonexistent"); assert!(r.is_err()); } #[test] fn load_from_dir_returns_typed_error() { let r = LoraSet::load_from_dir("/nonexistent", LoraConfig::default()); assert!(r.is_err()); } }