//! Activation-steering vectors for the Llama backbone. //! //! Sprint 2 of the post-research roadmap (Phase 8 plan), inspired by //! EmoSteer-TTS (arXiv 2508.03543) and the broader ActAdd / contrastive //! activation-steering literature (Turner et al.). //! //! ## Why this differs from EmoSteer //! //! EmoSteer is specific to **flow-matching** TTS: it hooks DiT layers, //! exploits 32 CFM steps, and runs a per-token attribution search using //! mel-spectrogram synthesis. CSM is **autoregressive** over Mimi RVQ //! tokens, so the flow-matching plumbing doesn't apply. What is portable //! is the underlying technique: a per-layer **difference-in-means** //! activation vector added to the residual stream at inference shifts //! generation toward the contrast direction. //! //! Our design is the simpler ActAdd form: //! //! ```text //! steering vector u^l = mean(activation | emotion_X) - mean(activation | neutral) //! ``` //! //! At inference, after layer `l`'s residual stream is computed: //! //! ```text //! x^l <- x^l + alpha * u^l //! ``` //! //! Vectors are stored in safetensors with keys `layer__steering` and //! shape `(embed_dim,)`. A separate "extract" binary (Sprint 2 follow-up) //! produces these from a labeled corpus; this module just defines the //! type and the apply hook. use crate::error::{CsmError, Result}; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use std::path::Path; /// Per-layer steering vectors plus a global scale. Bind one of these to /// a `LlamaModel` to shift residual-stream activations during forward. #[derive(Debug, Clone)] pub struct LayerSteering { /// `vectors[i]` = `Some((1, embed_dim))` tensor added after layer `i`'s /// forward, or `None` to skip steering at that layer. Length must equal /// the model's `num_layers`. vectors: Vec>, /// Scalar multiplier applied to every non-None entry at inference. The /// EmoSteer paper uses 2.0 for emotion conversion, 2.5 for erasure; /// 0.0 effectively disables steering without dropping the layout. scale: f32, } impl LayerSteering { /// Empty steering — `apply()` is a no-op until vectors are inserted. pub fn empty(num_layers: usize) -> Self { Self { vectors: vec![None; num_layers], scale: 1.0, } } /// Set the steering vector for one layer. `vec` must be a 1-D tensor of /// length `embed_dim`; we reshape to `(1, embed_dim)` for broadcasting /// across the (batch, seq, embed) residual stream. pub fn set_layer(&mut self, layer_idx: usize, vec: Tensor) -> Result<()> { if layer_idx >= self.vectors.len() { return Err(CsmError::Config(format!( "layer_idx {layer_idx} out of range (num_layers={})", self.vectors.len() ))); } let dims = vec.dims(); let reshaped = match dims.len() { 1 => vec .reshape((1, dims[0])) .map_err(|e| CsmError::Config(e.to_string()))?, 2 if dims[0] == 1 => vec, _ => { return Err(CsmError::Config(format!( "steering vector must be 1-D or (1, embed_dim); got {dims:?}" ))); } }; self.vectors[layer_idx] = Some(reshaped); Ok(()) } pub fn set_scale(&mut self, scale: f32) { self.scale = scale; } pub fn scale(&self) -> f32 { self.scale } pub fn num_layers(&self) -> usize { self.vectors.len() } /// Restrict steering to a subset of layers; clears every layer not in /// `layers`. Useful for the EmoSteer paper's "spaced middle-to-deep" /// recipe — for our 16-layer backbone the analogue is roughly /// `[4, 8, 12]` or `[2, 6, 10, 14]`. Indices outside `[0, num_layers)` /// are silently ignored. pub fn restrict_to_layers(&mut self, layers: &[usize]) { let allow: std::collections::HashSet = layers.iter().copied().collect(); for (i, v) in self.vectors.iter_mut().enumerate() { if !allow.contains(&i) { *v = None; } } } /// Indices of layers that currently have a non-None steering vector. pub fn active_layers(&self) -> Vec { self.vectors .iter() .enumerate() .filter_map(|(i, v)| v.as_ref().map(|_| i)) .collect() } /// Apply this layer's steering vector to `xs` (residual stream after /// layer `layer_idx`). Returns `xs` unchanged if there's no vector for /// this layer or `scale == 0`. pub fn apply(&self, layer_idx: usize, xs: &Tensor) -> candle_core::Result { if self.scale == 0.0 { return Ok(xs.clone()); } let v = match self.vectors.get(layer_idx).and_then(|o| o.as_ref()) { Some(v) => v, None => return Ok(xs.clone()), }; // Match dtype + device. The model may be running bf16/f16; the // steering vector arrives as whatever the safetensors had. let v = v.to_dtype(xs.dtype())?.to_device(xs.device())?; // Broadcast: xs is (batch, seq, embed); v is (1, embed) — add via // reshape to (1, 1, embed). let dims = v.dims(); let v3 = v.reshape((1, 1, dims[dims.len() - 1]))?; let scaled = (v3 * self.scale as f64)?; xs.broadcast_add(&scaled) } /// Load steering vectors from a safetensors file. Expected keys: /// `layer__steering`. Missing layers stay None. pub fn load_safetensors>( path: P, num_layers: usize, device: &Device, ) -> Result { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[path.as_ref()], DType::F32, device) .map_err(|e| CsmError::Config(format!("steering safetensors: {e}")))? }; let mut steering = Self::empty(num_layers); for i in 0..num_layers { let key = format!("layer_{i}_steering"); if let Ok(t) = vb.get_unchecked(&key) { steering.set_layer(i, t)?; } } Ok(steering) } } #[cfg(test)] mod tests { use super::*; use candle_core::Device; #[test] fn empty_steering_is_noop() { let dev = Device::Cpu; let s = LayerSteering::empty(8); let xs = Tensor::ones((2, 3, 16), DType::F32, &dev).unwrap(); let out = s.apply(0, &xs).unwrap(); let diff = (&out - &xs).unwrap().abs().unwrap().sum_all().unwrap(); let diff_v: f32 = diff.to_scalar().unwrap(); assert!(diff_v.abs() < 1e-6); } #[test] fn set_layer_validates_size() { let dev = Device::Cpu; let mut s = LayerSteering::empty(4); let v = Tensor::ones((16,), DType::F32, &dev).unwrap(); s.set_layer(2, v).unwrap(); assert!(s.vectors[2].is_some()); assert!(s.vectors[0].is_none()); } #[test] fn restrict_clears_layers_outside_allowlist() { let dev = Device::Cpu; let mut s = LayerSteering::empty(8); for i in 0..8 { s.set_layer(i, Tensor::ones((4,), DType::F32, &dev).unwrap()) .unwrap(); } assert_eq!(s.active_layers(), vec![0, 1, 2, 3, 4, 5, 6, 7]); s.restrict_to_layers(&[2, 5]); assert_eq!(s.active_layers(), vec![2, 5]); } #[test] fn apply_shifts_by_scaled_vector() { let dev = Device::Cpu; let mut s = LayerSteering::empty(2); let v = Tensor::full(2.0f32, (16,), &dev).unwrap(); s.set_layer(0, v).unwrap(); s.set_scale(0.5); let xs = Tensor::ones((1, 3, 16), DType::F32, &dev).unwrap(); let out = s.apply(0, &xs).unwrap(); // Each element should be 1 + 0.5*2 = 2 let mean: f32 = out.mean_all().unwrap().to_scalar().unwrap(); assert!((mean - 2.0).abs() < 1e-5); // Other layers untouched let out1 = s.apply(1, &xs).unwrap(); let mean1: f32 = out1.mean_all().unwrap().to_scalar().unwrap(); assert!((mean1 - 1.0).abs() < 1e-5); } }