Files
rustytorch/crates/models/rtx-csm/src/steering.rs
T
2026-05-07 16:30:04 +00:00

224 lines
8.0 KiB
Rust

//! 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_<i>_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<Option<Tensor>>,
/// 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<usize> = 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<usize> {
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<Tensor> {
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_<i>_steering`. Missing layers stay None.
pub fn load_safetensors<P: AsRef<Path>>(
path: P,
num_layers: usize,
device: &Device,
) -> Result<Self> {
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);
}
}