feat(batch19): Medusa heads, TIES+DARE model merging, Mixture of Depths
CI / Format Check (push) Failing after 13s
CI / Build (macos-latest) (push) Failing after 29s
CI / Clippy Check (push) Failing after 43s
Documentation / Build User Guide (push) Successful in 6s
GPU Tests / Check GPU Availability (push) Successful in 0s
GPU Tests / CUDA Tests (11.8) (push) Has been skipped
GPU Tests / CUDA Tests (12.1) (push) Has been skipped
CI / Build (ubuntu-latest) (push) Failing after 1m12s
CI / Test (macos-latest) (push) Has been skipped
CI / Test (ubuntu-latest) (push) Has been skipped
CI / Python Bindings (maturin) (macos-latest) (push) Has been skipped
CI / Python Bindings (maturin) (ubuntu-latest) (push) Has been skipped
CI / WASM Build + Size Check (push) Has been skipped
CI / Distributed Training Tests (push) Has been skipped
CI / Build CPU-Only (Explicit) (push) Failing after 1m40s
Documentation / Build API Documentation (push) Failing after 1m26s
CI / CI Success (push) Failing after 0s
Performance Benchmarks / Run Benchmarks (push) Successful in 2m10s
GPU Tests / Metal Tests (push) Has been skipped
CI / Format Check (push) Failing after 13s
CI / Build (macos-latest) (push) Failing after 29s
CI / Clippy Check (push) Failing after 43s
Documentation / Build User Guide (push) Successful in 6s
GPU Tests / Check GPU Availability (push) Successful in 0s
GPU Tests / CUDA Tests (11.8) (push) Has been skipped
GPU Tests / CUDA Tests (12.1) (push) Has been skipped
CI / Build (ubuntu-latest) (push) Failing after 1m12s
CI / Test (macos-latest) (push) Has been skipped
CI / Test (ubuntu-latest) (push) Has been skipped
CI / Python Bindings (maturin) (macos-latest) (push) Has been skipped
CI / Python Bindings (maturin) (ubuntu-latest) (push) Has been skipped
CI / WASM Build + Size Check (push) Has been skipped
CI / Distributed Training Tests (push) Has been skipped
CI / Build CPU-Only (Explicit) (push) Failing after 1m40s
Documentation / Build API Documentation (push) Failing after 1m26s
CI / CI Success (push) Failing after 0s
Performance Benchmarks / Run Benchmarks (push) Successful in 2m10s
GPU Tests / Metal Tests (push) Has been skipped
- MedusaHeads: K FFN draft heads (SiLU 2-layer); tree candidate generation via cartesian product of per-head top-k; path verification with oracle; CE training loss per head (arXiv:2401.10774); 23 tests - ModelMerger: TIES (task-vector trim+elect-sign+disjoint-merge, arXiv:2306.01708) + DARE sparse rescaling (arXiv:2311.03099); linear merge baseline; 29 tests - MoDLayer/MoDStack: per-token capacity routing (top-k by router score); residual bypass for skipped tokens; load-balancing aux loss; flops_reduction = product of capacity_fractions (arXiv:2404.02258); 22 tests Co-Authored-By: Claude Sonnet 4.6 <[email protected]>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
960fd73c82
commit
bff27c302f
@@ -0,0 +1,775 @@
|
||||
//! Medusa speculative decoding heads (Cai et al. 2024, arXiv:2401.10774).
|
||||
//!
|
||||
//! Medusa accelerates autoregressive decoding by attaching K small FFN "draft heads"
|
||||
//! to the last hidden state of the base model. Each head independently predicts
|
||||
//! tokens at offsets +1, +2, …, +K from the current position. At inference time the
|
||||
//! heads collectively generate a *tree* of candidate continuations, which the base
|
||||
//! model verifies in a single forward pass using tree attention.
|
||||
//!
|
||||
//! # Architecture
|
||||
//!
|
||||
//! Each [`MedusaHead`] is a two-layer FFN:
|
||||
//!
|
||||
//! ```text
|
||||
//! hidden [H] → SiLU(W1 @ hidden + b1) [H] → W2 @ act + b2 [V]
|
||||
//! ```
|
||||
//!
|
||||
//! The [`MedusaHeads`] ensemble holds K such heads and exposes:
|
||||
//!
|
||||
//! * [`MedusaHeads::generate_tree`] — cartesian-product candidate tree
|
||||
//! * [`MedusaHeads::verify`] — longest-accepted-prefix selection
|
||||
//! * [`MedusaHeads::training_loss`] — per-head cross-entropy
|
||||
//!
|
||||
//! # Example
|
||||
//!
|
||||
//! ```rust
|
||||
//! use rtx_inference::medusa::{MedusaHeads, MedusaConfig};
|
||||
//!
|
||||
//! let config = MedusaConfig::default();
|
||||
//! let heads = MedusaHeads::new(config);
|
||||
//! let hidden = vec![0.0_f32; 64];
|
||||
//! let tree = heads.generate_tree(&hidden);
|
||||
//! assert_eq!(tree.num_paths(), 3_usize.pow(4)); // top_k^num_heads
|
||||
//! ```
|
||||
|
||||
// ── LCG pseudo-random number generator ────────────────────────────────────────
|
||||
|
||||
/// Minimal LCG PRNG used for reproducible weight initialisation (no external deps).
|
||||
struct Lcg(u64);
|
||||
|
||||
impl Lcg {
|
||||
fn new(seed: u64) -> Self {
|
||||
Self(seed ^ 0x1234_5678_9abc_def0)
|
||||
}
|
||||
|
||||
/// Advance the state and return a value in `[0, 1)`.
|
||||
fn next_f32(&mut self) -> f32 {
|
||||
// Knuth's multiplier + addend, 64-bit
|
||||
self.0 = self
|
||||
.0
|
||||
.wrapping_mul(6_364_136_223_846_793_005)
|
||||
.wrapping_add(1_442_695_040_888_963_407);
|
||||
// Use the high 23 bits for the mantissa
|
||||
let bits = 0x3f80_0000_u32 | ((self.0 >> 41) as u32 & 0x007f_ffff);
|
||||
f32::from_bits(bits) - 1.0
|
||||
}
|
||||
|
||||
/// Return a value uniformly in `[-scale, scale]`.
|
||||
fn next_scaled(&mut self, scale: f32) -> f32 {
|
||||
(self.next_f32() * 2.0 - 1.0) * scale
|
||||
}
|
||||
}
|
||||
|
||||
// ── MedusaHead ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// One Medusa draft head: a 2-layer FFN applied to the last hidden state.
|
||||
///
|
||||
/// Computes `hidden [H] → SiLU(W1 @ hidden + b1) [H] → W2 @ act + b2 [V]`.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MedusaHead {
|
||||
/// W1 weight matrix, row-major `[hidden_dim × hidden_dim]`.
|
||||
pub w1: Vec<f32>,
|
||||
/// Bias for layer 1, `[hidden_dim]`.
|
||||
pub b1: Vec<f32>,
|
||||
/// W2 weight matrix, row-major `[vocab_size × hidden_dim]`.
|
||||
pub w2: Vec<f32>,
|
||||
/// Bias for layer 2, `[vocab_size]`.
|
||||
pub b2: Vec<f32>,
|
||||
/// Hidden dimension `H`.
|
||||
pub hidden_dim: usize,
|
||||
/// Vocabulary size `V`.
|
||||
pub vocab_size: usize,
|
||||
}
|
||||
|
||||
impl MedusaHead {
|
||||
/// Initialise with LCG-seeded random weights.
|
||||
///
|
||||
/// Weights are drawn from `U[-scale, scale]` where
|
||||
/// `scale = sqrt(2 / fan_in)` (Xavier-like). Biases are zeroed.
|
||||
pub fn new_random(hidden_dim: usize, vocab_size: usize, seed: u64) -> Self {
|
||||
let mut rng = Lcg::new(seed);
|
||||
|
||||
let scale_w1 = (2.0_f32 / hidden_dim as f32).sqrt();
|
||||
let w1: Vec<f32> = (0..hidden_dim * hidden_dim)
|
||||
.map(|_| rng.next_scaled(scale_w1))
|
||||
.collect();
|
||||
let b1 = vec![0.0_f32; hidden_dim];
|
||||
|
||||
let scale_w2 = (2.0_f32 / hidden_dim as f32).sqrt();
|
||||
let w2: Vec<f32> = (0..vocab_size * hidden_dim)
|
||||
.map(|_| rng.next_scaled(scale_w2))
|
||||
.collect();
|
||||
let b2 = vec![0.0_f32; vocab_size];
|
||||
|
||||
Self { w1, b1, w2, b2, hidden_dim, vocab_size }
|
||||
}
|
||||
|
||||
/// Run the head forward pass.
|
||||
///
|
||||
/// `hidden` must have length `hidden_dim`. Returns logits of length
|
||||
/// `vocab_size`.
|
||||
///
|
||||
/// # Panics
|
||||
///
|
||||
/// Panics in debug builds if `hidden.len() != self.hidden_dim`.
|
||||
pub fn forward(&self, hidden: &[f32]) -> Vec<f32> {
|
||||
debug_assert_eq!(
|
||||
hidden.len(),
|
||||
self.hidden_dim,
|
||||
"hidden vector length must equal hidden_dim"
|
||||
);
|
||||
let h = self.hidden_dim;
|
||||
|
||||
// Layer 1: act = SiLU(W1 @ hidden + b1)
|
||||
let mut act = vec![0.0_f32; h];
|
||||
for i in 0..h {
|
||||
let dot: f32 = (0..h).map(|j| self.w1[i * h + j] * hidden[j]).sum::<f32>()
|
||||
+ self.b1[i];
|
||||
// SiLU(x) = x * sigmoid(x) = x / (1 + e^{-x})
|
||||
act[i] = dot * (1.0 / (1.0 + (-dot).exp()));
|
||||
}
|
||||
|
||||
// Layer 2: logits = W2 @ act + b2
|
||||
let v = self.vocab_size;
|
||||
let mut logits = vec![0.0_f32; v];
|
||||
for i in 0..v {
|
||||
logits[i] =
|
||||
(0..h).map(|j| self.w2[i * h + j] * act[j]).sum::<f32>() + self.b2[i];
|
||||
}
|
||||
logits
|
||||
}
|
||||
|
||||
/// Return the indices of the top-`k` tokens sorted descending by logit.
|
||||
///
|
||||
/// `k` is clamped to `vocab_size` if it exceeds it.
|
||||
pub fn top_k(&self, hidden: &[f32], k: usize) -> Vec<u32> {
|
||||
let logits = self.forward(hidden);
|
||||
let k = k.min(logits.len());
|
||||
let mut idx: Vec<u32> = (0..logits.len() as u32).collect();
|
||||
idx.sort_unstable_by(|&a, &b| {
|
||||
logits[b as usize]
|
||||
.partial_cmp(&logits[a as usize])
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
idx[..k].to_vec()
|
||||
}
|
||||
}
|
||||
|
||||
// ── MedusaConfig ─────────────────────────────────────────────────────────────
|
||||
|
||||
/// Configuration for the Medusa head ensemble.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MedusaConfig {
|
||||
/// `K`: number of draft heads (each predicts a different future offset).
|
||||
pub num_heads: usize,
|
||||
/// Hidden dimension shared by all heads and the base model.
|
||||
pub hidden_dim: usize,
|
||||
/// Vocabulary size.
|
||||
pub vocab_size: usize,
|
||||
/// Number of candidate tokens kept per head per position.
|
||||
pub top_k_per_head: usize,
|
||||
}
|
||||
|
||||
impl Default for MedusaConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
num_heads: 4,
|
||||
hidden_dim: 64,
|
||||
vocab_size: 256,
|
||||
top_k_per_head: 3,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── MedusaTree ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// A tree of draft candidate paths generated by [`MedusaHeads::generate_tree`].
|
||||
///
|
||||
/// Each path is a `Vec<u32>` of length `num_heads`, representing predicted
|
||||
/// tokens at offsets +1, +2, …, +K from the current position.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MedusaTree {
|
||||
/// All candidate paths. Length == `top_k_per_head ^ num_heads`.
|
||||
pub paths: Vec<Vec<u32>>,
|
||||
/// Raw logit vectors produced by each head, indexed `[head][vocab]`.
|
||||
pub head_logits: Vec<Vec<f32>>,
|
||||
}
|
||||
|
||||
impl MedusaTree {
|
||||
/// Number of candidate paths in the tree.
|
||||
pub fn num_paths(&self) -> usize {
|
||||
self.paths.len()
|
||||
}
|
||||
|
||||
/// Length of each path (equals `num_heads`).
|
||||
pub fn num_heads(&self) -> usize {
|
||||
self.paths.first().map_or(0, Vec::len)
|
||||
}
|
||||
}
|
||||
|
||||
// ── MedusaVerifyResult ────────────────────────────────────────────────────────
|
||||
|
||||
/// Result returned by [`MedusaHeads::verify`].
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MedusaVerifyResult {
|
||||
/// Longest accepted token sequence (always includes at least `base_next_token`).
|
||||
pub accepted_tokens: Vec<u32>,
|
||||
/// Index into [`MedusaTree::paths`] of the accepted path, or `None` if the
|
||||
/// draft was entirely rejected and only the base token was kept.
|
||||
pub accepted_path_idx: Option<usize>,
|
||||
/// Number of *speculative* tokens accepted beyond the base model token
|
||||
/// (i.e. `accepted_tokens.len() - 1`).
|
||||
pub num_accepted: usize,
|
||||
}
|
||||
|
||||
// ── MedusaLossResult ─────────────────────────────────────────────────────────
|
||||
|
||||
/// Cross-entropy losses returned by [`MedusaHeads::training_loss`].
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MedusaLossResult {
|
||||
/// Per-head cross-entropy loss, indexed `[head]`.
|
||||
pub per_head_loss: Vec<f32>,
|
||||
/// Mean cross-entropy across all heads.
|
||||
pub total_loss: f32,
|
||||
}
|
||||
|
||||
// ── MedusaHeads ──────────────────────────────────────────────────────────────
|
||||
|
||||
/// The complete Medusa head ensemble.
|
||||
///
|
||||
/// Holds K [`MedusaHead`]s and orchestrates tree generation, candidate
|
||||
/// verification, and training-loss computation.
|
||||
pub struct MedusaHeads {
|
||||
/// Shared configuration.
|
||||
pub config: MedusaConfig,
|
||||
/// Individual draft heads.
|
||||
pub heads: Vec<MedusaHead>,
|
||||
}
|
||||
|
||||
impl MedusaHeads {
|
||||
/// Create a new ensemble with random weights, seeding each head differently.
|
||||
pub fn new(config: MedusaConfig) -> Self {
|
||||
let heads: Vec<MedusaHead> = (0..config.num_heads)
|
||||
.map(|i| {
|
||||
MedusaHead::new_random(
|
||||
config.hidden_dim,
|
||||
config.vocab_size,
|
||||
// Distinct seed per head so weights differ.
|
||||
(i as u64).wrapping_mul(0xdead_beef_cafe_babe).wrapping_add(42),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
Self { config, heads }
|
||||
}
|
||||
|
||||
/// Generate a draft candidate tree from the base model's last hidden state.
|
||||
///
|
||||
/// Each head produces a top-`top_k_per_head` list; the tree is the
|
||||
/// cartesian product of these lists, giving `top_k_per_head ^ num_heads`
|
||||
/// paths, each of length `num_heads`.
|
||||
pub fn generate_tree(&self, hidden: &[f32]) -> MedusaTree {
|
||||
let k = self.config.top_k_per_head;
|
||||
|
||||
// Collect logits and top-k token indices from every head.
|
||||
let head_logits: Vec<Vec<f32>> = self.heads.iter().map(|h| h.forward(hidden)).collect();
|
||||
|
||||
let head_topk: Vec<Vec<u32>> = self
|
||||
.heads
|
||||
.iter()
|
||||
.zip(head_logits.iter())
|
||||
.map(|(_, logits)| {
|
||||
let mut idx: Vec<u32> = (0..logits.len() as u32).collect();
|
||||
idx.sort_unstable_by(|&a, &b| {
|
||||
logits[b as usize]
|
||||
.partial_cmp(&logits[a as usize])
|
||||
.unwrap_or(std::cmp::Ordering::Equal)
|
||||
});
|
||||
idx[..k.min(idx.len())].to_vec()
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Cartesian product: start with a single empty path and extend.
|
||||
let mut paths: Vec<Vec<u32>> = vec![vec![]];
|
||||
for topk in &head_topk {
|
||||
paths = paths
|
||||
.iter()
|
||||
.flat_map(|path| {
|
||||
topk.iter().map(move |&tok| {
|
||||
let mut p = path.clone();
|
||||
p.push(tok);
|
||||
p
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
}
|
||||
|
||||
MedusaTree { paths, head_logits }
|
||||
}
|
||||
|
||||
/// Verify draft candidates against base model predictions.
|
||||
///
|
||||
/// For each candidate path the verifier checks, token by token, whether the
|
||||
/// draft agrees with `path_oracle` (a closure that returns the base model's
|
||||
/// greedy token given a path prefix). The first token in every path must
|
||||
/// match `base_next_token`; subsequent tokens are checked via `path_oracle`.
|
||||
///
|
||||
/// When a mismatch is found the oracle's correction is appended and the
|
||||
/// comparison stops. The path with the longest accepted sequence is returned.
|
||||
///
|
||||
/// The result always contains at least `base_next_token`.
|
||||
pub fn verify(
|
||||
&self,
|
||||
tree: &MedusaTree,
|
||||
base_next_token: u32,
|
||||
path_oracle: &dyn Fn(&[u32]) -> u32,
|
||||
) -> MedusaVerifyResult {
|
||||
let mut best: Vec<u32> = vec![];
|
||||
let mut best_path_idx: Option<usize> = None;
|
||||
|
||||
for (idx, path) in tree.paths.iter().enumerate() {
|
||||
// The first element of every candidate path must be the base token.
|
||||
if path.is_empty() || path[0] != base_next_token {
|
||||
continue;
|
||||
}
|
||||
|
||||
let mut accepted = vec![base_next_token];
|
||||
|
||||
// Verify tokens path[1], path[2], … against the oracle.
|
||||
for (pos, &draft_tok) in path[1..].iter().enumerate() {
|
||||
// Oracle: given path[0..=pos], what would the base model predict?
|
||||
let expected = path_oracle(&path[..pos + 1]);
|
||||
if draft_tok == expected {
|
||||
accepted.push(draft_tok);
|
||||
} else {
|
||||
// Accept the oracle's correction and stop.
|
||||
accepted.push(expected);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if accepted.len() > best.len() {
|
||||
best = accepted;
|
||||
best_path_idx = Some(idx);
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback: if no path started with base_next_token, accept it alone.
|
||||
if best.is_empty() {
|
||||
best = vec![base_next_token];
|
||||
}
|
||||
|
||||
let num_accepted = best.len().saturating_sub(1);
|
||||
MedusaVerifyResult {
|
||||
accepted_tokens: best,
|
||||
accepted_path_idx: best_path_idx,
|
||||
num_accepted,
|
||||
}
|
||||
}
|
||||
|
||||
/// Compute per-head cross-entropy training loss.
|
||||
///
|
||||
/// `hidden` is the base model's last hidden state `[hidden_dim]`.
|
||||
/// `targets[i]` is the ground-truth token that head `i` should predict
|
||||
/// (i.e. the token at offset `+i+1` in the training sequence).
|
||||
///
|
||||
/// Loss is computed with the numerically stable log-sum-exp trick.
|
||||
///
|
||||
/// # Panics
|
||||
///
|
||||
/// Panics if `targets.len() != self.heads.len()`.
|
||||
pub fn training_loss(&self, hidden: &[f32], targets: &[u32]) -> MedusaLossResult {
|
||||
assert_eq!(
|
||||
targets.len(),
|
||||
self.heads.len(),
|
||||
"targets length must equal number of heads"
|
||||
);
|
||||
|
||||
let per_head_loss: Vec<f32> = self
|
||||
.heads
|
||||
.iter()
|
||||
.zip(targets.iter())
|
||||
.map(|(head, &target)| {
|
||||
let logits = head.forward(hidden);
|
||||
// Numerically stable CE: loss = log_sum_exp(logits) - logits[target]
|
||||
let max_l = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
|
||||
let log_sum =
|
||||
max_l + logits.iter().map(|l| (l - max_l).exp()).sum::<f32>().ln();
|
||||
-(logits[target as usize] - log_sum)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let total_loss =
|
||||
per_head_loss.iter().sum::<f32>() / per_head_loss.len() as f32;
|
||||
|
||||
MedusaLossResult { per_head_loss, total_loss }
|
||||
}
|
||||
}
|
||||
|
||||
// ── Tests ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// ── helpers ──
|
||||
|
||||
fn default_heads() -> MedusaHeads {
|
||||
MedusaHeads::new(MedusaConfig::default())
|
||||
}
|
||||
|
||||
fn zero_hidden(dim: usize) -> Vec<f32> {
|
||||
vec![0.0_f32; dim]
|
||||
}
|
||||
|
||||
fn ones_hidden(dim: usize) -> Vec<f32> {
|
||||
vec![1.0_f32; dim]
|
||||
}
|
||||
|
||||
// ── config ───────────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let cfg = MedusaConfig::default();
|
||||
assert_eq!(cfg.num_heads, 4);
|
||||
assert_eq!(cfg.hidden_dim, 64);
|
||||
assert_eq!(cfg.vocab_size, 256);
|
||||
assert_eq!(cfg.top_k_per_head, 3);
|
||||
}
|
||||
|
||||
// ── MedusaHead::forward ───────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_head_forward_shape() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let head = MedusaHead::new_random(cfg.hidden_dim, cfg.vocab_size, 1);
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let logits = head.forward(&hidden);
|
||||
assert_eq!(logits.len(), cfg.vocab_size);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_head_forward_finite() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let head = MedusaHead::new_random(cfg.hidden_dim, cfg.vocab_size, 2);
|
||||
let hidden = ones_hidden(cfg.hidden_dim);
|
||||
let logits = head.forward(&hidden);
|
||||
assert!(logits.iter().all(|v| v.is_finite()), "all logits must be finite");
|
||||
}
|
||||
|
||||
// ── MedusaHead::top_k ────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_head_top_k_length() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let head = MedusaHead::new_random(cfg.hidden_dim, cfg.vocab_size, 3);
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let k = 5;
|
||||
let result = head.top_k(&hidden, k);
|
||||
assert_eq!(result.len(), k);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_head_top_k_sorted() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let head = MedusaHead::new_random(cfg.hidden_dim, cfg.vocab_size, 4);
|
||||
let hidden = ones_hidden(cfg.hidden_dim);
|
||||
let result = head.top_k(&hidden, 8);
|
||||
let logits = head.forward(&hidden);
|
||||
// Each successive index should have a logit <= the previous.
|
||||
for w in result.windows(2) {
|
||||
assert!(
|
||||
logits[w[0] as usize] >= logits[w[1] as usize],
|
||||
"top-k must be sorted descending by logit"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_head_top_k_unique() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let head = MedusaHead::new_random(cfg.hidden_dim, cfg.vocab_size, 5);
|
||||
let hidden = ones_hidden(cfg.hidden_dim);
|
||||
let k = 10;
|
||||
let result = head.top_k(&hidden, k);
|
||||
let mut sorted = result.clone();
|
||||
sorted.sort_unstable();
|
||||
sorted.dedup();
|
||||
assert_eq!(sorted.len(), result.len(), "top-k indices must be unique");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_head_top_k_clamped_to_vocab() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let head = MedusaHead::new_random(cfg.hidden_dim, cfg.vocab_size, 6);
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
// Request more than vocab_size — should not panic, returns vocab_size items.
|
||||
let result = head.top_k(&hidden, cfg.vocab_size + 100);
|
||||
assert_eq!(result.len(), cfg.vocab_size);
|
||||
}
|
||||
|
||||
// ── SiLU activation ──────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_head_silu_activation_at_zero() {
|
||||
// When the hidden state is zero AND biases are zero, W1@0 + b1 = 0,
|
||||
// so the pre-activation is 0 for every neuron. SiLU(0) = 0 * 0.5 = 0.
|
||||
// Build a head with zero weights to isolate this property.
|
||||
let h = 4;
|
||||
let v = 8;
|
||||
let head = MedusaHead {
|
||||
w1: vec![0.0; h * h],
|
||||
b1: vec![0.0; h],
|
||||
w2: vec![1.0; v * h], // uniform W2 so logits == sum(act) = 0
|
||||
b2: vec![0.0; v],
|
||||
hidden_dim: h,
|
||||
vocab_size: v,
|
||||
};
|
||||
let logits = head.forward(&vec![0.0; h]);
|
||||
// All activations are SiLU(0) = 0, so all logits should be 0.
|
||||
for l in &logits {
|
||||
assert!(
|
||||
l.abs() < 1e-6,
|
||||
"SiLU(0) path: expected logit ≈ 0, got {l}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_head_silu_activation_positive_input() {
|
||||
// With positive pre-activation x > 0, SiLU(x) = x * sigmoid(x) > 0.
|
||||
let h = 2;
|
||||
let v = 2;
|
||||
// Identity W1 with positive bias forces positive pre-activation.
|
||||
let w1 = vec![1.0_f32, 0.0, 0.0, 1.0]; // 2×2 identity
|
||||
let b1 = vec![1.0_f32, 1.0]; // shift pre-activation up
|
||||
let w2 = vec![1.0_f32; v * h];
|
||||
let b2 = vec![0.0_f32; v];
|
||||
let head = MedusaHead { w1, b1, w2, b2, hidden_dim: h, vocab_size: v };
|
||||
let logits = head.forward(&vec![0.0; h]);
|
||||
// Pre-activation = 1.0; SiLU(1.0) = 1 / (1 + e^{-1}) ≈ 0.731 > 0.
|
||||
for l in &logits {
|
||||
assert!(*l > 0.0, "SiLU of positive input should produce positive output, got {l}");
|
||||
}
|
||||
}
|
||||
|
||||
// ── MedusaTree ────────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_generate_tree_num_paths() {
|
||||
let cfg = MedusaConfig::default(); // top_k=3, num_heads=4
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
let expected = cfg.top_k_per_head.pow(cfg.num_heads as u32);
|
||||
assert_eq!(tree.num_paths(), expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generate_tree_path_length() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
for path in &tree.paths {
|
||||
assert_eq!(path.len(), cfg.num_heads, "each path must span all heads");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_generate_tree_head_logits_count() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
assert_eq!(tree.head_logits.len(), cfg.num_heads);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tree_num_paths_accessor() {
|
||||
let tree = MedusaTree {
|
||||
paths: vec![vec![1, 2], vec![3, 4], vec![5, 6]],
|
||||
head_logits: vec![],
|
||||
};
|
||||
assert_eq!(tree.num_paths(), 3);
|
||||
assert_eq!(tree.num_heads(), 2);
|
||||
}
|
||||
|
||||
// ── verify ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Build a simple 1-head tree (top_k=3) and control path_oracle precisely.
|
||||
fn small_heads() -> (MedusaHeads, MedusaConfig) {
|
||||
let cfg = MedusaConfig {
|
||||
num_heads: 1,
|
||||
hidden_dim: 4,
|
||||
vocab_size: 8,
|
||||
top_k_per_head: 3,
|
||||
};
|
||||
(MedusaHeads::new(cfg.clone()), cfg)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_accepts_base_token() {
|
||||
let (heads, cfg) = small_heads();
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
let base_token = 99_u32; // not in any path (paths ⊂ [0,7])
|
||||
// oracle never called because no path starts with 99
|
||||
let result = heads.verify(&tree, base_token, &|_| unreachable!());
|
||||
assert_eq!(result.accepted_tokens, vec![base_token]);
|
||||
assert_eq!(result.num_accepted, 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_perfect_draft() {
|
||||
// 1-head tree: a path [tok] matches base_next_token == tok.
|
||||
// Since num_heads==1, after accepting tok the loop over path[1..] is empty,
|
||||
// so the whole length-1 path is accepted.
|
||||
let (heads, cfg) = small_heads();
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
|
||||
// Find the first path's leading token and use it as base_next_token.
|
||||
let first_tok = tree.paths[0][0];
|
||||
let result = heads.verify(&tree, first_tok, &|_| first_tok);
|
||||
// accepted_tokens contains exactly [first_tok]; num_accepted == 0 (single head)
|
||||
assert!(result.accepted_tokens.contains(&first_tok));
|
||||
assert!(result.accepted_path_idx.is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_perfect_draft_multi_head() {
|
||||
// 2-head tree, oracle always agrees with draft → full path accepted.
|
||||
let cfg = MedusaConfig {
|
||||
num_heads: 2,
|
||||
hidden_dim: 4,
|
||||
vocab_size: 8,
|
||||
top_k_per_head: 2,
|
||||
};
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
|
||||
// Pick a path and set oracle to always confirm draft.
|
||||
let path = tree.paths[0].clone();
|
||||
let base = path[0];
|
||||
let draft_second = path[1];
|
||||
let result = heads.verify(&tree, base, &move |_prefix| draft_second);
|
||||
|
||||
// Should accept both tokens.
|
||||
assert_eq!(result.accepted_tokens.len(), 2);
|
||||
assert_eq!(result.num_accepted, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_mismatch_accepts_correction() {
|
||||
// 2-head tree: oracle disagrees at position 1 → correction is appended.
|
||||
let cfg = MedusaConfig {
|
||||
num_heads: 2,
|
||||
hidden_dim: 4,
|
||||
vocab_size: 8,
|
||||
top_k_per_head: 2,
|
||||
};
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
|
||||
let path = tree.paths[0].clone();
|
||||
let base = path[0];
|
||||
let correction = 200_u32; // guaranteed to differ from any vocab token (0..7)
|
||||
let result = heads.verify(&tree, base, &move |_prefix| correction);
|
||||
|
||||
// base accepted, then correction appended → length 2, num_accepted = 1
|
||||
assert_eq!(result.accepted_tokens.len(), 2);
|
||||
assert_eq!(*result.accepted_tokens.last().unwrap(), correction);
|
||||
assert_eq!(result.num_accepted, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_verify_num_accepted_equals_extra_tokens() {
|
||||
let (heads, cfg) = small_heads();
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let tree = heads.generate_tree(&hidden);
|
||||
let base = tree.paths[0][0];
|
||||
let result = heads.verify(&tree, base, &|_| base);
|
||||
// num_accepted == accepted_tokens.len() - 1
|
||||
assert_eq!(result.num_accepted, result.accepted_tokens.len().saturating_sub(1));
|
||||
}
|
||||
|
||||
// ── training_loss ─────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn test_training_loss_shape() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let targets: Vec<u32> = (0..cfg.num_heads as u32).collect();
|
||||
let result = heads.training_loss(&hidden, &targets);
|
||||
assert_eq!(result.per_head_loss.len(), cfg.num_heads);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_training_loss_positive() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = ones_hidden(cfg.hidden_dim);
|
||||
let targets = vec![0_u32; cfg.num_heads];
|
||||
let result = heads.training_loss(&hidden, &targets);
|
||||
for (i, &loss) in result.per_head_loss.iter().enumerate() {
|
||||
assert!(loss >= 0.0, "head {i}: CE loss must be non-negative, got {loss}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_training_loss_total_is_mean() {
|
||||
let cfg = MedusaConfig::default();
|
||||
let heads = MedusaHeads::new(cfg.clone());
|
||||
let hidden = zero_hidden(cfg.hidden_dim);
|
||||
let targets = vec![1_u32; cfg.num_heads];
|
||||
let result = heads.training_loss(&hidden, &targets);
|
||||
let expected_mean =
|
||||
result.per_head_loss.iter().sum::<f32>() / result.per_head_loss.len() as f32;
|
||||
assert!(
|
||||
(result.total_loss - expected_mean).abs() < 1e-5,
|
||||
"total_loss must equal mean of per_head_loss"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_training_loss_near_zero_for_dominant_logit() {
|
||||
// Construct a head whose W2 row for token 0 is large (+100)
|
||||
// while all other rows are small (0). CE loss for target=0 should be ≈ 0.
|
||||
let h = 4;
|
||||
let v = 8;
|
||||
let target: u32 = 0;
|
||||
|
||||
// Build: W1 = identity, b1 = 0 → act = SiLU(hidden)
|
||||
// W2 row 0 = [100, 100, 100, 100], other rows = 0 → logit[0] >> logit[k>0]
|
||||
let mut w2 = vec![0.0_f32; v * h];
|
||||
for j in 0..h {
|
||||
w2[target as usize * h + j] = 100.0;
|
||||
}
|
||||
let head = MedusaHead {
|
||||
w1: {
|
||||
let mut m = vec![0.0_f32; h * h];
|
||||
for i in 0..h { m[i * h + i] = 1.0; } // identity
|
||||
m
|
||||
},
|
||||
b1: vec![0.0; h],
|
||||
w2,
|
||||
b2: vec![0.0; v],
|
||||
hidden_dim: h,
|
||||
vocab_size: v,
|
||||
};
|
||||
|
||||
// hidden = ones so act = SiLU(1) per neuron, all positive.
|
||||
let hidden = vec![1.0_f32; h];
|
||||
let logits = head.forward(&hidden);
|
||||
let max_l = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
|
||||
let log_sum = max_l + logits.iter().map(|l| (l - max_l).exp()).sum::<f32>().ln();
|
||||
let ce = -(logits[target as usize] - log_sum);
|
||||
assert!(
|
||||
ce < 0.01,
|
||||
"CE loss should be near 0 when target logit dominates, got {ce}"
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user