Files
rustytorch/crates/models/rtx-csm/src/emotion2vec.rs
T
osobhandClaude Opus 4.7 a5cedfb46a rtx-csm: emotional_speech_guide — CREMA-D vs RAVDESS firdhokk verdict
8-gen bench (4 emotions × 2 corpora) at seed=42 against firdhokk
Whisper-LV3:

  target    RAVDESS              CREMA-D
  happy     happy (0.999) ✓      happy (0.999) ✓
  angry     neutral (0.92)       sad (0.99)
  fearful   happy (0.998)        fearful (0.984) ✓
  sad       angry (0.99)         fearful (0.99)

CREMA-D 2/4 vs RAVDESS 1/4. Larger / more naturalistic corpus
produces more class-pure fearful direction. Neither corpus solves
angry or sad — recipe shifts into 'vague expressivity' rather than
class-specific corners.

Practical: prefer CREMA-D when available; A/B both per emotion if
class precision matters.

Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
2026-04-30 00:01:02 -07:00

873 lines
34 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Phase 13.8 — emotion2vec_plus_base candle port.
//!
//! Multi-slice port of the FunASR/iic emotion2vec_plus_base model
//! (`emotion2vec/emotion2vec_plus_base` on HF Hub). Final goal: a real
//! [`crate::ser::EmotionDetector`] implementation that replaces the
//! prosody-rule placeholder used by `audio_to_manifest --auto-emotion-tag`
//! and `converse_server --reactive-emotion`.
//!
//! Architecture confirmed via `examples/emotion2vec_inspect`:
//! audio 16 kHz (1, T)
//! → local_encoder 7 × Conv1d (1→512 chans, T → T/320)
//! → project_features Linear 512 → 768
//! → relative_positional_encoder 5 × Conv1d (768→48, kernel=19)
//! → context_encoder 4-layer transformer prenet
//! → blocks.0..7 8-layer main transformer
//! → norm LayerNorm(768)
//! → mean-pool over time
//! → proj Linear 768 → 9
//! → softmax → 9-class probability
//!
//! Slicing:
//! - **Slice 1 (Phase 13.8)**: inspector + design notes ✅
//! - **Slice 2 (this module, growing)**: candle module scaffolding
//! - **Slice 3**: forward pass + shape verification on synthetic audio
//! - **Slice 4**: `EmotionDetector` impl + integration
//!
//! Full design notes live in `docs/emotion2vec_port_notes.md`.
use crate::error::Result as CsmResult;
use candle_core::{Module, Tensor};
use candle_nn::{
Conv1d, Conv1dConfig, LayerNorm, Linear, VarBuilder, conv1d_no_bias, layer_norm, linear,
};
/// Top-level architectural hyperparameters from `config.yaml`. Only the
/// fields needed by the inference path are surfaced; training-time fields
/// (mask probabilities, EMA decay, drop-path schedules, etc.) are
/// deliberately omitted because they have no effect on `forward()`.
#[derive(Debug, Clone)]
pub struct Emotion2VecConfig {
/// Hidden dim of every transformer block (768 for `_base`).
pub embed_dim: usize,
/// Heads in every transformer block (12 for `_base`).
pub num_heads: usize,
/// MLP expansion ratio (4.0 for `_base` → 768→3072→768).
pub mlp_ratio: f32,
/// Main encoder depth — `blocks.0..depth-1` (8 for `_base`).
pub depth: usize,
/// Prenet (`context_encoder`) depth (4 for `_base`).
pub prenet_depth: usize,
/// Number of output emotion classes (9 for `_base`):
/// angry/disgusted/fearful/happy/neutral/other/sad/surprised/`<unk>`.
pub num_classes: usize,
/// LayerNorm epsilon (1e-5).
pub norm_eps: f64,
/// Output channels of the local conv encoder (512 for `_base`).
pub feature_dim: usize,
/// Conv stack spec: `(out_channels, kernel_size, stride)` per layer.
/// For `_base`: `[(512,10,5), (512,3,2), (512,3,2), (512,3,2),
/// (512,3,2), (512,2,2), (512,2,2)]` — total stride product = 320,
/// so 16 kHz → 50 Hz feature rate (one frame per 20 ms).
pub conv_layers: Vec<(usize, usize, usize)>,
}
impl Emotion2VecConfig {
pub fn plus_base() -> Self {
Self {
embed_dim: 768,
num_heads: 12,
mlp_ratio: 4.0,
depth: 8,
prenet_depth: 4,
num_classes: 9,
norm_eps: 1e-5,
feature_dim: 512,
conv_layers: vec![
(512, 10, 5),
(512, 3, 2),
(512, 3, 2),
(512, 3, 2),
(512, 3, 2),
(512, 2, 2),
(512, 2, 2),
],
}
}
/// Total stride product of the conv stack — the audio-rate divisor.
/// 320 for `_base`: 16 kHz → 50 Hz feature frames.
pub fn conv_stride_total(&self) -> usize {
self.conv_layers.iter().map(|(_, _, s)| *s).product()
}
}
/// One Conv1d → LayerNorm → GELU block in the local encoder.
///
/// In the upstream PyTorch, each block is
/// `Sequential(Conv1d, Dropout, Sequential(TransposeLast, LayerNorm,
/// TransposeLast), GELU)`
/// with state-dict keys `.0.weight` (Conv1d, no bias), `.2.1.weight/bias`
/// (the LayerNorm). The Dropout has no params and is inactive at
/// inference, so we skip it. The TransposeLast wrappers move the channel
/// axis to last for LayerNorm and back; in candle we call `.transpose(1,2)`
/// before/after the norm.
#[derive(Debug)]
pub struct LocalConvBlock {
conv: Conv1d,
norm: LayerNorm,
}
impl LocalConvBlock {
/// Build from a VarBuilder rooted at e.g.
/// `d2v_model.modality_encoders.AUDIO.local_encoder.conv_layers.{i}`.
/// `vb.pp("0")` reaches the Conv1d weight; `vb.pp("2").pp("1")` reaches
/// the LayerNorm weight/bias. Mirrors the upstream nesting exactly.
pub fn new(
in_channels: usize,
out_channels: usize,
kernel_size: usize,
stride: usize,
norm_eps: f64,
vb: VarBuilder,
) -> CsmResult<Self> {
let cfg = Conv1dConfig {
stride,
..Default::default()
};
let conv = conv1d_no_bias(in_channels, out_channels, kernel_size, cfg, vb.pp("0"))
.map_err(|e| crate::CsmError::Config(format!("local conv: {e}")))?;
let norm = layer_norm(out_channels, norm_eps, vb.pp("2").pp("1"))
.map_err(|e| crate::CsmError::Config(format!("local norm: {e}")))?;
Ok(Self { conv, norm })
}
}
impl Module for LocalConvBlock {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
// (B, C_in, T_in)
let xs = self.conv.forward(xs)?;
// → (B, C_out, T_out). LayerNorm applies over the last dim, so
// move channels to the last axis, normalize, swap back.
let xs = xs.transpose(1, 2)?.contiguous()?; // (B, T, C)
let xs = self.norm.forward(&xs)?;
let xs = xs.transpose(1, 2)?.contiguous()?; // (B, C, T)
xs.gelu()
}
}
/// The 7-block local conv feature extractor.
///
/// Input: 16 kHz mono PCM tensor of shape `(B, 1, T)`.
/// Output: `(B, feature_dim, T / conv_stride_total())` — 512 channels,
/// 50 Hz frame rate for `_base` (T → T/320).
#[derive(Debug)]
pub struct LocalEncoder {
blocks: Vec<LocalConvBlock>,
}
impl LocalEncoder {
/// Build from a VarBuilder rooted at
/// `d2v_model.modality_encoders.AUDIO.local_encoder.conv_layers`.
pub fn new(cfg: &Emotion2VecConfig, vb: VarBuilder) -> CsmResult<Self> {
let mut blocks = Vec::with_capacity(cfg.conv_layers.len());
let mut in_ch = 1usize; // raw audio is mono
for (i, (out_ch, k, s)) in cfg.conv_layers.iter().enumerate() {
let block =
LocalConvBlock::new(in_ch, *out_ch, *k, *s, cfg.norm_eps, vb.pp(i.to_string()))?;
blocks.push(block);
in_ch = *out_ch;
}
Ok(Self { blocks })
}
}
impl Module for LocalEncoder {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
let mut xs = xs.clone();
for block in self.blocks.iter() {
xs = block.forward(&xs)?;
}
Ok(xs)
}
}
/// Linear `feature_dim → embed_dim` (512 → 768 for `_base`) — projects the
/// LocalEncoder's per-frame output into the transformer's working
/// dimension. State-dict keys: `project_features.2.weight/bias` (path
/// includes a LayerNorm at index 1 that we surface separately).
///
/// Per the inspector output the upstream stores:
/// `project_features.1.weight/bias` → LayerNorm(512)
/// `project_features.2.weight/bias` → Linear(512 → 768)
///
/// So the projection module is `LayerNorm → Linear`, both at the feature
/// dim (channel-last after the LocalEncoder transposes back to (B, T, C)
/// for the transformer).
#[derive(Debug)]
pub struct ProjectFeatures {
norm: LayerNorm,
proj: Linear,
}
impl ProjectFeatures {
/// Build from a VarBuilder rooted at
/// `d2v_model.modality_encoders.AUDIO.project_features`.
pub fn new(in_dim: usize, out_dim: usize, norm_eps: f64, vb: VarBuilder) -> CsmResult<Self> {
let norm = layer_norm(in_dim, norm_eps, vb.pp("1"))
.map_err(|e| crate::CsmError::Config(format!("project_features norm: {e}")))?;
let proj = linear(in_dim, out_dim, vb.pp("2"))
.map_err(|e| crate::CsmError::Config(format!("project_features linear: {e}")))?;
Ok(Self { norm, proj })
}
}
impl Module for ProjectFeatures {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
// xs: (B, T, in_dim) — LocalEncoder output already transposed
// to channel-last. Both ops are channel-last.
let xs = self.norm.forward(xs)?;
self.proj.forward(&xs)
}
}
/// One transformer block — pre-norm, fused QKV attention, MLP residual.
///
/// State-dict layout (per `blocks.{i}` and `context_encoder.blocks.{i}`):
/// `norm1.weight/bias` LayerNorm(embed_dim)
/// `attn.qkv.weight/bias` Linear(embed_dim → 3 × embed_dim)
/// `attn.proj.weight/bias` Linear(embed_dim → embed_dim)
/// `norm2.weight/bias` LayerNorm(embed_dim)
/// `mlp.fc1.weight/bias` Linear(embed_dim → mlp_dim)
/// `mlp.fc2.weight/bias` Linear(mlp_dim → embed_dim)
///
/// MLP activation is GELU. No positional encoding is added inside the
/// block — the upstream applies a conv-based positional bias at the
/// encoder boundary (slice 2c).
#[derive(Debug)]
pub struct Block {
norm1: LayerNorm,
qkv: Linear,
attn_proj: Linear,
norm2: LayerNorm,
fc1: Linear,
fc2: Linear,
num_heads: usize,
head_dim: usize,
scale: f64,
}
impl Block {
/// Build from a VarBuilder rooted at e.g. `d2v_model.blocks.{i}` or
/// `d2v_model.modality_encoders.AUDIO.context_encoder.blocks.{i}`.
pub fn new(
embed_dim: usize,
num_heads: usize,
mlp_dim: usize,
norm_eps: f64,
vb: VarBuilder,
) -> CsmResult<Self> {
let head_dim = embed_dim / num_heads;
let scale = 1.0 / (head_dim as f64).sqrt();
let norm1 = layer_norm(embed_dim, norm_eps, vb.pp("norm1"))
.map_err(|e| crate::CsmError::Config(format!("block norm1: {e}")))?;
let qkv = linear(embed_dim, embed_dim * 3, vb.pp("attn").pp("qkv"))
.map_err(|e| crate::CsmError::Config(format!("block qkv: {e}")))?;
let attn_proj = linear(embed_dim, embed_dim, vb.pp("attn").pp("proj"))
.map_err(|e| crate::CsmError::Config(format!("block attn proj: {e}")))?;
let norm2 = layer_norm(embed_dim, norm_eps, vb.pp("norm2"))
.map_err(|e| crate::CsmError::Config(format!("block norm2: {e}")))?;
let fc1 = linear(embed_dim, mlp_dim, vb.pp("mlp").pp("fc1"))
.map_err(|e| crate::CsmError::Config(format!("block fc1: {e}")))?;
let fc2 = linear(mlp_dim, embed_dim, vb.pp("mlp").pp("fc2"))
.map_err(|e| crate::CsmError::Config(format!("block fc2: {e}")))?;
Ok(Self {
norm1,
qkv,
attn_proj,
norm2,
fc1,
fc2,
num_heads,
head_dim,
scale,
})
}
fn attention(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
// xs: (B, T, D)
let (b, t, d) = xs.dims3()?;
let qkv = self.qkv.forward(xs)?; // (B, T, 3D)
// Reshape to (B, T, 3, H, head_dim), permute to (3, B, H, T, hd)
let qkv = qkv
.reshape((b, t, 3, self.num_heads, self.head_dim))?
.permute((2, 0, 3, 1, 4))?
.contiguous()?;
// Split into Q, K, V along the leading 3-axis.
let q = qkv.narrow(0, 0, 1)?.squeeze(0)?; // (B, H, T, hd)
let k = qkv.narrow(0, 1, 1)?.squeeze(0)?;
let v = qkv.narrow(0, 2, 1)?.squeeze(0)?;
// Attention scores. For Metal compatibility (the 4D matmul issue
// we hit on Moonshine, Phase 8.8), collapse (B, H) to a single
// batch dim before matmul.
let bh = b * self.num_heads;
let q3 = q.reshape((bh, t, self.head_dim))?;
let k3 = k.reshape((bh, t, self.head_dim))?;
let v3 = v.reshape((bh, t, self.head_dim))?;
let scores = (q3.matmul(&k3.transpose(1, 2)?)? * self.scale)?;
let attn = candle_nn::ops::softmax_last_dim(&scores)?;
let out = attn.matmul(&v3)?; // (B*H, T, hd)
let out = out
.reshape((b, self.num_heads, t, self.head_dim))?
.transpose(1, 2)?
.reshape((b, t, d))?
.contiguous()?;
self.attn_proj.forward(&out)
}
fn mlp(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
let xs = self.fc1.forward(xs)?;
let xs = xs.gelu()?;
self.fc2.forward(&xs)
}
}
impl Module for Block {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
// Pre-norm: x = x + attn(norm1(x)); x = x + mlp(norm2(x))
let h = self.norm1.forward(xs)?;
let h = self.attention(&h)?;
let xs = (xs + h)?;
let h = self.norm2.forward(&xs)?;
let h = self.mlp(&h)?;
xs + h
}
}
/// 4-layer transformer prenet sitting between `ProjectFeatures` (or
/// rather: features + positional bias) and the main 8-layer encoder.
/// Has its own final LayerNorm — pickle keys
/// `context_encoder.blocks.0..3.*` + `context_encoder.norm.weight/bias`.
#[derive(Debug)]
pub struct ContextEncoder {
blocks: Vec<Block>,
norm: LayerNorm,
}
impl ContextEncoder {
/// Build from a VarBuilder rooted at
/// `d2v_model.modality_encoders.AUDIO.context_encoder`.
pub fn new(cfg: &Emotion2VecConfig, vb: VarBuilder) -> CsmResult<Self> {
let mlp_dim = (cfg.embed_dim as f32 * cfg.mlp_ratio) as usize;
let mut blocks = Vec::with_capacity(cfg.prenet_depth);
let vb_blocks = vb.pp("blocks");
for i in 0..cfg.prenet_depth {
blocks.push(Block::new(
cfg.embed_dim,
cfg.num_heads,
mlp_dim,
cfg.norm_eps,
vb_blocks.pp(i.to_string()),
)?);
}
let norm = layer_norm(cfg.embed_dim, cfg.norm_eps, vb.pp("norm"))
.map_err(|e| crate::CsmError::Config(format!("context_encoder norm: {e}")))?;
Ok(Self { blocks, norm })
}
}
impl Module for ContextEncoder {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
let mut xs = xs.clone();
for block in self.blocks.iter() {
xs = block.forward(&xs)?;
}
self.norm.forward(&xs)
}
}
/// 8-layer main transformer encoder (`d2v_model.blocks.0..7`). No final
/// LayerNorm — the upstream pickle has none, and the pre-norm pattern's
/// per-block `norm2` keeps residuals well-conditioned without one.
#[derive(Debug)]
pub struct MainEncoder {
blocks: Vec<Block>,
}
impl MainEncoder {
/// Build from a VarBuilder rooted at `d2v_model.blocks`.
pub fn new(cfg: &Emotion2VecConfig, vb: VarBuilder) -> CsmResult<Self> {
let mlp_dim = (cfg.embed_dim as f32 * cfg.mlp_ratio) as usize;
let mut blocks = Vec::with_capacity(cfg.depth);
for i in 0..cfg.depth {
blocks.push(Block::new(
cfg.embed_dim,
cfg.num_heads,
mlp_dim,
cfg.norm_eps,
vb.pp(i.to_string()),
)?);
}
Ok(Self { blocks })
}
}
impl Module for MainEncoder {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
let mut xs = xs.clone();
for block in self.blocks.iter() {
xs = block.forward(&xs)?;
}
Ok(xs)
}
}
/// Final 9-class emotion classifier — a single Linear layer applied to
/// the time-mean-pooled hidden state. No LayerNorm; pickle keys are just
/// `proj.weight` (`[9, 768]`) and `proj.bias` (`[9]`) at the TOP level
/// of the state dict (NOT under the `d2v_model.` prefix).
///
/// Output classes (per `tokens.txt`):
/// 0=angry 1=disgusted 2=fearful 3=happy 4=neutral 5=other 6=sad
/// 7=surprised 8=`<unk>`. Map to [`crate::ser::EmotionLabel`] via
/// [`Self::tag_for_class`].
#[derive(Debug)]
pub struct Classifier {
proj: Linear,
}
impl Classifier {
/// Build from a VarBuilder rooted at the *root* of the model state
/// dict (NOT at `d2v_model.`) — `proj.weight/bias` live there.
pub fn new(embed_dim: usize, num_classes: usize, vb: VarBuilder) -> CsmResult<Self> {
let proj = linear(embed_dim, num_classes, vb.pp("proj"))
.map_err(|e| crate::CsmError::Config(format!("classifier proj: {e}")))?;
Ok(Self { proj })
}
/// Map a 0..=8 emotion2vec class index directly to one of the 9 raw
/// [`crate::ser::EmotionLabel`] variants — NO fold.
///
/// The earlier 9→5 fold (`happy`/`surprised`/`other` → Excited;
/// `disgusted`/`fearful` → Sad) collapsed emotion2vec's resolution
/// to a single tag in practice: motivational speech, technical
/// talks, and movie dialogue all bucketed to `[excited]`. With
/// raw 9-class tags the LoRA trainer sees real emotional variety
/// in its prompts.
pub fn tag_for_class(idx: u32) -> crate::ser::EmotionLabel {
match idx {
0 => crate::ser::EmotionLabel::Angry,
1 => crate::ser::EmotionLabel::Disgusted,
2 => crate::ser::EmotionLabel::Fearful,
3 => crate::ser::EmotionLabel::Happy,
4 => crate::ser::EmotionLabel::Neutral,
5 => crate::ser::EmotionLabel::Excited, // emotion2vec's "other"
6 => crate::ser::EmotionLabel::Sad,
7 => crate::ser::EmotionLabel::Surprised,
_ => crate::ser::EmotionLabel::Unk, // 8 = <unk> + any unknown
}
}
}
impl Module for Classifier {
fn forward(&self, pooled: &Tensor) -> candle_core::Result<Tensor> {
// pooled: (B, embed_dim) — the time-mean-pooled hidden state.
// Output: (B, num_classes) raw logits; caller applies softmax.
self.proj.forward(pooled)
}
}
/// Conv-based relative positional encoder. 5 grouped Conv1d layers with
/// kernel 19 and groups 16 — produces a positional bias of the same
/// shape as the input which is added back into the sequence.
///
/// State-dict keys: `relative_positional_encoder.{1..=5}.0.weight/bias`.
/// Indices are 1-based (no `.0.*` exists in the upstream pickle); the
/// inner `.0` is the Conv1d in a Sequential wrapper that also includes
/// a GELU at `.1` (no params).
///
/// The encoder uses **same-padding**: kernel 19 with `(19-1)/2 = 9`
/// padding on each side preserves the time dim through every layer.
#[derive(Debug)]
pub struct RelativePositionalEncoder {
convs: Vec<Conv1d>,
}
impl RelativePositionalEncoder {
/// Build from a VarBuilder rooted at
/// `d2v_model.modality_encoders.AUDIO.relative_positional_encoder`.
pub fn new(
embed_dim: usize,
depth: usize,
kernel: usize,
groups: usize,
vb: VarBuilder,
) -> CsmResult<Self> {
if kernel.is_multiple_of(2) {
return Err(crate::CsmError::Config(format!(
"relative_positional_encoder kernel must be odd for same-padding, got {kernel}"
)));
}
let cfg = Conv1dConfig {
padding: (kernel - 1) / 2,
stride: 1,
dilation: 1,
groups,
cudnn_fwd_algo: None,
};
let mut convs = Vec::with_capacity(depth);
// Upstream indexing: 1..=depth (no `.0`).
for i in 1..=depth {
let conv = candle_nn::conv1d(
embed_dim,
embed_dim,
kernel,
cfg,
vb.pp(i.to_string()).pp("0"),
)
.map_err(|e| crate::CsmError::Config(format!("rel_pos_enc layer {i}: {e}")))?;
convs.push(conv);
}
Ok(Self { convs })
}
}
impl Module for RelativePositionalEncoder {
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
// xs: (B, T, C). Conv1d operates on (B, C, T), so transpose.
let mut h = xs.transpose(1, 2)?.contiguous()?;
for conv in self.convs.iter() {
h = conv.forward(&h)?;
h = h.gelu()?;
}
// Transpose back; the result is the positional bias added to xs.
let bias = h.transpose(1, 2)?.contiguous()?;
xs + bias
}
}
/// Emotion2Vec stores its inference device alongside the modules so
/// [`crate::ser::EmotionDetector::classify`] (which only takes `&[f32]`)
/// can build input tensors without an out-of-band device handle.
fn _device_struct_doc() {}
/// Top-level emotion2vec_plus_base model. Wires:
/// 1. [`LocalEncoder`] (raw 16 kHz audio → 50 Hz, 512-d features)
/// 2. [`ProjectFeatures`] (LN + 512→768 Linear)
/// 3. [`RelativePositionalEncoder`] (additive conv positional bias)
/// 4. [`ContextEncoder`] (4-block prenet + LN)
/// 5. [`MainEncoder`] (8 blocks)
/// 6. mean-pool over the time axis
/// 7. [`Classifier`] (Linear 768→9)
///
/// Input: 16 kHz mono PCM as a `(B, 1, T)` tensor.
/// Output: `(B, 9)` raw logits — caller applies softmax + argmax.
#[derive(Debug)]
pub struct Emotion2Vec {
cfg: Emotion2VecConfig,
local_encoder: LocalEncoder,
project_features: ProjectFeatures,
rel_pos_enc: RelativePositionalEncoder,
context_encoder: ContextEncoder,
main_encoder: MainEncoder,
classifier: Classifier,
device: candle_core::Device,
}
impl Emotion2Vec {
/// Build with random init from a VarBuilder rooted at the state-dict
/// root. Useful for tests; production use should call
/// [`Self::load_from_pickle`].
///
/// `kernel` and `groups` for the relative positional encoder default
/// to the `_base` config (kernel 19, groups 16).
pub fn new(
cfg: Emotion2VecConfig,
vb: VarBuilder,
device: candle_core::Device,
) -> CsmResult<Self> {
// d2v_model.* sub-builder
let vb_d = vb.pp("d2v_model");
let vb_mod = vb_d.pp("modality_encoders").pp("AUDIO");
let local_encoder = LocalEncoder::new(&cfg, vb_mod.pp("local_encoder").pp("conv_layers"))?;
let project_features = ProjectFeatures::new(
cfg.feature_dim,
cfg.embed_dim,
cfg.norm_eps,
vb_mod.pp("project_features"),
)?;
let rel_pos_enc = RelativePositionalEncoder::new(
cfg.embed_dim,
5, // _base: 5 conv layers
19, // kernel
16, // groups
vb_mod.pp("relative_positional_encoder"),
)?;
let context_encoder = ContextEncoder::new(&cfg, vb_mod.pp("context_encoder"))?;
let main_encoder = MainEncoder::new(&cfg, vb_d.pp("blocks"))?;
// proj.* lives at the ROOT of the state dict, so we use `vb`,
// not `vb_d`.
let classifier = Classifier::new(cfg.embed_dim, cfg.num_classes, vb)?;
Ok(Self {
cfg,
local_encoder,
project_features,
rel_pos_enc,
context_encoder,
main_encoder,
classifier,
device,
})
}
/// Load the upstream `emotion2vec/emotion2vec_plus_base` checkpoint
/// from a local `.pt` path. The pickle is a fairseq-style nested
/// dict; the actual state lives under the `"model"` key.
pub fn load_from_pickle<P: AsRef<std::path::Path>>(
path: P,
device: &candle_core::Device,
) -> CsmResult<Self> {
let vb = VarBuilder::from_pth_with_state(
path.as_ref(),
candle_core::DType::F32,
"model",
device,
)
.map_err(|e| crate::CsmError::Config(format!("emotion2vec pth load: {e}")))?;
Self::new(Emotion2VecConfig::plus_base(), vb, device.clone())
}
/// Forward pass: raw audio → 9-class logits.
pub fn forward(&self, audio_16k: &Tensor) -> candle_core::Result<Tensor> {
// (B, 1, T) → (B, 512, T/320)
let feats = self.local_encoder.forward(audio_16k)?;
// Transpose to channel-last for the transformer path.
let feats = feats.transpose(1, 2)?.contiguous()?;
// → (B, T', 768)
let h = self.project_features.forward(&feats)?;
// Add conv-based positional bias.
let h = self.rel_pos_enc.forward(&h)?;
// 4-block prenet.
let h = self.context_encoder.forward(&h)?;
// 8-block main encoder.
let h = self.main_encoder.forward(&h)?;
// Mean-pool over time → (B, 768).
let pooled = h.mean(1)?;
// (B, 9) logits.
self.classifier.forward(&pooled)
}
pub fn config(&self) -> &Emotion2VecConfig {
&self.cfg
}
pub fn device(&self) -> &candle_core::Device {
&self.device
}
}
/// Implementation of the [`crate::ser::EmotionDetector`] trait so
/// `Emotion2Vec` is a drop-in replacement for the Phase 13.3
/// [`crate::ser::ProsodyDetector`] placeholder. Used by
/// `audio_to_manifest --auto-emotion-tag` and
/// `converse_server --reactive-emotion` once the user opts in via
/// `--use-emotion2vec`.
impl crate::ser::EmotionDetector for Emotion2Vec {
fn classify(&self, samples_16k: &[f32]) -> CsmResult<crate::ser::EmotionLabel> {
if samples_16k.is_empty() {
return Ok(crate::ser::EmotionLabel::Neutral);
}
// Per-utterance zero-mean unit-variance normalization
// (config.yaml: `normalize: true`). Without this the model
// produces near-constant logits regardless of input — every
// clip in the corpus would get the same label. Verified
// empirically via examples/emotion2vec_probe before/after.
let n = samples_16k.len();
let mean = samples_16k.iter().sum::<f32>() / n as f32;
let var = samples_16k.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / n as f32;
let std = var.sqrt().max(1e-7);
let normed: Vec<f32> = samples_16k.iter().map(|x| (x - mean) / std).collect();
let audio = Tensor::from_vec(normed, (1, 1, n), &self.device)
.map_err(|e| crate::CsmError::Config(format!("emotion2vec audio tensor: {e}")))?;
let logits = self
.forward(&audio)
.map_err(|e| crate::CsmError::Config(format!("emotion2vec forward: {e}")))?;
let logits = logits
.flatten_all()
.and_then(|t| t.to_vec1::<f32>())
.map_err(|e| crate::CsmError::Config(format!("emotion2vec logits: {e}")))?;
let argmax = logits
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(i, _)| i)
.unwrap_or(0);
Ok(Classifier::tag_for_class(argmax as u32))
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::{DType, Device};
use candle_nn::VarMap;
#[test]
fn config_stride_product_matches_320() {
let cfg = Emotion2VecConfig::plus_base();
assert_eq!(cfg.conv_stride_total(), 320, "16 kHz → 50 Hz");
assert_eq!(cfg.conv_layers.len(), 7);
}
#[test]
fn project_features_shape_check() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let pf =
ProjectFeatures::new(cfg.feature_dim, cfg.embed_dim, cfg.norm_eps, vb).expect("build");
// (B=1, T=50, in=512) → expect (1, 50, 768)
let xs = Tensor::zeros((1, 50, cfg.feature_dim), DType::F32, &dev).unwrap();
let ys = <ProjectFeatures as Module>::forward(&pf, &xs).expect("forward");
assert_eq!(ys.dims(), &[1, 50, cfg.embed_dim]);
}
#[test]
fn block_residual_shape_check() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let mlp_dim = (cfg.embed_dim as f32 * cfg.mlp_ratio) as usize;
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let block =
Block::new(cfg.embed_dim, cfg.num_heads, mlp_dim, cfg.norm_eps, vb).expect("build");
// 8-frame input of zeros — residual should pass through, output
// shape unchanged, and the LayerNorm guarantees no NaNs even when
// input is all zeros.
let xs = Tensor::randn(0f32, 1.0, (2, 8, cfg.embed_dim), &dev).unwrap();
let ys = <Block as Module>::forward(&block, &xs).expect("forward");
assert_eq!(ys.dims(), &[2, 8, cfg.embed_dim]);
// Sanity: output is finite (not NaN/Inf).
let v = ys.flatten_all().unwrap().to_vec1::<f32>().unwrap();
assert!(
v.iter().all(|x| x.is_finite()),
"block produced non-finite output"
);
}
#[test]
fn context_encoder_chains_4_blocks_with_final_norm() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let enc = ContextEncoder::new(&cfg, vb).expect("build");
let xs = Tensor::randn(0f32, 1.0, (1, 16, cfg.embed_dim), &dev).unwrap();
let ys = <ContextEncoder as Module>::forward(&enc, &xs).expect("forward");
assert_eq!(ys.dims(), &[1, 16, cfg.embed_dim]);
let v = ys.flatten_all().unwrap().to_vec1::<f32>().unwrap();
assert!(v.iter().all(|x| x.is_finite()));
}
#[test]
fn main_encoder_chains_8_blocks_no_final_norm() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let enc = MainEncoder::new(&cfg, vb).expect("build");
let xs = Tensor::randn(0f32, 1.0, (1, 16, cfg.embed_dim), &dev).unwrap();
let ys = <MainEncoder as Module>::forward(&enc, &xs).expect("forward");
assert_eq!(ys.dims(), &[1, 16, cfg.embed_dim]);
}
#[test]
fn classifier_emits_9_logits() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let cls = Classifier::new(cfg.embed_dim, cfg.num_classes, vb).expect("build");
let pooled = Tensor::zeros((2, cfg.embed_dim), DType::F32, &dev).unwrap();
let logits = <Classifier as Module>::forward(&cls, &pooled).expect("forward");
assert_eq!(logits.dims(), &[2, cfg.num_classes]);
}
#[test]
fn relative_positional_encoder_preserves_shape() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let rpe = RelativePositionalEncoder::new(cfg.embed_dim, 5, 19, 16, vb).expect("build");
let xs = Tensor::randn(0f32, 1.0, (1, 32, cfg.embed_dim), &dev).unwrap();
let ys = <RelativePositionalEncoder as Module>::forward(&rpe, &xs).expect("forward");
// Same-padding (kernel 19 with pad 9) preserves time dim.
assert_eq!(ys.dims(), &[1, 32, cfg.embed_dim]);
let v = ys.flatten_all().unwrap().to_vec1::<f32>().unwrap();
assert!(v.iter().all(|x| x.is_finite()));
}
#[test]
fn emotion2vec_random_init_end_to_end_shape() {
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let model = Emotion2Vec::new(cfg.clone(), vb, dev.clone()).expect("build");
// 1 second of synthetic audio. Random init won't give meaningful
// emotion predictions but shape + finiteness should hold.
let audio = Tensor::randn(0f32, 0.1, (1, 1, 16_000), &dev).unwrap();
let logits = model.forward(&audio).expect("forward");
assert_eq!(logits.dims(), &[1, cfg.num_classes]);
let v = logits.flatten_all().unwrap().to_vec1::<f32>().unwrap();
assert!(v.iter().all(|x| x.is_finite()), "logits had NaN/Inf");
}
#[test]
fn classifier_class_to_emotion_label_mapping() {
use crate::ser::EmotionLabel;
// Phase 13.10: 9 raw classes map directly (no fold). The earlier
// 9→5 fold was found to collapse all real audio to [excited]
// empirically; raw labels give the LoRA trainer real prosodic
// variety in its prompts.
assert_eq!(Classifier::tag_for_class(0), EmotionLabel::Angry);
assert_eq!(Classifier::tag_for_class(1), EmotionLabel::Disgusted);
assert_eq!(Classifier::tag_for_class(2), EmotionLabel::Fearful);
assert_eq!(Classifier::tag_for_class(3), EmotionLabel::Happy);
assert_eq!(Classifier::tag_for_class(4), EmotionLabel::Neutral);
assert_eq!(Classifier::tag_for_class(5), EmotionLabel::Excited); // emotion2vec "other"
assert_eq!(Classifier::tag_for_class(6), EmotionLabel::Sad);
assert_eq!(Classifier::tag_for_class(7), EmotionLabel::Surprised);
assert_eq!(Classifier::tag_for_class(8), EmotionLabel::Unk);
}
#[test]
fn local_encoder_random_init_shape_check() {
// Build with a fresh VarMap → random Kaiming init weights, then
// run a forward pass on synthetic audio and verify the output
// shape matches the expected (B, 512, T/320).
let dev = Device::Cpu;
let cfg = Emotion2VecConfig::plus_base();
let vm = VarMap::new();
let vb = VarBuilder::from_varmap(&vm, DType::F32, &dev);
let enc = LocalEncoder::new(&cfg, vb).expect("build");
// 1 second of 16 kHz = 16_000 samples; 16_000 / 320 = 50 frames.
let batch = 1;
let t_in = 16_000;
let xs = Tensor::zeros((batch, 1, t_in), DType::F32, &dev).unwrap();
let ys = <LocalEncoder as Module>::forward(&enc, &xs).expect("forward");
let dims = ys.dims();
assert_eq!(dims.len(), 3, "expected 3D output, got {dims:?}");
assert_eq!(dims[0], batch, "batch dim");
assert_eq!(dims[1], cfg.feature_dim, "feature dim should be 512");
// The first conv has kernel 10 stride 5; subsequent strides
// multiply, but kernel-padding effects shrink the time dim a few
// samples at each layer. The exact output length can be computed
// from the conv arithmetic; for 16 000 samples the upstream
// implementation produces 49 frames (one short of T/320 due to
// valid-padding kernel effects).
// Tolerance: within 2 frames of the ideal T/320=50.
let ideal = t_in / cfg.conv_stride_total();
assert!(
dims[2] <= ideal && dims[2] >= ideal.saturating_sub(2),
"time dim {} out of expected range [{}, {}]",
dims[2],
ideal.saturating_sub(2),
ideal,
);
}
}