//! Moonshine v2 candle port — encoder-decoder ASR transformer designed //! for low-latency English transcription. //! //! See `docs/moonshine_port_notes.md` for the architectural plan, the //! tensor layout from `examples/moonshine_inspect`, and the porting //! task list. //! //! ## Status //! //! - [x] Config + Conv1d audio stem (this commit, Phase 8.5) //! - [ ] Encoder transformer layer (partial RoPE, GELU MLP) //! - [ ] Decoder transformer layer (self-attn + cross-attn, SwiGLU MLP) //! - [ ] Generation loop (greedy + KV cache) //! - [ ] Tokenizer wiring (HF `tokenizer.json` via `tokenizers` crate) //! - [ ] Weight mapping + smoke test against HF reference output //! //! Each item is roughly an hour of focused work; total ~12-15 h. //! //! ## Architecture (from HF `UsefulSensors/moonshine-tiny/config.json`) //! //! - hidden_size = 288, intermediate_size = 1152 //! - 6 encoder layers + 6 decoder layers, 8 heads each (head_dim = 36, //! pad to 8 → 40 — see `pad_head_dim_to_multiple_of`) //! - vocab_size = 32_768, max_position_embeddings = 194 (decoder text) //! - encoder MLP: GELU, no gating //! - decoder MLP: SiLU SwiGLU (fc1 outputs 2× intermediate, split gate/up) //! - RoPE: theta=10000, partial_rotary_factor=0.9 (rotary on first 90 % //! of head_dim) //! - bos = 1, eos = 2, pad = 2, decoder_start = 1 use candle_core::{DType, Device, Module, Result, Tensor}; use candle_nn::{Conv1d, Conv1dConfig, LayerNorm, Linear, VarBuilder}; /// Hyperparameters matching `UsefulSensors/moonshine-tiny/config.json`. #[derive(Debug, Clone)] pub struct MoonshineConfig { pub hidden_size: usize, pub intermediate_size: usize, pub vocab_size: usize, pub max_position_embeddings: usize, pub encoder_num_hidden_layers: usize, pub encoder_num_attention_heads: usize, pub decoder_num_hidden_layers: usize, pub decoder_num_attention_heads: usize, pub rope_theta: f64, pub partial_rotary_factor: f64, pub pad_head_dim_to_multiple_of: usize, pub bos_token_id: u32, pub eos_token_id: u32, pub decoder_start_token_id: u32, } impl MoonshineConfig { /// Tiny variant (27.1 M params). pub fn tiny() -> Self { Self { hidden_size: 288, intermediate_size: 1152, vocab_size: 32_768, max_position_embeddings: 194, encoder_num_hidden_layers: 6, encoder_num_attention_heads: 8, decoder_num_hidden_layers: 6, decoder_num_attention_heads: 8, rope_theta: 10000.0, partial_rotary_factor: 0.9, pad_head_dim_to_multiple_of: 8, bos_token_id: 1, eos_token_id: 2, decoder_start_token_id: 1, } } /// Padded head dim (the actual size used for RoPE + attention math). /// `head_dim_raw = hidden_size / num_heads`, rounded up to the /// nearest multiple of `pad_head_dim_to_multiple_of`. pub fn padded_head_dim(&self, num_heads: usize) -> usize { let raw = self.hidden_size / num_heads; let pad = self.pad_head_dim_to_multiple_of; raw.div_ceil(pad) * pad } } /// Three-layer Conv1d audio stem mapping raw 16 kHz waveform to a /// sequence of 288-d hidden vectors at ~42 Hz frame rate. /// /// Layout (from HF `modeling_moonshine.py`): /// - conv1: in=1, out=288, kernel=127, stride=64, no bias /// - conv2: in=288, out=576, kernel=7, stride=3, bias /// - conv3: in=576, out=288, kernel=3, stride=2, bias /// /// Total downsampling: 64 × 3 × 2 = 384 ×. /// 16 kHz audio → output ~42 Hz frame rate (≈ 24 ms / frame). /// /// Activations between convs: HF source uses GELU after conv1 and a /// LayerNorm + GELU after conv2 (TODO: verify the exact ordering when /// the encoder transformer block is wired). pub struct ConvStem { conv1: Conv1d, conv2: Conv1d, conv3: Conv1d, } impl ConvStem { pub fn new(vb: VarBuilder) -> Result { let conv1 = candle_nn::conv1d_no_bias( 1, 288, 127, Conv1dConfig { padding: 0, stride: 64, dilation: 1, groups: 1, cudnn_fwd_algo: None, }, vb.pp("conv1"), )?; let conv2 = candle_nn::conv1d( 288, 576, 7, Conv1dConfig { padding: 0, stride: 3, dilation: 1, groups: 1, cudnn_fwd_algo: None, }, vb.pp("conv2"), )?; let conv3 = candle_nn::conv1d( 576, 288, 3, Conv1dConfig { padding: 0, stride: 2, dilation: 1, groups: 1, cudnn_fwd_algo: None, }, vb.pp("conv3"), )?; Ok(Self { conv1, conv2, conv3, }) } /// Forward: `(B, 1, T_audio)` raw waveform → `(B, T_seq, 288)` where /// `T_seq ≈ T_audio / 384`. Returns hidden states ready for the /// encoder transformer stack. pub fn forward(&self, pcm: &Tensor) -> Result { // Conv1: raw audio → 288 channels at ~250 Hz let h = self.conv1.forward(pcm)?; let h = h.tanh()?; // HF source uses tanh after conv1 // Conv2: 288 → 576, stride 3 let h = self.conv2.forward(&h)?; let h = h.gelu_erf()?; // Conv3: 576 → 288, stride 2 let h = self.conv3.forward(&h)?; let h = h.gelu_erf()?; // Output is (B, 288, T_seq); transpose to (B, T_seq, 288) for // the transformer encoder layers. h.transpose(1, 2)?.contiguous() } } /// Loader: open the HF safetensors and construct a `ConvStem`. Useful /// for the standalone Phase 8.5 smoke test. pub fn load_conv_stem(weights_path: &std::path::Path, device: &Device) -> Result { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device) }?; ConvStem::new(vb.pp("model").pp("encoder")) } // --------------------------------------------------------------------- // Phase 8.6 — encoder transformer block // --------------------------------------------------------------------- /// Cosine/sine tables for partial RoPE. /// /// Moonshine uses `partial_rotary_factor = 0.9` with `head_dim = 36` /// (288 / 8 heads). `rotary_dim = floor(36 * 0.9 / 2) * 2 = 32` (must /// be even — RoPE rotates pairs). The remaining 4 dims pass through /// unchanged. RoPE base is `theta = 10_000`. struct RotaryCache { cos: Tensor, // (max_seq, rotary_dim/2) sin: Tensor, // (max_seq, rotary_dim/2) rotary_dim: usize, } impl RotaryCache { fn new( rotary_dim: usize, max_seq: usize, theta: f64, dtype: DType, dev: &Device, ) -> Result { assert!(rotary_dim.is_multiple_of(2), "rotary_dim must be even"); let inv_freq: Vec = (0..rotary_dim) .step_by(2) .map(|i| (1.0 / theta.powf(i as f64 / rotary_dim as f64)) as f32) .collect(); let inv_freq = Tensor::new(inv_freq, dev)?; let positions = Tensor::arange(0u32, max_seq as u32, dev)? .to_dtype(DType::F32)? .reshape((max_seq, 1))?; let freqs = positions.matmul(&inv_freq.reshape((1, rotary_dim / 2))?)?; let cos = freqs.cos()?.to_dtype(dtype)?; let sin = freqs.sin()?.to_dtype(dtype)?; Ok(Self { cos, sin, rotary_dim, }) } /// Apply partial RoPE to `q` of shape `(B, H, T, head_dim)`. /// Rotates the first `rotary_dim` channels; leaves the rest as-is. /// Uses positions `0..T` (prefill mode). fn apply(&self, x: &Tensor) -> Result { let (_b, _h, t, _) = x.dims4()?; self.apply_at(x, 0, t) } /// Apply RoPE for a window starting at `position`, length `t`. /// Used by the cached single-token step (`t=1`, `position=cache_len`). fn apply_at(&self, x: &Tensor, position: usize, t: usize) -> Result { let (_b, _h, _t, head_dim) = x.dims4()?; let cos = self.cos.narrow(0, position, t)?; let sin = self.sin.narrow(0, position, t)?; if head_dim == self.rotary_dim { return candle_nn::rotary_emb::rope_i(x, &cos, &sin); } let rot = x.narrow(3, 0, self.rotary_dim)?.contiguous()?; let pass = x.narrow(3, self.rotary_dim, head_dim - self.rotary_dim)?; let rot = candle_nn::rotary_emb::rope_i(&rot, &cos, &sin)?; Tensor::cat(&[&rot, &pass], 3) } } /// Self-attention block matching `model.encoder.layers.X.self_attn`. /// 8 heads × 36 head_dim, no bias (`attention_bias=false`), partial /// RoPE on first 32 head_dim channels. struct EncoderAttention { q_proj: Linear, k_proj: Linear, v_proj: Linear, o_proj: Linear, n_heads: usize, head_dim: usize, } impl EncoderAttention { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let h = cfg.hidden_size; let n_heads = cfg.encoder_num_attention_heads; let head_dim = h / n_heads; let q_proj = candle_nn::linear_no_bias(h, h, vb.pp("q_proj"))?; let k_proj = candle_nn::linear_no_bias(h, h, vb.pp("k_proj"))?; let v_proj = candle_nn::linear_no_bias(h, h, vb.pp("v_proj"))?; let o_proj = candle_nn::linear_no_bias(h, h, vb.pp("o_proj"))?; Ok(Self { q_proj, k_proj, v_proj, o_proj, n_heads, head_dim, }) } fn forward(&self, xs: &Tensor, rope: &RotaryCache) -> Result { let (b, t, _h) = xs.dims3()?; let q = self.q_proj.forward(xs)?; let k = self.k_proj.forward(xs)?; let v = self.v_proj.forward(xs)?; // (B, T, H) -> (B, H_heads, T, head_dim) let q = q .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let k = k .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let v = v .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let q = rope.apply(&q)?; let k = rope.apply(&k)?; // Collapse (B, H, T, D) -> (B*H, T, D) for the matmul. candle's // 4D batched matmul on Metal had shape-mismatch issues for our // pattern; the 3D form is unambiguous. let bh = b * self.n_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 scale = 1.0 / (self.head_dim as f64).sqrt(); let scores = (q3.matmul(&k3.transpose(1, 2)?.contiguous()?)? * scale)?; let probs = candle_nn::ops::softmax_last_dim(&scores)?; let out = probs.matmul(&v3)?; // (B*H, T, D) -> (B, H, T, D) -> (B, T, H*D) let out = out .reshape((b, self.n_heads, t, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((b, t, self.n_heads * self.head_dim))?; self.o_proj.forward(&out) } } /// FFN block matching `model.encoder.layers.X.mlp`. /// 288 → 1152 (with bias) → GELU → 1152 → 288 (with bias). struct EncoderMlp { fc1: Linear, fc2: Linear, } impl EncoderMlp { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let fc1 = candle_nn::linear(cfg.hidden_size, cfg.intermediate_size, vb.pp("fc1"))?; let fc2 = candle_nn::linear(cfg.intermediate_size, cfg.hidden_size, vb.pp("fc2"))?; Ok(Self { fc1, fc2 }) } fn forward(&self, xs: &Tensor) -> Result { let h = self.fc1.forward(xs)?; let h = h.gelu_erf()?; self.fc2.forward(&h) } } /// One encoder layer: Pre-LN attn + Pre-LN MLP, residuals as standard. /// Tensors are `LayerNorm` weight-only (no bias) per the inspector /// dump — implemented via `LayerNorm` with bias zeros. struct EncoderLayer { input_ln: LayerNorm, self_attn: EncoderAttention, post_attn_ln: LayerNorm, mlp: EncoderMlp, } impl EncoderLayer { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let h = cfg.hidden_size; let input_ln = layer_norm_weight_only(h, 1e-5, vb.pp("input_layernorm"))?; let self_attn = EncoderAttention::new(cfg, vb.pp("self_attn"))?; let post_attn_ln = layer_norm_weight_only(h, 1e-5, vb.pp("post_attention_layernorm"))?; let mlp = EncoderMlp::new(cfg, vb.pp("mlp"))?; Ok(Self { input_ln, self_attn, post_attn_ln, mlp, }) } fn forward(&self, xs: &Tensor, rope: &RotaryCache) -> Result { let h = self.input_ln.forward(xs)?; let h = self.self_attn.forward(&h, rope)?; let xs = (xs + h)?; let h = self.post_attn_ln.forward(&xs)?; let h = self.mlp.forward(&h)?; xs + h } } /// Helper: `LayerNorm` with weight only (zeros bias). Moonshine's /// `*_layernorm.weight [288]` tensors lack a corresponding `.bias`. fn layer_norm_weight_only(size: usize, eps: f64, vb: VarBuilder) -> Result { let weight = vb.get((size,), "weight")?; let bias = Tensor::zeros((size,), weight.dtype(), weight.device())?; Ok(LayerNorm::new(weight, bias, eps)) } /// Full encoder: conv stem → 6 transformer layers → final layer norm. /// Output shape: `(B, T_seq, hidden=288)` ready to feed cross-attn in /// the decoder (Phase 8.7). pub struct Encoder { stem: ConvStem, layers: Vec, final_ln: LayerNorm, rope: RotaryCache, } impl Encoder { pub fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let stem = ConvStem::new(vb.pp("encoder"))?; let head_dim = cfg.hidden_size / cfg.encoder_num_attention_heads; let rotary_dim = ((head_dim as f64 * cfg.partial_rotary_factor) as usize / 2) * 2; // max_seq for the encoder is bounded by the input audio length / // total stride 384. For 30 s of 16 kHz audio that's ~1250. // Allocate a generous table. let rope = RotaryCache::new(rotary_dim, 4096, cfg.rope_theta, vb.dtype(), vb.device())?; let mut layers = Vec::with_capacity(cfg.encoder_num_hidden_layers); let layer_vb = vb.pp("encoder").pp("layers"); for i in 0..cfg.encoder_num_hidden_layers { layers.push(EncoderLayer::new(cfg, layer_vb.pp(i))?); } let final_ln = layer_norm_weight_only(cfg.hidden_size, 1e-5, vb.pp("encoder").pp("layer_norm"))?; Ok(Self { stem, layers, final_ln, rope, }) } pub fn forward(&self, pcm: &Tensor) -> Result { let mut h = self.stem.forward(pcm)?; for layer in &self.layers { h = layer.forward(&h, &self.rope)?; } self.final_ln.forward(&h) } } /// Loader: open the HF safetensors and construct an `Encoder` with /// weights at `model.*` (the canonical HF prefix). pub fn load_encoder( weights_path: &std::path::Path, device: &Device, cfg: &MoonshineConfig, ) -> Result { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device) }?; Encoder::new(cfg, vb.pp("model")) } // --------------------------------------------------------------------- // Phase 8.7 — decoder transformer block // --------------------------------------------------------------------- /// Cross-attention block matching `model.decoder.layers.X.encoder_attn`. /// Queries come from the decoder hidden state; keys and values come /// from the encoder output (computed once per generation step, cached /// across generation tokens — caching deferred to Phase 8.8). struct CrossAttention { q_proj: Linear, k_proj: Linear, v_proj: Linear, o_proj: Linear, n_heads: usize, head_dim: usize, } impl CrossAttention { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let h = cfg.hidden_size; let n_heads = cfg.decoder_num_attention_heads; let head_dim = h / n_heads; Ok(Self { q_proj: candle_nn::linear_no_bias(h, h, vb.pp("q_proj"))?, k_proj: candle_nn::linear_no_bias(h, h, vb.pp("k_proj"))?, v_proj: candle_nn::linear_no_bias(h, h, vb.pp("v_proj"))?, o_proj: candle_nn::linear_no_bias(h, h, vb.pp("o_proj"))?, n_heads, head_dim, }) } /// `xs`: decoder hidden `(B, T_dec, H)`. /// `enc`: encoder output `(B, T_enc, H)`. /// Returns: `(B, T_dec, H)`. fn forward(&self, xs: &Tensor, enc: &Tensor) -> Result { let (b, t_dec, _h) = xs.dims3()?; let (_, t_enc, _) = enc.dims3()?; let q = self.q_proj.forward(xs)?; let k = self.k_proj.forward(enc)?; let v = self.v_proj.forward(enc)?; // (B, T, H) -> (B*H_heads, T, head_dim) — same 3D matmul // pattern as EncoderAttention to dodge the candle 4D Metal bug. let bh = b * self.n_heads; let q = q .reshape((b, t_dec, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((bh, t_dec, self.head_dim))?; let k = k .reshape((b, t_enc, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((bh, t_enc, self.head_dim))?; let v = v .reshape((b, t_enc, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((bh, t_enc, self.head_dim))?; let scale = 1.0 / (self.head_dim as f64).sqrt(); // Cross-attention has NO causal mask: each decoder token can // attend to every encoder position. let scores = (q.matmul(&k.transpose(1, 2)?.contiguous()?)? * scale)?; let probs = candle_nn::ops::softmax_last_dim(&scores)?; let out = probs.matmul(&v)?; let out = out .reshape((b, self.n_heads, t_dec, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((b, t_dec, self.n_heads * self.head_dim))?; self.o_proj.forward(&out) } } /// Decoder self-attention. Same shape as encoder attention but applies /// a causal mask so token `i` only attends to tokens `0..=i`. Partial /// RoPE on q/k as in the encoder. struct DecoderSelfAttention { q_proj: Linear, k_proj: Linear, v_proj: Linear, o_proj: Linear, n_heads: usize, head_dim: usize, } impl DecoderSelfAttention { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let h = cfg.hidden_size; let n_heads = cfg.decoder_num_attention_heads; let head_dim = h / n_heads; Ok(Self { q_proj: candle_nn::linear_no_bias(h, h, vb.pp("q_proj"))?, k_proj: candle_nn::linear_no_bias(h, h, vb.pp("k_proj"))?, v_proj: candle_nn::linear_no_bias(h, h, vb.pp("v_proj"))?, o_proj: candle_nn::linear_no_bias(h, h, vb.pp("o_proj"))?, n_heads, head_dim, }) } fn forward(&self, xs: &Tensor, rope: &RotaryCache) -> Result { let (b, t, _h) = xs.dims3()?; let q = self.q_proj.forward(xs)?; let k = self.k_proj.forward(xs)?; let v = self.v_proj.forward(xs)?; let q = q .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let k = k .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let v = v .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let q = rope.apply(&q)?; let k = rope.apply(&k)?; let bh = b * self.n_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 scale = 1.0 / (self.head_dim as f64).sqrt(); let mut scores = (q3.matmul(&k3.transpose(1, 2)?.contiguous()?)? * scale)?; // Causal mask: build a (T, T) upper-triangular -inf mask. let mask = causal_mask(t, scores.dtype(), scores.device())?; scores = scores.broadcast_add(&mask)?; let probs = candle_nn::ops::softmax_last_dim(&scores)?; let out = probs.matmul(&v3)?; let out = out .reshape((b, self.n_heads, t, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((b, t, self.n_heads * self.head_dim))?; self.o_proj.forward(&out) } } fn causal_mask(t: usize, dtype: DType, device: &Device) -> Result { // Upper-triangular `-inf` mask: value at (i, j) is 0 if j <= i, else -inf. let mut data = vec![0.0f32; t * t]; for i in 0..t { for j in (i + 1)..t { data[i * t + j] = f32::NEG_INFINITY; } } Tensor::from_vec(data, (t, t), device)?.to_dtype(dtype) } /// SwiGLU MLP for the decoder. `fc1` is the fused gate+up projection: /// weight shape `[2 * intermediate_size, hidden_size] = [2304, 288]`. /// Output is split on dim -1 into `gate` and `up`, both `[..., 1152]`. /// Activation: `silu(gate) * up`. Then `fc2: [hidden, intermediate] = /// [288, 1152]` projects back to `hidden`. struct DecoderMlp { fc1: Linear, fc2: Linear, intermediate: usize, } impl DecoderMlp { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { // fc1 has bias; fc2 has bias (consistent with encoder MLP — the // inspector dump shows `.bias` on both). let fc1 = candle_nn::linear(cfg.hidden_size, cfg.intermediate_size * 2, vb.pp("fc1"))?; let fc2 = candle_nn::linear(cfg.intermediate_size, cfg.hidden_size, vb.pp("fc2"))?; Ok(Self { fc1, fc2, intermediate: cfg.intermediate_size, }) } fn forward(&self, xs: &Tensor) -> Result { // (B, T, H) -> (B, T, 2 * intermediate) let h = self.fc1.forward(xs)?; let dims = h.dims(); let last = dims.len() - 1; // HF Moonshine MLP (verified against modeling_moonshine.py): // hidden, gate = fc1(x).chunk(2, dim=-1) // out = silu(gate) * hidden // i.e., FIRST half is `up` (hidden), SECOND half is `gate`. // We had this reversed initially, which produced a degenerate // repetition loop after 1-2 tokens. let up = h.narrow(last, 0, self.intermediate)?; let gate = h.narrow(last, self.intermediate, self.intermediate)?; let activated = candle_nn::ops::silu(&gate)?.mul(&up)?; self.fc2.forward(&activated) } } /// One decoder layer: Pre-LN self-attn, Pre-LN cross-attn, Pre-LN MLP. /// Three layer norms per layer (input_layernorm, post_attention_layernorm, /// final_layernorm), all weight-only. struct DecoderLayer { input_ln: LayerNorm, self_attn: DecoderSelfAttention, post_attn_ln: LayerNorm, cross_attn: CrossAttention, final_ln: LayerNorm, mlp: DecoderMlp, } impl DecoderLayer { fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let h = cfg.hidden_size; Ok(Self { input_ln: layer_norm_weight_only(h, 1e-5, vb.pp("input_layernorm"))?, self_attn: DecoderSelfAttention::new(cfg, vb.pp("self_attn"))?, post_attn_ln: layer_norm_weight_only(h, 1e-5, vb.pp("post_attention_layernorm"))?, cross_attn: CrossAttention::new(cfg, vb.pp("encoder_attn"))?, final_ln: layer_norm_weight_only(h, 1e-5, vb.pp("final_layernorm"))?, mlp: DecoderMlp::new(cfg, vb.pp("mlp"))?, }) } fn forward(&self, xs: &Tensor, enc: &Tensor, rope: &RotaryCache) -> Result { // Self-attn (causal) let h = self.input_ln.forward(xs)?; let h = self.self_attn.forward(&h, rope)?; let xs = (xs + h)?; // Cross-attn let h = self.post_attn_ln.forward(&xs)?; let h = self.cross_attn.forward(&h, enc)?; let xs = (xs + h)?; // MLP let h = self.final_ln.forward(&xs)?; let h = self.mlp.forward(&h)?; xs + h } } /// Full decoder: token embedding → 6 layers → final LN → tied LM head. pub struct Decoder { embed: candle_nn::Embedding, layers: Vec, final_ln: LayerNorm, rope: RotaryCache, /// Cached transposed embedding for the tied LM head. /// Computing logits: hidden @ embed.weight.T -> (B, T, vocab). embed_weight: Tensor, } impl Decoder { pub fn new(cfg: &MoonshineConfig, vb: VarBuilder) -> Result { let embed = candle_nn::embedding( cfg.vocab_size, cfg.hidden_size, vb.pp("decoder").pp("embed_tokens"), )?; let head_dim = cfg.hidden_size / cfg.decoder_num_attention_heads; let rotary_dim = ((head_dim as f64 * cfg.partial_rotary_factor) as usize / 2) * 2; let rope = RotaryCache::new( rotary_dim, cfg.max_position_embeddings.max(512), cfg.rope_theta, vb.dtype(), vb.device(), )?; let mut layers = Vec::with_capacity(cfg.decoder_num_hidden_layers); let layer_vb = vb.pp("decoder").pp("layers"); for i in 0..cfg.decoder_num_hidden_layers { layers.push(DecoderLayer::new(cfg, layer_vb.pp(i))?); } // Decoder uses `decoder.norm.weight` (the encoder uses // `encoder.layer_norm.weight`). Inspector dump confirmed. let final_ln = layer_norm_weight_only(cfg.hidden_size, 1e-5, vb.pp("decoder").pp("norm"))?; let embed_weight = embed.embeddings().clone(); Ok(Self { embed, layers, final_ln, rope, embed_weight, }) } /// Forward over `tokens` shape `(B, T)` with encoder output `enc` /// shape `(B, T_enc, H)`. Returns logits `(B, T, vocab_size)`. /// **No KV cache** — this is the prefill / single-step path. /// The full streaming-generation loop with KV cache lands in /// Phase 8.8. pub fn forward(&self, tokens: &Tensor, enc: &Tensor) -> Result { let mut h = self.embed.forward(tokens)?; for layer in &self.layers { h = layer.forward(&h, enc, &self.rope)?; } h = self.final_ln.forward(&h)?; // LM head tied to embedding: logits = hidden @ embed.T. // embed.weight shape: (vocab, hidden). transpose -> (hidden, vocab). let lm_w = self.embed_weight.transpose(0, 1)?.contiguous()?; h.broadcast_matmul(&lm_w) } } /// Loader: open the HF safetensors and construct an encoder + decoder. pub fn load_full( weights_path: &std::path::Path, device: &Device, cfg: &MoonshineConfig, ) -> Result<(Encoder, Decoder)> { let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, device) }?; let encoder = Encoder::new(cfg, vb.pp("model"))?; let decoder = Decoder::new(cfg, vb.pp("model"))?; Ok((encoder, decoder)) } // --------------------------------------------------------------------- // Phase 8.8 — generation loop + tokenizer // --------------------------------------------------------------------- impl Decoder { /// Greedy autoregressive decode given an encoder output. Stops on /// `eos_token_id` or after `max_tokens`. Returns the predicted /// token ids (excluding the initial start token). /// /// **No KV cache** — each step re-runs the decoder on the full /// growing token sequence (O(T^2) total work). For Moonshine-tiny /// at ~85 ms / step prefill, a 30-token transcript takes a few /// seconds. KV-cached `step()` is a follow-up if needed. pub fn generate( &self, encoder_output: &Tensor, cfg: &MoonshineConfig, max_tokens: usize, ) -> Result> { let device = encoder_output.device(); let mut tokens: Vec = vec![cfg.decoder_start_token_id]; let mut out = Vec::with_capacity(max_tokens); for _ in 0..max_tokens { let input = Tensor::from_vec(tokens.clone(), (1, tokens.len()), device)?; let logits = self.forward(&input, encoder_output)?; // logits: (1, T_so_far, vocab). Take last position. let last_t = tokens.len() - 1; let last = logits.narrow(1, last_t, 1)?.squeeze(1)?; // (1, vocab) let argmax = last.argmax(1)?; let next_id: u32 = argmax.to_dtype(DType::U32)?.to_vec1::()?[0]; if next_id == cfg.eos_token_id { break; } out.push(next_id); tokens.push(next_id); if tokens.len() >= cfg.max_position_embeddings { break; } } Ok(out) } } /// Convenience: load the HF tokenizer.json for Moonshine. Caller passes /// the path returned by hf_hub. pub fn load_tokenizer( path: &std::path::Path, ) -> std::result::Result> { tokenizers::Tokenizer::from_file(path) } // --------------------------------------------------------------------- // Phase 8.9 — KV cache for single-token autoregressive decoding // --------------------------------------------------------------------- /// Per-layer K/V cache for one decode session. /// - `self_k[i]` / `self_v[i]`: cumulative self-attn K/V for layer i, /// shape `(B*H, position, head_dim)`. Grows by 1 per `step()`. /// - `cross_k[i]` / `cross_v[i]`: encoder cross-attn K/V for layer i, /// computed once via `precompute_cross_kv` and reused every step. pub struct DecoderCache { self_k: Vec>, self_v: Vec>, cross_k: Vec, cross_v: Vec, pub position: usize, } impl DecoderSelfAttention { /// Single-token cached self-attention. Appends new K/V to the /// cache and attends across the full accumulated history. /// `xs`: `(B=1, T=1, H)`. Returns `(B=1, T=1, H)`. fn forward_step( &self, xs: &Tensor, cache_k: &mut Option, cache_v: &mut Option, rope: &RotaryCache, position: usize, ) -> Result { let (b, t, _h) = xs.dims3()?; debug_assert_eq!(t, 1, "self-attn step expects T=1"); let q = self.q_proj.forward(xs)?; let k = self.k_proj.forward(xs)?; let v = self.v_proj.forward(xs)?; // (B, T, H) -> (B, H, T, head_dim) let q = q .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let k = k .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; let v = v .reshape((b, t, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()?; // RoPE for q and the new k at this position. let q = rope.apply_at(&q, position, 1)?; let k_new = rope.apply_at(&k, position, 1)?; // Collapse to (B*H, T, head_dim) for the matmul. let bh = b * self.n_heads; let q3 = q.reshape((bh, 1, self.head_dim))?; let k_new3 = k_new.reshape((bh, 1, self.head_dim))?; let v3 = v.reshape((bh, 1, self.head_dim))?; // Append to cache (or initialize on first step). let k_full = match cache_k.take() { Some(prev) => Tensor::cat(&[&prev, &k_new3], 1)?, None => k_new3, }; let v_full = match cache_v.take() { Some(prev) => Tensor::cat(&[&prev, &v3], 1)?, None => v3, }; // No causal mask needed: K/V only contains positions <= current. let scale = 1.0 / (self.head_dim as f64).sqrt(); let scores = (q3.matmul(&k_full.transpose(1, 2)?.contiguous()?)? * scale)?; let probs = candle_nn::ops::softmax_last_dim(&scores)?; let out = probs.matmul(&v_full)?; // Write cache back. *cache_k = Some(k_full); *cache_v = Some(v_full); // (B*H, 1, head_dim) -> (B, 1, H) let out = out .reshape((b, self.n_heads, 1, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((b, 1, self.n_heads * self.head_dim))?; self.o_proj.forward(&out) } } impl CrossAttention { /// Pre-compute and cache encoder K/V projections. Returns /// `(K, V)` each shaped `(B*H, T_enc, head_dim)`. fn precompute_kv(&self, enc: &Tensor) -> Result<(Tensor, Tensor)> { let (b, t_enc, _) = enc.dims3()?; let k = self.k_proj.forward(enc)?; let v = self.v_proj.forward(enc)?; let bh = b * self.n_heads; let k = k .reshape((b, t_enc, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((bh, t_enc, self.head_dim))?; let v = v .reshape((b, t_enc, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((bh, t_enc, self.head_dim))?; Ok((k, v)) } /// Cached cross-attention step. K/V come from `precompute_kv`. /// `xs`: `(B=1, T=1, H)`. Returns `(B=1, T=1, H)`. fn forward_step(&self, xs: &Tensor, k: &Tensor, v: &Tensor) -> Result { let (b, t_dec, _) = xs.dims3()?; debug_assert_eq!(t_dec, 1, "cross-attn step expects T=1"); let q = self.q_proj.forward(xs)?; let bh = b * self.n_heads; let q = q .reshape((b, t_dec, self.n_heads, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((bh, 1, self.head_dim))?; let scale = 1.0 / (self.head_dim as f64).sqrt(); let scores = (q.matmul(&k.transpose(1, 2)?.contiguous()?)? * scale)?; let probs = candle_nn::ops::softmax_last_dim(&scores)?; let out = probs.matmul(v)?; let out = out .reshape((b, self.n_heads, 1, self.head_dim))? .transpose(1, 2)? .contiguous()? .reshape((b, 1, self.n_heads * self.head_dim))?; self.o_proj.forward(&out) } } impl Decoder { /// Pre-compute the encoder cross-attn K/V for every layer. Returns /// an empty `DecoderCache` ready for `step()` calls. pub fn precompute_cross_kv(&self, encoder_output: &Tensor) -> Result { let n = self.layers.len(); let mut cross_k = Vec::with_capacity(n); let mut cross_v = Vec::with_capacity(n); for layer in &self.layers { let (k, v) = layer.cross_attn.precompute_kv(encoder_output)?; cross_k.push(k); cross_v.push(v); } Ok(DecoderCache { self_k: (0..n).map(|_| None).collect(), self_v: (0..n).map(|_| None).collect(), cross_k, cross_v, position: 0, }) } /// Single-token cached forward. `token_id` is the *new* token to /// process. Returns logits for that one position, shape /// `(1, vocab)`. Updates `cache.position` and per-layer `self_k`/ /// `self_v` in place. pub fn step(&self, token_id: u32, cache: &mut DecoderCache) -> Result { let device = self.embed_weight.device(); let token = Tensor::from_vec(vec![token_id], (1, 1), device)?; let mut h = self.embed.forward(&token)?; let position = cache.position; for (i, layer) in self.layers.iter().enumerate() { let normed = layer.input_ln.forward(&h)?; let attn_out = layer.self_attn.forward_step( &normed, &mut cache.self_k[i], &mut cache.self_v[i], &self.rope, position, )?; h = (h + attn_out)?; let normed = layer.post_attn_ln.forward(&h)?; let cross_out = layer .cross_attn .forward_step(&normed, &cache.cross_k[i], &cache.cross_v[i])?; h = (h + cross_out)?; let normed = layer.final_ln.forward(&h)?; let mlp_out = layer.mlp.forward(&normed)?; h = (h + mlp_out)?; } h = self.final_ln.forward(&h)?; cache.position += 1; // Logits: (1, 1, H) -> (1, H) -> (1, vocab) via tied LM head. let lm_w = self.embed_weight.transpose(0, 1)?.contiguous()?; let logits_3d = h.broadcast_matmul(&lm_w)?; logits_3d.squeeze(1) } /// Cached greedy decode. Equivalent semantics to /// [`Self::generate`] but with O(T) total work instead of O(T²). pub fn generate_cached( &self, encoder_output: &Tensor, cfg: &MoonshineConfig, max_tokens: usize, ) -> Result> { let mut cache = self.precompute_cross_kv(encoder_output)?; let mut next = cfg.decoder_start_token_id; let mut out = Vec::with_capacity(max_tokens); for _ in 0..max_tokens { let logits = self.step(next, &mut cache)?; // (1, vocab) let argmax = logits.argmax(1)?; let id: u32 = argmax.to_dtype(DType::U32)?.to_vec1::()?[0]; if id == cfg.eos_token_id { break; } out.push(id); next = id; if cache.position >= cfg.max_position_embeddings { break; } } Ok(out) } } #[cfg(test)] mod tests { use super::*; #[test] fn config_padded_head_dim_pads_to_multiple_of_8() { let cfg = MoonshineConfig::tiny(); // 288 / 8 = 36, padded up to 40 (next multiple of 8). assert_eq!(cfg.padded_head_dim(cfg.encoder_num_attention_heads), 40); // Sanity: smaller-head config still pads. let mut alt = cfg.clone(); alt.hidden_size = 256; assert_eq!(alt.padded_head_dim(8), 32); } }