//! CSM dual-transformer model wrapper. //! //! Imports `candle_transformers::models::csm::Model` directly — Stage 1 does //! not vendor or fork. If we hit a blocker (need backbone hidden states for a //! probe, ring-buffer KV cache, etc.) Stage 2 may copy `csm.rs` here. use crate::config::{BackboneFlavor, DecoderFlavor, ModelConfig}; use crate::csm_fork as ccsm; use crate::csm_quantized as ccsmq; use crate::error::{CsmError, Result}; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use candle_transformers::generation::LogitsProcessor; use candle_transformers::quantized_var_builder::VarBuilder as QVarBuilder; use std::path::Path; pub type Inner = ccsm::Model; /// Backend dispatch over the two model variants. `Fp` uses memory-mapped /// safetensors weights at runtime dtype (BF16/F16/F32). `Quantized` loads /// from a GGUF emitted by `quantize::convert_to_quantized` with QMatMul on /// backbone projections. pub enum ModelBackend { Fp(ccsm::Model), Quantized(ccsmq::Model), } impl ModelBackend { pub fn cfg_enabled(&self) -> bool { match self { Self::Fp(m) => m.cfg_enabled(), Self::Quantized(m) => m.cfg_enabled(), } } pub fn clear_kv_cache(&mut self) { match self { Self::Fp(m) => m.clear_kv_cache(), Self::Quantized(m) => m.clear_kv_cache(), } } pub fn audio_tokens_and_mask( &self, frame: Vec, ) -> std::result::Result<(Tensor, Tensor), candle_core::Error> { match self { Self::Fp(m) => m.audio_tokens_and_mask(frame), Self::Quantized(m) => m.audio_tokens_and_mask(frame), } } pub fn text_tokens_and_mask( &self, ids: &[u32], ) -> std::result::Result<(Tensor, Tensor), candle_core::Error> { match self { Self::Fp(m) => m.text_tokens_and_mask(ids), Self::Quantized(m) => m.text_tokens_and_mask(ids), } } pub fn generate_frame( &mut self, tokens: &Tensor, mask: &Tensor, input_pos: usize, lp: &mut LogitsProcessor, ) -> std::result::Result, candle_core::Error> { match self { Self::Fp(m) => m.generate_frame(tokens, mask, input_pos, lp), Self::Quantized(m) => m.generate_frame(tokens, mask, input_pos, lp), } } #[allow(clippy::too_many_arguments)] pub fn generate_frame_cfg( &mut self, cond_tokens: &Tensor, cond_mask: &Tensor, cond_pos: usize, uncond_tokens: &Tensor, uncond_mask: &Tensor, uncond_pos: usize, cfg_scale: f64, lp: &mut LogitsProcessor, ) -> std::result::Result, candle_core::Error> { match self { Self::Fp(m) => m.generate_frame_cfg( cond_tokens, cond_mask, cond_pos, uncond_tokens, uncond_mask, uncond_pos, cfg_scale, lp, ), Self::Quantized(m) => m.generate_frame_cfg( cond_tokens, cond_mask, cond_pos, uncond_tokens, uncond_mask, uncond_pos, cfg_scale, lp, ), } } /// Inject LoRA adapters into the backbone. Works on both FP (for training /// and inference) and Quantized (inference-only) backends. The adapter /// delta path is the same in both cases — base output + LoRA delta. pub fn add_lora_to_backbone( &mut self, cfg: &crate::lora::LoraConfig, vm: &candle_nn::VarMap, ) -> std::result::Result<(), candle_core::Error> { match self { Self::Fp(m) => m.add_lora_to_backbone(cfg, vm), Self::Quantized(m) => m.add_lora_to_backbone(cfg, vm), } } /// Refresh LoRA adapter tensor handles after a VarMap mutation (e.g. /// AdamW step or `load_lora_adapter`). No-op for the quantized backend's /// LoraDelta if it wasn't injected, but harmless to call. pub fn refresh_lora( &mut self, vm: &candle_nn::VarMap, ) -> std::result::Result<(), candle_core::Error> { match self { Self::Fp(m) => m.refresh_lora(vm), Self::Quantized(m) => m.refresh_lora(vm), } } /// Teacher-forced loss for one frame. Currently only the FP backend /// implements this — the quantized backend returns a typed error since /// training through quantized weights isn't a supported workflow (you'd /// instead wrap LoRA adapters around the FP base for fine-tuning). pub fn forward_loss( &mut self, tokens: &Tensor, tokens_mask: &Tensor, input_pos: usize, target_codes: &[u32], ) -> std::result::Result { match self { Self::Fp(m) => m.forward_loss(tokens, tokens_mask, input_pos, target_codes), Self::Quantized(_) => Err(candle_core::Error::Msg( "forward_loss not implemented for quantized backend; use the FP path with LoRA wrapping for training".into(), )), } } /// Backbone activation capture for steering-vector extraction. /// FP-only — quantized path returns an error. pub fn capture_backbone_activations( &mut self, tokens: &Tensor, tokens_mask: &Tensor, ) -> std::result::Result, candle_core::Error> { match self { Self::Fp(m) => m .capture_backbone_activations(tokens, tokens_mask) .map_err(|e| candle_core::Error::Msg(e.to_string())), Self::Quantized(_) => Err(candle_core::Error::Msg( "capture_backbone_activations only supported on FP backbone".into(), )), } } /// Depth-decoder activation capture for steering-vector extraction. /// FP-only. pub fn capture_decoder_activations( &mut self, tokens: &Tensor, tokens_mask: &Tensor, target_c0: u32, ) -> std::result::Result, candle_core::Error> { match self { Self::Fp(m) => m .capture_decoder_activations(tokens, tokens_mask, target_c0) .map_err(|e| candle_core::Error::Msg(e.to_string())), Self::Quantized(_) => Err(candle_core::Error::Msg( "capture_decoder_activations only supported on FP backbone".into(), )), } } } pub struct CsmModel { pub inner: ModelBackend, pub config: ModelConfig, pub dtype: DType, pub device: Device, } impl CsmModel { pub fn load_from_safetensors>( path: P, config: ModelConfig, dtype: DType, device: &Device, ) -> Result { Self::load_from_safetensors_with_cfg(path, config, dtype, device, false) } /// Load CSM weights and optionally allocate a second backbone for CFG. pub fn load_from_safetensors_with_cfg>( path: P, config: ModelConfig, dtype: DType, device: &Device, enable_cfg: bool, ) -> Result { let path = path.as_ref(); tracing::info!( "loading CSM safetensors from {} (cfg={enable_cfg})", path.display() ); // SAFETY: memory-map the safetensors file — standard pattern in candle examples. let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[path], dtype, device)? }; let ccfg = to_candle_config(&config); let mut inner = ccsm::Model::new(&ccfg, vb.clone())?; if enable_cfg { inner.enable_cfg(vb)?; } Ok(Self { inner: ModelBackend::Fp(inner), config, dtype, device: device.clone(), }) } /// Load a quantized model from a GGUF file (the artifact emitted by /// `quantize::convert_to_quantized`). `runtime_dtype` is the precision /// for activations + dequantized kept-native tensors (F16 on Metal, F32 on CPU). pub fn load_from_gguf>( path: P, config: ModelConfig, runtime_dtype: DType, device: &Device, enable_cfg: bool, ) -> Result { let path = path.as_ref(); tracing::info!( "loading quantized CSM GGUF from {} (cfg={enable_cfg})", path.display() ); let qcfg = to_quantized_config(&config); let vb = QVarBuilder::from_gguf(path, device).map_err(|e| { CsmError::Other(anyhow::anyhow!( "QVarBuilder::from_gguf {}: {e}", path.display() )) })?; let mut inner = ccsmq::Model::new(&qcfg, runtime_dtype, vb.clone())?; if enable_cfg { inner.enable_cfg(runtime_dtype, vb)?; } Ok(Self { inner: ModelBackend::Quantized(inner), config, dtype: runtime_dtype, device: device.clone(), }) } pub fn clear_kv_cache(&mut self) { self.inner.clear_kv_cache(); } } pub fn to_candle_config(cfg: &ModelConfig) -> ccsm::Config { ccsm::Config { audio_num_codebooks: cfg.audio_num_codebooks, audio_vocab_size: cfg.audio_vocab_size, backbone_flavor: match cfg.backbone { BackboneFlavor::Llama1B => ccsm::Flavor::Llama1B, // Fork only implements Llama1B + Llama100M (matching candle upstream // and the released CSM weights). 3B/8B variants would require // extending the fork's `Flavor` enum. BackboneFlavor::Llama3B | BackboneFlavor::Llama8B => ccsm::Flavor::Llama1B, }, decoder_flavor: match cfg.decoder { DecoderFlavor::Llama100M => ccsm::Flavor::Llama100M, DecoderFlavor::Llama250M | DecoderFlavor::Llama300M => ccsm::Flavor::Llama100M, }, text_vocab_size: cfg.text_vocab_size, } } pub fn to_quantized_config(cfg: &ModelConfig) -> ccsmq::Config { ccsmq::Config { audio_num_codebooks: cfg.audio_num_codebooks, audio_vocab_size: cfg.audio_vocab_size, backbone_flavor: match cfg.backbone { BackboneFlavor::Llama1B => ccsmq::Flavor::Llama1B, BackboneFlavor::Llama3B | BackboneFlavor::Llama8B => ccsmq::Flavor::Llama1B, }, decoder_flavor: match cfg.decoder { DecoderFlavor::Llama100M => ccsmq::Flavor::Llama100M, DecoderFlavor::Llama250M | DecoderFlavor::Llama300M => ccsmq::Flavor::Llama100M, }, text_vocab_size: cfg.text_vocab_size, } } /// Read the safetensors header without loading tensors into device memory. /// Returns (name, shape, dtype) for every tensor in the file. Used by Step B /// to verify that the HuggingFace checkpoint uses the naming convention /// candle expects (`backbone.*`, `decoder.*`, `audio_embeddings.weight`, etc.). pub fn dump_safetensors_keys>(path: P) -> Result> { let path = path.as_ref(); let bytes = std::fs::read(path)?; let st = safetensors::SafeTensors::deserialize(&bytes)?; let mut out = Vec::with_capacity(st.names().len()); for name in st.names() { let info = st.tensor(name)?; out.push(TensorDescriptor { name: name.to_string(), shape: info.shape().to_vec(), dtype: format!("{:?}", info.dtype()), }); } out.sort_by(|a, b| a.name.cmp(&b.name)); Ok(out) } #[derive(Debug, Clone)] pub struct TensorDescriptor { pub name: String, pub shape: Vec, pub dtype: String, } impl TensorDescriptor { /// The critical head shapes — caller can check these match /// `(audio_vocab_size=2051, embed_dim)` for `codebook0_head.weight` etc. pub fn is_head(&self) -> bool { matches!( self.name.as_str(), "codebook0_head.weight" | "audio_head" | "audio_embeddings.weight" | "text_embeddings.weight" | "projection.weight" ) } } /// Sanity-check that the safetensors file contains the keys candle's CSM /// `Model::new` will ask for. Returns the list of **missing** keys (empty = OK). pub fn audit_csm_keys(descriptors: &[TensorDescriptor]) -> Vec { let required_exact = [ "audio_embeddings.weight", "text_embeddings.weight", "projection.weight", "codebook0_head.weight", "audio_head", ]; let present: std::collections::HashSet<&str> = descriptors.iter().map(|d| d.name.as_str()).collect(); let mut missing: Vec = required_exact .iter() .filter(|k| !present.contains(**k)) .map(|s| s.to_string()) .collect(); if !descriptors.iter().any(|d| d.name.starts_with("backbone.")) { missing.push("backbone.*".into()); } if !descriptors.iter().any(|d| d.name.starts_with("decoder.")) { missing.push("decoder.*".into()); } missing } impl std::fmt::Debug for CsmModel { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("CsmModel") .field("config", &self.config) .field("dtype", &self.dtype) .field("device", &self.device) .finish_non_exhaustive() } } impl From for candle_core::Error { fn from(e: CsmError) -> Self { candle_core::Error::Msg(e.to_string()) } }