224 lines
8.0 KiB
Rust
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);
|
|
}
|
|
}
|