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

- 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:
Omar Sobh
2026-06-27 13:07:39 +00:00
co-authored by Claude Sonnet 4.6
parent 960fd73c82
commit bff27c302f
6 changed files with 2480 additions and 0 deletions
@@ -24,6 +24,16 @@ pub use logit_processors::{
LogitProcessor, LogitProcessorList, MinPProcessor, PresencePenaltyProcessor, LogitProcessor, LogitProcessorList, MinPProcessor, PresencePenaltyProcessor,
RepetitionPenaltyProcessor, TemperatureProcessor, TopKProcessor, TopPProcessor, RepetitionPenaltyProcessor, TemperatureProcessor, TopKProcessor, TopPProcessor,
}; };
pub mod medusa;
pub use medusa::{
MedusaHead,
MedusaConfig as MedusaHeadsConfig,
MedusaHeads,
MedusaLossResult,
MedusaTree,
MedusaVerifyResult,
};
pub mod batch_processor; pub mod batch_processor;
pub mod cache; pub mod cache;
pub mod gqa; pub mod gqa;
@@ -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}"
);
}
}
@@ -0,0 +1,807 @@
//! Mixture of Depths (MoD) — dynamic token routing through transformer layers.
//!
//! Implements the method from Raposo et al. 2024 (arXiv:2404.02258).
//! A learned router selects a fraction of tokens to process at each layer;
//! the rest skip the layer via a residual bypass. This reduces total FLOPs
//! proportionally to the skipped fraction without degrading model quality.
//!
//! # Examples
//!
//! ```rust
//! use rtx_transformers::layers::mixture_of_depths::{MoDConfig, MoDLayer};
//!
//! let config = MoDConfig::new(64);
//! let layer = MoDLayer::with_zero_router(config);
//!
//! // Identity layer_fn: output should equal input
//! let hidden: Vec<f32> = (0..8 * 64).map(|i| i as f32).collect();
//! let (output, stats) = layer.forward(&hidden, 8, |tokens, _k| tokens.to_vec());
//! assert_eq!(output.len(), hidden.len());
//! assert!(stats.aux_loss > 0.0);
//! ```
// ============================================================================
// Router
// ============================================================================
/// Configuration for a Mixture of Depths layer.
#[derive(Debug, Clone)]
pub struct MoDConfig {
/// Fraction of tokens to process at this layer (0.0, 1.0].
pub capacity_fraction: f32,
/// Hidden dimension of each token embedding.
pub hidden_dim: usize,
/// Weight for the load-balancing auxiliary loss term.
pub aux_loss_weight: f32,
}
impl MoDConfig {
/// Create a config with default `capacity_fraction = 0.5` and
/// `aux_loss_weight = 0.01`.
pub fn new(hidden_dim: usize) -> Self {
Self {
capacity_fraction: 0.5,
hidden_dim,
aux_loss_weight: 0.01,
}
}
}
// ─────────────────────────────────────────────────────────────────────────────
/// A lightweight linear router: `[hidden_dim → 1]`.
///
/// Computes a scalar score per token via a dot product plus bias. No
/// activation is applied to the raw score; `sigmoid` is used separately to
/// produce probabilities for the auxiliary loss.
#[derive(Debug, Clone)]
pub struct MoDRouter {
/// Weight vector of length `hidden_dim`.
pub weight: Vec<f32>,
/// Scalar bias term.
pub bias: f32,
/// Input dimension.
pub hidden_dim: usize,
}
impl MoDRouter {
/// Initialise router weights with a simple LCG pseudo-random sequence
/// scaled to `[-0.1, 0.1]` using the given `seed`.
///
/// This gives deterministic, non-zero weights without pulling in external
/// crates.
pub fn new_random(hidden_dim: usize, seed: u64) -> Self {
// Linear congruential generator — good enough for weight initialisation.
let mut state = seed.wrapping_add(1);
let scale = 0.2 / u64::MAX as f64;
let weight: Vec<f32> = (0..hidden_dim)
.map(|_| {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(state as f64 * scale - 0.1) as f32
})
.collect();
Self {
weight,
bias: 0.0,
hidden_dim,
}
}
/// Initialise router with all-zero weights (useful for unit tests that
/// need fully deterministic, equal scores).
pub fn new_zeros(hidden_dim: usize) -> Self {
Self {
weight: vec![0.0_f32; hidden_dim],
bias: 0.0,
hidden_dim,
}
}
/// Compute the scalar router score for a single token embedding.
///
/// `hidden` must have length `self.hidden_dim`.
pub fn score(&self, hidden: &[f32]) -> f32 {
debug_assert_eq!(hidden.len(), self.hidden_dim);
let dot: f32 = hidden
.iter()
.zip(self.weight.iter())
.map(|(&h, &w)| h * w)
.sum();
dot + self.bias
}
/// Compute scores for every token in a sequence.
///
/// `hidden` is row-major with shape `[seq_len × hidden_dim]`.
/// Returns a `Vec<f32>` of length `seq_len`.
pub fn scores(&self, hidden: &[f32], seq_len: usize) -> Vec<f32> {
let d = self.hidden_dim;
debug_assert_eq!(hidden.len(), seq_len * d);
(0..seq_len)
.map(|i| self.score(&hidden[i * d..(i + 1) * d]))
.collect()
}
/// Compute sigmoid-activated routing probabilities for every token.
///
/// Returns a `Vec<f32>` of length `seq_len` with values in `(0, 1)`.
pub fn probs(&self, hidden: &[f32], seq_len: usize) -> Vec<f32> {
self.scores(hidden, seq_len)
.into_iter()
.map(sigmoid)
.collect()
}
/// Select the top-`capacity` token indices by router score.
///
/// Returns the indices in **ascending** (original sequence) order so that
/// the processed sub-sequence preserves positional structure.
pub fn select_tokens(&self, hidden: &[f32], seq_len: usize, capacity: usize) -> Vec<usize> {
let scores = self.scores(hidden, seq_len);
let capacity = capacity.min(seq_len);
// Build (index, score) pairs and sort descending by score.
let mut indexed: Vec<(usize, f32)> = scores
.iter()
.enumerate()
.map(|(i, &s)| (i, s))
.collect();
indexed.sort_unstable_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
});
// Take the top-capacity entries and restore original order.
let mut selected: Vec<usize> = indexed[..capacity].iter().map(|(i, _)| *i).collect();
selected.sort_unstable();
selected
}
}
// ─────────────────────────────────────────────────────────────────────────────
/// Result of the token selection step for one MoD layer.
#[derive(Debug, Clone)]
pub struct TokenSelection {
/// Indices of tokens that will be processed by the layer (ascending).
pub selected_indices: Vec<usize>,
/// Indices of tokens that bypass the layer unchanged (ascending).
pub skipped_indices: Vec<usize>,
/// Raw router scores, length `seq_len`.
pub router_scores: Vec<f32>,
/// Sigmoid of router scores, length `seq_len`.
pub router_probs: Vec<f32>,
}
// ─────────────────────────────────────────────────────────────────────────────
/// Per-layer statistics from a single MoD forward pass.
#[derive(Debug, Clone)]
pub struct MoDStats {
/// Sequence length.
pub seq_len: usize,
/// Maximum number of tokens the layer is allowed to process.
pub capacity: usize,
/// Actual number of tokens routed to the layer.
pub selected_count: usize,
/// Number of tokens that bypassed the layer.
pub skipped_count: usize,
/// `selected_count / seq_len`.
pub selection_fraction: f32,
/// Load-balancing auxiliary loss value.
pub aux_loss: f32,
/// Mean raw router score across the full sequence.
pub avg_router_score: f32,
}
// ============================================================================
// MoDLayer
// ============================================================================
/// A Mixture of Depths layer wrapper.
///
/// Wraps any token-processing function with MoD routing. The router selects
/// `floor(seq_len * capacity_fraction)` tokens; the selected tokens are passed
/// to the inner `layer_fn`; all other tokens bypass unchanged.
pub struct MoDLayer {
/// Configuration for this layer.
pub config: MoDConfig,
/// The learned router.
pub router: MoDRouter,
}
impl MoDLayer {
/// Create a new `MoDLayer` with randomly-initialised router weights.
pub fn new(config: MoDConfig, seed: u64) -> Self {
let router = MoDRouter::new_random(config.hidden_dim, seed);
Self { config, router }
}
/// Create a `MoDLayer` whose router weights are all zero.
///
/// With zero weights every token receives the same score, so the first
/// `capacity` tokens in sequence order will be selected. This is useful
/// for deterministic unit tests.
pub fn with_zero_router(config: MoDConfig) -> Self {
let router = MoDRouter::new_zeros(config.hidden_dim);
Self { config, router }
}
// ── helpers ─────────────────────────────────────────────────────────────
/// Compute capacity (minimum 1) for a given sequence length.
fn capacity_for(&self, seq_len: usize) -> usize {
((seq_len as f32 * self.config.capacity_fraction) as usize).max(1)
}
// ── public API ──────────────────────────────────────────────────────────
/// Determine which tokens to process.
pub fn select(&self, hidden: &[f32], seq_len: usize) -> TokenSelection {
let capacity = self.capacity_for(seq_len);
let router_scores = self.router.scores(hidden, seq_len);
let router_probs: Vec<f32> = router_scores.iter().copied().map(sigmoid).collect();
let selected_indices = self
.router
.select_tokens(hidden, seq_len, capacity);
// Build the skipped set as the complement of selected.
let mut selected_set = vec![false; seq_len];
for &i in &selected_indices {
selected_set[i] = true;
}
let skipped_indices: Vec<usize> = (0..seq_len).filter(|&i| !selected_set[i]).collect();
TokenSelection {
selected_indices,
skipped_indices,
router_scores,
router_probs,
}
}
/// Forward pass with MoD routing.
///
/// - `hidden`: row-major `[seq_len × hidden_dim]`
/// - `layer_fn`: processes the selected token sub-batch
/// `(tokens_flat: &[f32], num_tokens: usize) -> Vec<f32>` where the
/// output must have the same length as the input.
///
/// Returns `(output, stats)` where `output` has the same shape as
/// `hidden`: selected tokens are replaced with `layer_fn` output; skipped
/// tokens carry their original values unchanged.
pub fn forward<F>(
&self,
hidden: &[f32],
seq_len: usize,
layer_fn: F,
) -> (Vec<f32>, MoDStats)
where
F: Fn(&[f32], usize) -> Vec<f32>,
{
let d = self.config.hidden_dim;
let capacity = self.capacity_for(seq_len);
let selection = self.select(hidden, seq_len);
// Gather selected tokens into a contiguous flat buffer.
let selected_flat: Vec<f32> = selection
.selected_indices
.iter()
.flat_map(|&i| &hidden[i * d..(i + 1) * d])
.copied()
.collect();
let k = selection.selected_indices.len();
let processed = if k > 0 {
layer_fn(&selected_flat, k)
} else {
selected_flat
};
// Scatter processed tokens back; skipped tokens keep original values.
let mut output = hidden.to_vec();
for (slot, &orig_idx) in selection.selected_indices.iter().enumerate() {
output[orig_idx * d..(orig_idx + 1) * d]
.copy_from_slice(&processed[slot * d..(slot + 1) * d]);
}
let aux = self.aux_loss(&selection, seq_len);
let avg_score = selection.router_scores.iter().sum::<f32>() / seq_len as f32;
let stats = MoDStats {
seq_len,
capacity,
selected_count: selection.selected_indices.len(),
skipped_count: selection.skipped_indices.len(),
selection_fraction: selection.selected_indices.len() as f32 / seq_len as f32,
aux_loss: aux,
avg_router_score: avg_score,
};
(output, stats)
}
/// Compute the load-balancing auxiliary loss.
///
/// Encourages the router to select tokens uniformly rather than always
/// favouring the same positions:
///
/// ```text
/// aux_loss = aux_loss_weight * mean(router_probs) * seq_len / capacity
/// ```
pub fn aux_loss(&self, selection: &TokenSelection, seq_len: usize) -> f32 {
let capacity = self.capacity_for(seq_len) as f32;
let mean_prob = selection.router_probs.iter().sum::<f32>() / seq_len as f32;
self.config.aux_loss_weight * mean_prob * seq_len as f32 / capacity
}
}
// ============================================================================
// MoDStack
// ============================================================================
/// A sequential stack of MoD layers, each with independent routers.
///
/// Useful for wrapping a full transformer's layers with per-layer capacity
/// budgets.
pub struct MoDStack {
/// The individual MoD layers.
pub layers: Vec<MoDLayer>,
}
impl MoDStack {
/// Create a stack where every layer uses the same `capacity_fraction`.
///
/// Each layer receives a distinct seed derived from the base `seed` so
/// that their routers are initialised independently.
pub fn new_uniform(
num_layers: usize,
hidden_dim: usize,
capacity_fraction: f32,
seed: u64,
) -> Self {
let layers = (0..num_layers)
.map(|i| {
let config = MoDConfig {
capacity_fraction,
hidden_dim,
aux_loss_weight: 0.01,
};
MoDLayer::new(config, seed.wrapping_add(i as u64))
})
.collect();
Self { layers }
}
/// Estimated FLOPs reduction relative to processing every token at every
/// layer.
///
/// With uniform `capacity_fraction = c` across `L` layers the reduction
/// factor is `c^L` (each layer independently reduces computation by `c`).
///
/// For non-uniform fractions this returns the product of all
/// `capacity_fraction` values.
pub fn flops_reduction(&self) -> f32 {
self.layers
.iter()
.map(|l| l.config.capacity_fraction)
.product()
}
/// Forward pass through every layer in sequence.
///
/// `layer_fns` must have the same length as `self.layers`. Each entry is
/// the inner computation to be conditionally applied at that depth.
///
/// Returns `(output, stats_per_layer)`.
pub fn forward<F>(
&self,
hidden: &[f32],
seq_len: usize,
layer_fns: &[F],
) -> (Vec<f32>, Vec<MoDStats>)
where
F: Fn(&[f32], usize) -> Vec<f32>,
{
assert_eq!(
self.layers.len(),
layer_fns.len(),
"layer_fns must have the same length as the stack"
);
let mut current = hidden.to_vec();
let mut all_stats = Vec::with_capacity(self.layers.len());
for (layer, layer_fn) in self.layers.iter().zip(layer_fns.iter()) {
let (next, stats) = layer.forward(&current, seq_len, layer_fn);
current = next;
all_stats.push(stats);
}
(current, all_stats)
}
}
// ============================================================================
// Internal helpers
// ============================================================================
/// Numerically stable sigmoid: `1 / (1 + exp(-x))`.
#[inline]
fn sigmoid(x: f32) -> f32 {
if x >= 0.0 {
let e = (-x).exp();
1.0 / (1.0 + e)
} else {
let e = x.exp();
e / (1.0 + e)
}
}
// ============================================================================
// Tests
// ============================================================================
#[cfg(test)]
mod tests {
use super::*;
// ── helpers ─────────────────────────────────────────────────────────────
fn make_hidden(seq_len: usize, hidden_dim: usize) -> Vec<f32> {
(0..seq_len * hidden_dim).map(|i| i as f32).collect()
}
fn identity_fn(tokens: &[f32], _k: usize) -> Vec<f32> {
tokens.to_vec()
}
fn double_fn(tokens: &[f32], _k: usize) -> Vec<f32> {
tokens.iter().map(|&v| v * 2.0).collect()
}
// ── MoDConfig ────────────────────────────────────────────────────────────
#[test]
fn test_config_default() {
let cfg = MoDConfig::new(128);
assert_eq!(cfg.capacity_fraction, 0.5);
assert_eq!(cfg.hidden_dim, 128);
assert!((cfg.aux_loss_weight - 0.01).abs() < 1e-6);
}
// ── MoDRouter ────────────────────────────────────────────────────────────
#[test]
fn test_router_score_shape() {
let seq_len = 10;
let hidden_dim = 8;
let router = MoDRouter::new_random(hidden_dim, 42);
let hidden = make_hidden(seq_len, hidden_dim);
let scores = router.scores(&hidden, seq_len);
assert_eq!(scores.len(), seq_len);
}
#[test]
fn test_router_zero_scores_all_equal() {
let seq_len = 6;
let hidden_dim = 4;
let router = MoDRouter::new_zeros(hidden_dim);
let hidden = make_hidden(seq_len, hidden_dim);
let scores = router.scores(&hidden, seq_len);
// All weights zero, bias zero → every dot product is 0.
for &s in &scores {
assert!((s - 0.0).abs() < 1e-6, "expected 0, got {s}");
}
}
#[test]
fn test_router_probs_in_range() {
let seq_len = 12;
let hidden_dim = 16;
let router = MoDRouter::new_random(hidden_dim, 7);
let hidden = make_hidden(seq_len, hidden_dim);
let probs = router.probs(&hidden, seq_len);
assert_eq!(probs.len(), seq_len);
for &p in &probs {
assert!(p >= 0.0 && p <= 1.0, "prob out of range: {p}");
}
}
#[test]
fn test_router_probs_sigmoid() {
// sigmoid(0) == 0.5
let hidden_dim = 4;
let router = MoDRouter::new_zeros(hidden_dim);
// Input doesn't matter; weights are zero so score is always 0.
let hidden = vec![1.0_f32; hidden_dim];
let p = router.probs(&hidden, 1);
assert!((p[0] - 0.5).abs() < 1e-6, "expected 0.5, got {}", p[0]);
}
// ── select_tokens ────────────────────────────────────────────────────────
#[test]
fn test_select_tokens_count() {
let seq_len = 8;
let hidden_dim = 4;
let capacity = 3;
let router = MoDRouter::new_random(hidden_dim, 99);
let hidden = make_hidden(seq_len, hidden_dim);
let selected = router.select_tokens(&hidden, seq_len, capacity);
assert_eq!(selected.len(), capacity);
}
#[test]
fn test_select_tokens_sorted() {
let seq_len = 10;
let hidden_dim = 8;
let capacity = 4;
let router = MoDRouter::new_random(hidden_dim, 13);
let hidden = make_hidden(seq_len, hidden_dim);
let selected = router.select_tokens(&hidden, seq_len, capacity);
for w in selected.windows(2) {
assert!(w[0] < w[1], "indices must be in ascending order");
}
}
#[test]
fn test_select_tokens_highest_score() {
// Build a router whose weight is 1 for dim-0, 0 elsewhere.
// Token scores then equal the first element of each token embedding.
let hidden_dim = 4;
let seq_len = 5;
let router = MoDRouter {
weight: vec![1.0, 0.0, 0.0, 0.0],
bias: 0.0,
hidden_dim,
};
// Tokens: each token's first element is its index, so token 4 has the
// highest score.
let hidden: Vec<f32> = (0..seq_len)
.flat_map(|i| vec![i as f32, 0.0, 0.0, 0.0])
.collect();
let selected = router.select_tokens(&hidden, seq_len, 1);
assert_eq!(selected, vec![4], "token 4 should have the highest score");
}
#[test]
fn test_select_tokens_cap_at_seq_len() {
// When capacity >= seq_len every token is selected.
let seq_len = 4;
let hidden_dim = 4;
let router = MoDRouter::new_random(hidden_dim, 1);
let hidden = make_hidden(seq_len, hidden_dim);
let selected = router.select_tokens(&hidden, seq_len, seq_len + 10);
assert_eq!(selected.len(), seq_len);
}
// ── TokenSelection ───────────────────────────────────────────────────────
#[test]
fn test_selection_disjoint() {
let seq_len = 8;
let hidden_dim = 4;
let config = MoDConfig::new(hidden_dim);
let layer = MoDLayer::new(config, 55);
let hidden = make_hidden(seq_len, hidden_dim);
let sel = layer.select(&hidden, seq_len);
let mut all: Vec<usize> = sel
.selected_indices
.iter()
.chain(sel.skipped_indices.iter())
.copied()
.collect();
all.sort_unstable();
assert_eq!(all, (0..seq_len).collect::<Vec<_>>());
}
// ── MoDLayer::forward ────────────────────────────────────────────────────
#[test]
fn test_forward_shape() {
let seq_len = 6;
let hidden_dim = 8;
let config = MoDConfig::new(hidden_dim);
let layer = MoDLayer::new(config, 1);
let hidden = make_hidden(seq_len, hidden_dim);
let (output, _) = layer.forward(&hidden, seq_len, identity_fn);
assert_eq!(output.len(), hidden.len());
}
#[test]
fn test_forward_skipped_unchanged() {
let seq_len = 8;
let hidden_dim = 4;
let config = MoDConfig::new(hidden_dim);
let layer = MoDLayer::new(config, 77);
let hidden = make_hidden(seq_len, hidden_dim);
let sel = layer.select(&hidden, seq_len);
let (output, _) = layer.forward(&hidden, seq_len, double_fn);
for &i in &sel.skipped_indices {
for j in 0..hidden_dim {
assert_eq!(
output[i * hidden_dim + j],
hidden[i * hidden_dim + j],
"skipped token {i} element {j} must be unchanged"
);
}
}
}
#[test]
fn test_forward_selected_modified() {
// Use double_fn so we know selected tokens are different from input.
let seq_len = 8;
let hidden_dim = 4;
let config = MoDConfig {
capacity_fraction: 0.5,
hidden_dim,
aux_loss_weight: 0.01,
};
let layer = MoDLayer::new(config, 88);
let hidden = make_hidden(seq_len, hidden_dim);
let sel = layer.select(&hidden, seq_len);
let (output, _) = layer.forward(&hidden, seq_len, double_fn);
// At least one selected element must differ from the original.
let any_changed = sel.selected_indices.iter().any(|&i| {
(0..hidden_dim).any(|j| {
(output[i * hidden_dim + j] - hidden[i * hidden_dim + j]).abs() > 1e-6
})
});
assert!(
any_changed,
"at least one selected token must be modified by double_fn"
);
}
#[test]
fn test_forward_identity_layer() {
// Identity layer_fn: output must equal input everywhere.
let seq_len = 6;
let hidden_dim = 8;
let config = MoDConfig::new(hidden_dim);
let layer = MoDLayer::new(config, 2);
let hidden = make_hidden(seq_len, hidden_dim);
let (output, _) = layer.forward(&hidden, seq_len, identity_fn);
for (a, b) in hidden.iter().zip(output.iter()) {
assert!((a - b).abs() < 1e-6);
}
}
#[test]
fn test_forward_capacity_fraction() {
// With capacity_fraction = 0.5 and seq_len = 10, exactly 5 tokens
// should be selected.
let seq_len = 10;
let hidden_dim = 4;
let config = MoDConfig {
capacity_fraction: 0.5,
hidden_dim,
aux_loss_weight: 0.01,
};
let layer = MoDLayer::new(config, 3);
let hidden = make_hidden(seq_len, hidden_dim);
let (_, stats) = layer.forward(&hidden, seq_len, identity_fn);
assert_eq!(stats.selected_count, 5);
assert_eq!(stats.skipped_count, 5);
}
// ── aux_loss ─────────────────────────────────────────────────────────────
#[test]
fn test_aux_loss_positive() {
let seq_len = 8;
let hidden_dim = 4;
let config = MoDConfig::new(hidden_dim);
let layer = MoDLayer::new(config, 5);
let hidden = make_hidden(seq_len, hidden_dim);
let sel = layer.select(&hidden, seq_len);
let loss = layer.aux_loss(&sel, seq_len);
assert!(loss > 0.0, "aux_loss must be positive");
}
#[test]
fn test_aux_loss_uniform_probs() {
// When all router probabilities are 0.5 (zero weights → sigmoid(0)):
// aux_loss = aux_loss_weight * 0.5 * seq_len / capacity
// With capacity_fraction = 0.5: capacity = seq_len/2
// → aux_loss = 0.01 * 0.5 * 2 = 0.01
let seq_len = 8;
let hidden_dim = 4;
let config = MoDConfig {
capacity_fraction: 0.5,
hidden_dim,
aux_loss_weight: 0.01,
};
let layer = MoDLayer::with_zero_router(config);
let hidden = make_hidden(seq_len, hidden_dim);
let sel = layer.select(&hidden, seq_len);
let loss = layer.aux_loss(&sel, seq_len);
let expected = 0.01_f32 * 0.5 * seq_len as f32 / (seq_len as f32 * 0.5);
assert!(
(loss - expected).abs() < 1e-5,
"expected {expected}, got {loss}"
);
}
// ── MoDStack ─────────────────────────────────────────────────────────────
#[test]
fn test_mod_stack_creation() {
let num_layers = 6;
let stack = MoDStack::new_uniform(num_layers, 32, 0.5, 10);
assert_eq!(stack.layers.len(), num_layers);
}
#[test]
fn test_mod_stack_flops_reduction() {
let num_layers = 4;
let fraction = 0.5_f32;
let stack = MoDStack::new_uniform(num_layers, 16, fraction, 0);
let expected = fraction.powi(num_layers as i32);
let got = stack.flops_reduction();
assert!(
(got - expected).abs() < 1e-5,
"expected {expected}, got {got}"
);
}
#[test]
fn test_mod_stack_forward_output_shape() {
let seq_len = 8;
let hidden_dim = 4;
let num_layers = 3;
let stack = MoDStack::new_uniform(num_layers, hidden_dim, 0.5, 20);
let hidden = make_hidden(seq_len, hidden_dim);
let fns: Vec<fn(&[f32], usize) -> Vec<f32>> = vec![identity_fn; num_layers];
let (output, _) = stack.forward(&hidden, seq_len, &fns);
assert_eq!(output.len(), hidden.len());
}
#[test]
fn test_mod_stack_forward_stats_count() {
let seq_len = 8;
let hidden_dim = 4;
let num_layers = 5;
let stack = MoDStack::new_uniform(num_layers, hidden_dim, 0.5, 30);
let hidden = make_hidden(seq_len, hidden_dim);
let fns: Vec<fn(&[f32], usize) -> Vec<f32>> = vec![identity_fn; num_layers];
let (_, stats) = stack.forward(&hidden, seq_len, &fns);
assert_eq!(stats.len(), num_layers);
}
// ── selection_fraction ───────────────────────────────────────────────────
#[test]
fn test_selection_fraction_near_capacity() {
// selection_fraction should be within one token of capacity_fraction.
let seq_len = 20;
let hidden_dim = 4;
let fraction = 0.5_f32;
let config = MoDConfig {
capacity_fraction: fraction,
hidden_dim,
aux_loss_weight: 0.01,
};
let layer = MoDLayer::new(config, 99);
let hidden = make_hidden(seq_len, hidden_dim);
let (_, stats) = layer.forward(&hidden, seq_len, identity_fn);
// Allow ±1 token tolerance due to floor.
let tolerance = 1.0 / seq_len as f32 + 1e-4;
assert!(
(stats.selection_fraction - fraction).abs() <= tolerance,
"selection_fraction {} far from capacity_fraction {}",
stats.selection_fraction,
fraction
);
}
}
@@ -68,6 +68,9 @@ pub mod ring_attention;
// Token Merging (ToMe) — bipartite soft matching for ViT-style speedup // Token Merging (ToMe) — bipartite soft matching for ViT-style speedup
pub mod token_merging; pub mod token_merging;
// Mixture of Depths (MoD) — dynamic token routing (Raposo et al. 2024)
pub mod mixture_of_depths;
// Cross-layer parameter sharing (ALBERT-style) // Cross-layer parameter sharing (ALBERT-style)
pub mod shared_layers; pub mod shared_layers;
// pub mod mamba_integration; // pub mod mamba_integration;
@@ -242,6 +245,9 @@ pub use sliding_window_attn::{
AttentionStats, SlidingWindowAttention, SlidingWindowConfig, WindowMask, AttentionStats, SlidingWindowAttention, SlidingWindowConfig, WindowMask,
}; };
pub use rope_scaling::{RopeScalingConfig, RopeScalingMode, RopeTable, RopeScaler}; pub use rope_scaling::{RopeScalingConfig, RopeScalingMode, RopeTable, RopeScaler};
pub use mixture_of_depths::{
MoDConfig, MoDLayer, MoDRouter, MoDStack, MoDStats, TokenSelection,
};
// TransformerConfig is defined above and available for import // TransformerConfig is defined above and available for import
@@ -81,6 +81,9 @@ pub use data_samplers::{
StratifiedStrategy, TemperatureSampler, StratifiedStrategy, TemperatureSampler,
}; };
pub mod model_merging;
pub use model_merging::{ModelMerger, TiesConfig, DareConfig, MergingStats};
/// Training state structure /// Training state structure
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct TrainingState { pub struct TrainingState {
@@ -0,0 +1,879 @@
//! Model merging algorithms: TIES and DARE
//!
//! Combines multiple fine-tuned models into a single model without additional training.
//!
//! # Algorithms
//!
//! - **TIES** (Yadav et al. 2023, arXiv:2306.01708): Trim, Elect Sign, Merge.
//! Trims small task-vector magnitudes, resolves sign conflicts via majority vote,
//! and averages only the parameters that agree with the elected sign.
//!
//! - **DARE** (Yu et al. 2023, arXiv:2311.03099): Drop And RE-scale.
//! Randomly sparsifies each task vector and rescales surviving elements to
//! preserve the expected value; can be combined with TIES sign election.
//!
//! # Example
//!
//! ```rust
//! use rtx_transformers::training::model_merging::{ModelMerger, TiesConfig, DareConfig};
//!
//! let base = vec![1.0_f32, 2.0, 3.0, 4.0];
//! let finetuned_a = vec![1.5_f32, 2.5, 2.5, 4.5];
//! let finetuned_b = vec![1.3_f32, 2.3, 3.3, 4.3];
//!
//! let (merged, stats) = ModelMerger::ties_merge(
//! &base,
//! &[finetuned_a.as_slice(), finetuned_b.as_slice()],
//! &TiesConfig::default(),
//! );
//!
//! assert_eq!(stats.num_models, 2);
//! assert_eq!(stats.num_params, 4);
//! ```
/// Configuration for TIES merging.
#[derive(Debug, Clone)]
pub struct TiesConfig {
/// Fraction of task-vector elements to retain (by magnitude).
/// `0.2` keeps the top 20 % — everything else is zeroed.
pub keep_ratio: f32,
/// Scalar applied to the merged task vector before adding it back to the base.
pub lambda: f32,
}
impl Default for TiesConfig {
fn default() -> Self {
Self {
keep_ratio: 0.2,
lambda: 1.0,
}
}
}
/// Configuration for DARE merging.
#[derive(Debug, Clone)]
pub struct DareConfig {
/// Fraction of task-vector elements to *drop* (set to zero).
/// `0.9` drops 90 % and scales survivors by 10×.
pub drop_rate: f32,
/// Scalar applied to the final merged task vector.
pub lambda: f32,
/// When `true`, apply TIES sign election after DARE sparsification.
pub use_ties: bool,
}
impl Default for DareConfig {
fn default() -> Self {
Self {
drop_rate: 0.9,
lambda: 1.0,
use_ties: false,
}
}
}
/// Diagnostics returned by the merge operations.
#[derive(Debug, Clone)]
pub struct MergingStats {
/// Number of fine-tuned models that were merged.
pub num_models: usize,
/// Total number of parameters in each model.
pub num_params: usize,
/// Fraction of the merged task vector that is nonzero.
pub nonzero_fraction: f32,
/// Fraction of parameter positions where a strict majority agreed on sign (TIES).
/// For DARE-only merges this is set to `0.0`.
pub sign_agreement_rate: f32,
}
/// Entry point for all model-merging algorithms.
pub struct ModelMerger;
impl ModelMerger {
// -----------------------------------------------------------------------
// Public primitives
// -----------------------------------------------------------------------
/// Compute the task vector τ = finetuned base (element-wise).
///
/// # Panics
///
/// Panics when `base` and `finetuned` differ in length.
pub fn task_vector(base: &[f32], finetuned: &[f32]) -> Vec<f32> {
assert_eq!(
base.len(),
finetuned.len(),
"base and finetuned must have the same length"
);
base.iter()
.zip(finetuned.iter())
.map(|(&b, &f)| f - b)
.collect()
}
/// Trim a task vector by zeroing elements whose absolute value falls below
/// the `keep_ratio` magnitude threshold.
///
/// `keep_ratio = 0.2` retains the top 20 % by magnitude; the other 80 %
/// become zero. The surviving values are left **unchanged** (no clipping).
///
/// # Edge cases
///
/// - `keep_ratio >= 1.0` — all elements are retained.
/// - `keep_ratio <= 0.0` — all elements are zeroed.
/// - Empty slice — returns an empty vector.
pub fn trim(task_vec: &[f32], keep_ratio: f32) -> Vec<f32> {
if task_vec.is_empty() {
return Vec::new();
}
let threshold = Self::percentile_threshold(task_vec, keep_ratio);
task_vec
.iter()
.map(|&v| if v.abs() >= threshold { v } else { 0.0 })
.collect()
}
/// Elect a sign (+1.0 or 1.0) for each parameter position via magnitude-
/// weighted majority vote across all task vectors.
///
/// For position `i`:
/// - Accumulate the absolute values of all positive contributions → `pos_sum`.
/// - Accumulate the absolute values of all negative contributions → `neg_sum`.
/// - Elected sign = +1.0 if `pos_sum >= neg_sum`, else 1.0.
///
/// This naturally breaks ties in favour of the larger-magnitude side and
/// reduces to a simple count majority when all magnitudes are equal.
///
/// # Panics
///
/// Panics when `task_vectors` is empty or when the inner slices differ in
/// length.
pub fn elect_sign(task_vectors: &[&[f32]]) -> Vec<f32> {
assert!(!task_vectors.is_empty(), "need at least one task vector");
let n = task_vectors[0].len();
for tv in task_vectors.iter().skip(1) {
assert_eq!(tv.len(), n, "all task vectors must have the same length");
}
(0..n)
.map(|i| {
let pos_sum: f32 = task_vectors
.iter()
.filter(|tv| tv[i] > 0.0)
.map(|tv| tv[i])
.sum();
let neg_sum: f32 = task_vectors
.iter()
.filter(|tv| tv[i] < 0.0)
.map(|tv| tv[i].abs())
.sum();
if pos_sum >= neg_sum {
1.0_f32
} else {
-1.0_f32
}
})
.collect()
}
/// Disjoint merge: for each parameter, average only the task-vector entries
/// whose sign matches the elected sign at that position.
///
/// Zero entries (pruned during trimming) are excluded from the average,
/// since `0.0.signum() == 0.0` never equals ±1.0. If no entry matches the
/// elected sign at a given position the output is `0.0`.
///
/// # Panics
///
/// Panics when `task_vectors` is empty or when lengths are inconsistent.
pub fn disjoint_merge(task_vectors: &[&[f32]], signs: &[f32]) -> Vec<f32> {
assert!(!task_vectors.is_empty(), "need at least one task vector");
let n = signs.len();
(0..n)
.map(|i| {
let matching: Vec<f32> = task_vectors
.iter()
.filter(|tv| tv[i] != 0.0 && tv[i].signum() == signs[i])
.map(|tv| tv[i])
.collect();
if matching.is_empty() {
0.0_f32
} else {
matching.iter().sum::<f32>() / matching.len() as f32
}
})
.collect()
}
// -----------------------------------------------------------------------
// Full TIES merge
// -----------------------------------------------------------------------
/// Perform a full TIES merge.
///
/// Steps:
/// 1. Compute task vectors τ_k = θ_k^ft θ_base.
/// 2. Trim each task vector (zero out low-magnitude entries).
/// 3. Elect a sign per position via magnitude-weighted majority vote.
/// 4. Disjoint-merge: average entries that agree with elected sign.
/// 5. Return θ_base + λ·τ_merged together with [`MergingStats`].
///
/// # Panics
///
/// Panics when `finetuned_models` is empty or when any model's length
/// differs from `base`.
pub fn ties_merge(
base: &[f32],
finetuned_models: &[&[f32]],
config: &TiesConfig,
) -> (Vec<f32>, MergingStats) {
assert!(
!finetuned_models.is_empty(),
"need at least one fine-tuned model"
);
let n = base.len();
let k = finetuned_models.len();
// Step 1: task vectors
let task_vecs: Vec<Vec<f32>> = finetuned_models
.iter()
.map(|ft| Self::task_vector(base, ft))
.collect();
// Step 2: trim
let trimmed: Vec<Vec<f32>> = task_vecs
.iter()
.map(|tv| Self::trim(tv, config.keep_ratio))
.collect();
// Step 3: elect sign — compute sign agreement rate while at it
let trimmed_refs: Vec<&[f32]> = trimmed.iter().map(Vec::as_slice).collect();
let signs = Self::elect_sign(&trimmed_refs);
// Measure strict majority agreement (count > k/2 on winning side)
let sign_agreement_rate = if n == 0 {
0.0
} else {
let agreed = (0..n)
.filter(|&i| {
let pos_count = trimmed_refs.iter().filter(|tv| tv[i] > 0.0).count();
let neg_count = trimmed_refs.iter().filter(|tv| tv[i] < 0.0).count();
pos_count != neg_count // strict majority exists
})
.count();
agreed as f32 / n as f32
};
// Step 4: disjoint merge
let tau_merged = Self::disjoint_merge(&trimmed_refs, &signs);
// Nonzero fraction of merged task vector
let nonzero_fraction = if n == 0 {
0.0
} else {
tau_merged.iter().filter(|&&v| v != 0.0).count() as f32 / n as f32
};
// Step 5: θ_base + λ·τ_merged
let merged: Vec<f32> = base
.iter()
.zip(tau_merged.iter())
.map(|(&b, &tau)| b + config.lambda * tau)
.collect();
let stats = MergingStats {
num_models: k,
num_params: n,
nonzero_fraction,
sign_agreement_rate,
};
(merged, stats)
}
// -----------------------------------------------------------------------
// DARE sparsification
// -----------------------------------------------------------------------
/// Randomly drop `drop_rate` fraction of task-vector elements and rescale
/// the survivors so that the expected value is preserved.
///
/// The rescaling factor is `1 / (1 drop_rate)`. A fast xorshift-style
/// LCG seeded by `seed` drives the Bernoulli draws, so results are fully
/// deterministic given the same seed and input.
///
/// # Panics
///
/// Panics when `drop_rate >= 1.0` (nothing would survive).
pub fn dare_sparsify(task_vec: &[f32], drop_rate: f32, seed: u64) -> Vec<f32> {
assert!(drop_rate < 1.0, "drop_rate must be < 1.0");
let keep_prob = (1.0 - drop_rate).max(1e-8_f32);
let mut state: u64 = seed;
task_vec
.iter()
.map(|&v| {
// LCG: Knuth's multiplicative constants (same as used in many
// well-known implementations of PCG/LCG PRNGs).
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
// Convert top 53 bits to a float in [0, 1).
let r = (state >> 11) as f32 / (1u64 << 53) as f32;
if r < keep_prob {
v / keep_prob // rescale to preserve E[sparse] = E[original]
} else {
0.0
}
})
.collect()
}
// -----------------------------------------------------------------------
// Full DARE merge
// -----------------------------------------------------------------------
/// Perform a full DARE merge.
///
/// Steps:
/// 1. Compute task vectors τ_k = θ_k^ft θ_base.
/// 2. Sparsify each task vector with [`Self::dare_sparsify`].
/// 3a. If `config.use_ties` is `true`: apply TIES sign election and
/// disjoint merge on the sparsified vectors.
/// 3b. Otherwise: plain average of sparsified task vectors.
/// 4. Return θ_base + λ·τ_merged together with [`MergingStats`].
///
/// `seeds` must contain one `u64` per entry in `finetuned_models`.
///
/// # Panics
///
/// Panics when `finetuned_models` is empty, lengths mismatch, or
/// `seeds.len() != finetuned_models.len()`.
pub fn dare_merge(
base: &[f32],
finetuned_models: &[&[f32]],
config: &DareConfig,
seeds: &[u64],
) -> (Vec<f32>, MergingStats) {
assert!(
!finetuned_models.is_empty(),
"need at least one fine-tuned model"
);
assert_eq!(
seeds.len(),
finetuned_models.len(),
"one seed required per fine-tuned model"
);
let n = base.len();
let k = finetuned_models.len();
// Step 1: task vectors
let task_vecs: Vec<Vec<f32>> = finetuned_models
.iter()
.map(|ft| Self::task_vector(base, ft))
.collect();
// Step 2: DARE sparsification
let sparsified: Vec<Vec<f32>> = task_vecs
.iter()
.zip(seeds.iter())
.map(|(tv, &seed)| Self::dare_sparsify(tv, config.drop_rate, seed))
.collect();
let sparse_refs: Vec<&[f32]> = sparsified.iter().map(Vec::as_slice).collect();
// Step 3: merge
let (tau_merged, sign_agreement_rate) = if config.use_ties {
let signs = Self::elect_sign(&sparse_refs);
let agreement = if n == 0 {
0.0
} else {
let agreed = (0..n)
.filter(|&i| {
let pos_count = sparse_refs.iter().filter(|tv| tv[i] > 0.0).count();
let neg_count = sparse_refs.iter().filter(|tv| tv[i] < 0.0).count();
pos_count != neg_count
})
.count();
agreed as f32 / n as f32
};
(Self::disjoint_merge(&sparse_refs, &signs), agreement)
} else {
// Plain average of sparsified task vectors
let tau: Vec<f32> = if n == 0 {
Vec::new()
} else {
(0..n)
.map(|i| {
sparse_refs.iter().map(|tv| tv[i]).sum::<f32>() / k as f32
})
.collect()
};
(tau, 0.0_f32)
};
// Nonzero fraction
let nonzero_fraction = if n == 0 {
0.0
} else {
tau_merged.iter().filter(|&&v| v != 0.0).count() as f32 / n as f32
};
// Step 4: θ_base + λ·τ_merged
let merged: Vec<f32> = base
.iter()
.zip(tau_merged.iter())
.map(|(&b, &tau)| b + config.lambda * tau)
.collect();
let stats = MergingStats {
num_models: k,
num_params: n,
nonzero_fraction,
sign_agreement_rate,
};
(merged, stats)
}
// -----------------------------------------------------------------------
// Simple linear merge
// -----------------------------------------------------------------------
/// Linearly interpolate fine-tuned models into the base model.
///
/// Computes `θ_base + Σ_k weights[k] * (θ_k^ft θ_base)`. With equal
/// weights that sum to 1.0 this reduces to a simple weighted average of
/// the fine-tuned parameters.
///
/// # Panics
///
/// Panics when `finetuned_models` is empty or `weights` has a different
/// length from `finetuned_models`.
pub fn linear_merge(base: &[f32], finetuned_models: &[&[f32]], weights: &[f32]) -> Vec<f32> {
assert!(
!finetuned_models.is_empty(),
"need at least one fine-tuned model"
);
assert_eq!(
weights.len(),
finetuned_models.len(),
"one weight per fine-tuned model required"
);
let n = base.len();
let mut merged = base.to_vec();
for (ft, &w) in finetuned_models.iter().zip(weights.iter()) {
assert_eq!(ft.len(), n, "model length must match base");
for (m, (&b, &f)) in merged.iter_mut().zip(base.iter().zip(ft.iter())) {
*m += w * (f - b);
}
}
merged
}
// -----------------------------------------------------------------------
// Private helpers
// -----------------------------------------------------------------------
/// Return the magnitude threshold such that `keep_ratio` fraction of
/// `values` have `|v| >= threshold`.
///
/// Implementation: sort absolute values ascending, then index at
/// `floor((1 keep_ratio) * len)`. Clamped so that keep_ratio ≥ 1.0
/// always returns 0.0 (keep everything) and keep_ratio ≤ 0.0 returns a
/// value that will zero every element.
fn percentile_threshold(values: &[f32], keep_ratio: f32) -> f32 {
if values.is_empty() {
return 0.0;
}
if keep_ratio >= 1.0 {
return 0.0; // keep all
}
if keep_ratio <= 0.0 {
// Return a value strictly greater than the max absolute value so
// that every element is zeroed.
let max_abs = values
.iter()
.map(|v| v.abs())
.fold(0.0_f32, f32::max);
return max_abs + 1.0;
}
let mut abs_vals: Vec<f32> = values.iter().map(|v| v.abs()).collect();
abs_vals.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let cutoff_idx = ((1.0 - keep_ratio) * abs_vals.len() as f32) as usize;
let idx = cutoff_idx.min(abs_vals.len().saturating_sub(1));
abs_vals[idx]
}
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
const EPS: f32 = 1e-5;
fn approx_eq(a: f32, b: f32) -> bool {
(a - b).abs() < EPS
}
// -----------------------------------------------------------------------
// task_vector
// -----------------------------------------------------------------------
#[test]
fn test_task_vector_basic() {
let base = vec![1.0_f32, 1.0];
let ft = vec![2.0_f32, 3.0];
let tv = ModelMerger::task_vector(&base, &ft);
assert!(approx_eq(tv[0], 1.0));
assert!(approx_eq(tv[1], 2.0));
}
#[test]
fn test_task_vector_negative() {
let base = vec![3.0_f32, 5.0];
let ft = vec![1.0_f32, 2.0];
let tv = ModelMerger::task_vector(&base, &ft);
assert!(approx_eq(tv[0], -2.0));
assert!(approx_eq(tv[1], -3.0));
}
// -----------------------------------------------------------------------
// trim
// -----------------------------------------------------------------------
#[test]
fn test_trim_keeps_top_fraction() {
// [1, 2, 3, 4]: keep_ratio=0.5 → keep top 2 by magnitude
let tv = vec![1.0_f32, 2.0, 3.0, 4.0];
let trimmed = ModelMerger::trim(&tv, 0.5);
let nonzero = trimmed.iter().filter(|&&v| v != 0.0).count();
// Threshold sits at the 50th-percentile absolute value, so exactly 2 should survive.
assert_eq!(nonzero, 2, "expected exactly 2 nonzero entries");
// The surviving values must be 3.0 and 4.0
assert!(approx_eq(trimmed[2], 3.0));
assert!(approx_eq(trimmed[3], 4.0));
}
#[test]
fn test_trim_keep_all() {
let tv = vec![0.1_f32, 0.5, 1.0, -2.0];
let trimmed = ModelMerger::trim(&tv, 1.0);
for (a, b) in tv.iter().zip(trimmed.iter()) {
assert!(approx_eq(*a, *b));
}
}
#[test]
fn test_trim_keep_none() {
let tv = vec![1.0_f32, 2.0, 3.0, 4.0];
let trimmed = ModelMerger::trim(&tv, 0.0);
assert!(trimmed.iter().all(|&v| v == 0.0));
}
#[test]
fn test_trim_preserves_values() {
// Surviving entries must be the original values, not clipped versions.
let tv = vec![10.0_f32, 0.01, 9.0, 0.02];
let trimmed = ModelMerger::trim(&tv, 0.5);
// 10.0 and 9.0 are the top 2; verify they are untouched.
assert!(approx_eq(trimmed[0], 10.0));
assert!(approx_eq(trimmed[2], 9.0));
assert!(approx_eq(trimmed[1], 0.0));
assert!(approx_eq(trimmed[3], 0.0));
}
// -----------------------------------------------------------------------
// elect_sign
// -----------------------------------------------------------------------
#[test]
fn test_elect_sign_majority_positive() {
// 2 positive, 1 negative → +1.0
let tv1 = vec![1.0_f32];
let tv2 = vec![2.0_f32];
let tv3 = vec![-0.5_f32];
let signs = ModelMerger::elect_sign(&[&tv1, &tv2, &tv3]);
assert!(approx_eq(signs[0], 1.0));
}
#[test]
fn test_elect_sign_majority_negative() {
// 2 negative, 1 positive → 1.0
let tv1 = vec![-1.0_f32];
let tv2 = vec![-2.0_f32];
let tv3 = vec![0.5_f32];
let signs = ModelMerger::elect_sign(&[&tv1, &tv2, &tv3]);
assert!(approx_eq(signs[0], -1.0));
}
#[test]
fn test_elect_sign_tie_breaks_by_magnitude() {
// Equal count (1 vs 1) but negative side has larger sum of magnitudes → 1.0
let tv1 = vec![0.1_f32]; // positive, small magnitude
let tv2 = vec![-5.0_f32]; // negative, large magnitude
let signs = ModelMerger::elect_sign(&[&tv1, &tv2]);
assert!(approx_eq(signs[0], -1.0));
}
#[test]
fn test_elect_sign_all_positive() {
let tv1 = vec![1.0_f32, 2.0, 3.0];
let tv2 = vec![4.0_f32, 5.0, 6.0];
let signs = ModelMerger::elect_sign(&[&tv1, &tv2]);
assert!(signs.iter().all(|&s| approx_eq(s, 1.0)));
}
// -----------------------------------------------------------------------
// disjoint_merge
// -----------------------------------------------------------------------
#[test]
fn test_disjoint_merge_agreement() {
// Both task vectors positive; elected sign +1 → average of both values.
let tv1 = vec![2.0_f32, 4.0];
let tv2 = vec![4.0_f32, 6.0];
let signs = vec![1.0_f32, 1.0];
let merged = ModelMerger::disjoint_merge(&[&tv1, &tv2], &signs);
assert!(approx_eq(merged[0], 3.0)); // (2+4)/2
assert!(approx_eq(merged[1], 5.0)); // (4+6)/2
}
#[test]
fn test_disjoint_merge_disagreement() {
// tv1 positive, tv2 negative; sign elected +1 → only tv1 contributes.
let tv1 = vec![3.0_f32];
let tv2 = vec![-2.0_f32];
let signs = vec![1.0_f32];
let merged = ModelMerger::disjoint_merge(&[&tv1, &tv2], &signs);
assert!(approx_eq(merged[0], 3.0));
}
#[test]
fn test_disjoint_merge_all_zero() {
let tv1 = vec![0.0_f32, 0.0];
let tv2 = vec![0.0_f32, 0.0];
let signs = vec![1.0_f32, -1.0];
let merged = ModelMerger::disjoint_merge(&[&tv1, &tv2], &signs);
assert!(merged.iter().all(|&v| approx_eq(v, 0.0)));
}
// -----------------------------------------------------------------------
// ties_merge
// -----------------------------------------------------------------------
#[test]
fn test_ties_merge_two_models() {
let base = vec![0.0_f32, 0.0, 0.0, 0.0];
let ft_a = vec![1.0_f32, -1.0, 2.0, -2.0];
let ft_b = vec![2.0_f32, -2.0, 1.0, -1.0];
let cfg = TiesConfig { keep_ratio: 1.0, lambda: 1.0 };
let (merged, _) = ModelMerger::ties_merge(&base, &[&ft_a, &ft_b], &cfg);
// Both task vectors agree in sign for every element; result should differ from base.
assert!(merged.iter().any(|&v| v != 0.0));
}
#[test]
fn test_ties_merge_stats_num_models() {
let base = vec![0.0_f32; 8];
let ft_a = vec![1.0_f32; 8];
let ft_b = vec![2.0_f32; 8];
let ft_c = vec![3.0_f32; 8];
let cfg = TiesConfig::default();
let (_, stats) = ModelMerger::ties_merge(&base, &[&ft_a, &ft_b, &ft_c], &cfg);
assert_eq!(stats.num_models, 3);
assert_eq!(stats.num_params, 8);
}
#[test]
fn test_ties_merge_single_model() {
// With a single fine-tuned model and keep_ratio=1.0, the merge is
// base + lambda * (finetuned - base). With lambda=1.0 this equals finetuned.
let base = vec![1.0_f32, 2.0, 3.0, 4.0];
let ft = vec![2.0_f32, 4.0, 6.0, 8.0];
let cfg = TiesConfig { keep_ratio: 1.0, lambda: 1.0 };
let (merged, stats) = ModelMerger::ties_merge(&base, &[&ft], &cfg);
assert_eq!(stats.num_models, 1);
for (m, f) in merged.iter().zip(ft.iter()) {
assert!(approx_eq(*m, *f), "expected merged[i] == ft[i], got {m} vs {f}");
}
}
// -----------------------------------------------------------------------
// dare_sparsify
// -----------------------------------------------------------------------
#[test]
fn test_dare_sparsify_fraction() {
// With drop_rate=0.5 we expect roughly 50 % nonzero.
let tv: Vec<f32> = (0..10_000).map(|i| i as f32).collect();
let sparse = ModelMerger::dare_sparsify(&tv, 0.5, 42);
let nonzero = sparse.iter().filter(|&&v| v != 0.0).count();
let fraction = nonzero as f32 / tv.len() as f32;
// Allow ±5 % tolerance around the expected 50 %.
assert!(
(fraction - 0.5).abs() < 0.05,
"nonzero fraction {fraction} not within 5 % of 0.5"
);
}
#[test]
fn test_dare_sparsify_rescales() {
// E[sparse] should approximate E[original] due to rescaling.
let tv: Vec<f32> = vec![1.0_f32; 100_000];
let sparse = ModelMerger::dare_sparsify(&tv, 0.7, 7);
let mean_sparse: f32 = sparse.iter().sum::<f32>() / sparse.len() as f32;
// Expected value = 1.0; allow ±3 % relative error.
assert!(
(mean_sparse - 1.0).abs() < 0.03,
"rescaled mean {mean_sparse} deviates from 1.0 by more than 3 %"
);
}
#[test]
fn test_dare_sparsify_seed_determinism() {
let tv: Vec<f32> = (0..256).map(|i| i as f32 * 0.1).collect();
let s1 = ModelMerger::dare_sparsify(&tv, 0.6, 99);
let s2 = ModelMerger::dare_sparsify(&tv, 0.6, 99);
assert_eq!(s1, s2, "same seed must produce identical output");
}
// -----------------------------------------------------------------------
// dare_merge
// -----------------------------------------------------------------------
#[test]
fn test_dare_merge_with_ties() {
// Verify the use_ties path executes without panicking and produces
// a plausible result (merged differs from base when task vectors are nonzero).
let base = vec![0.0_f32; 16];
let ft_a: Vec<f32> = (0..16).map(|i| i as f32).collect();
let ft_b: Vec<f32> = (0..16).map(|i| -(i as f32)).collect();
let cfg = DareConfig { drop_rate: 0.5, lambda: 1.0, use_ties: true };
let (merged, stats) = ModelMerger::dare_merge(&base, &[&ft_a, &ft_b], &cfg, &[1, 2]);
assert_eq!(stats.num_models, 2);
assert_eq!(stats.num_params, 16);
// Result is finite and not all NaN
assert!(merged.iter().all(|v| v.is_finite()));
}
// -----------------------------------------------------------------------
// linear_merge
// -----------------------------------------------------------------------
#[test]
fn test_linear_merge_weights() {
// Equal weights summing to 1.0 → weighted average of fine-tuned params.
let base = vec![0.0_f32, 0.0, 0.0];
let ft_a = vec![2.0_f32, 4.0, 6.0];
let ft_b = vec![4.0_f32, 6.0, 8.0];
let merged = ModelMerger::linear_merge(&base, &[&ft_a, &ft_b], &[0.5, 0.5]);
assert!(approx_eq(merged[0], 3.0));
assert!(approx_eq(merged[1], 5.0));
assert!(approx_eq(merged[2], 7.0));
}
#[test]
fn test_linear_merge_single() {
// One model with weight=1.0 and base=0 → result equals finetuned.
let base = vec![0.0_f32, 0.0, 0.0];
let ft = vec![3.0_f32, 6.0, 9.0];
let merged = ModelMerger::linear_merge(&base, &[&ft], &[1.0]);
for (m, f) in merged.iter().zip(ft.iter()) {
assert!(approx_eq(*m, *f));
}
}
// -----------------------------------------------------------------------
// Additional edge-case and integration tests
// -----------------------------------------------------------------------
#[test]
fn test_task_vector_zero_diff() {
// Identical base and finetuned → all-zero task vector
let base = vec![5.0_f32, 3.0, 1.0];
let tv = ModelMerger::task_vector(&base, &base);
assert!(tv.iter().all(|&v| approx_eq(v, 0.0)));
}
#[test]
fn test_trim_single_element() {
let tv = vec![7.0_f32];
let trimmed = ModelMerger::trim(&tv, 0.5);
assert!(approx_eq(trimmed[0], 7.0));
}
#[test]
fn test_ties_merge_stats_nonzero_fraction_bounded() {
let base = vec![0.0_f32; 100];
let ft_a: Vec<f32> = (0..100).map(|i| i as f32).collect();
let ft_b: Vec<f32> = (0..100).map(|i| i as f32 * 0.5).collect();
let cfg = TiesConfig { keep_ratio: 0.2, lambda: 1.0 };
let (_, stats) = ModelMerger::ties_merge(&base, &[&ft_a, &ft_b], &cfg);
assert!(stats.nonzero_fraction >= 0.0 && stats.nonzero_fraction <= 1.0);
assert!(stats.sign_agreement_rate >= 0.0 && stats.sign_agreement_rate <= 1.0);
}
#[test]
fn test_dare_merge_no_ties() {
// Plain averaging path (use_ties=false) should be deterministic given seeds.
let base = vec![1.0_f32; 8];
let ft_a = vec![2.0_f32; 8];
let ft_b = vec![3.0_f32; 8];
let cfg = DareConfig { drop_rate: 0.0, lambda: 1.0, use_ties: false };
// drop_rate=0 → every element kept, rescale = 1.0; average of task vecs = 1.5
let (merged, _) = ModelMerger::dare_merge(&base, &[&ft_a, &ft_b], &cfg, &[0, 1]);
// Each ft task vector: ft_a-base=[1,1,...], ft_b-base=[2,2,...]
// Average = 1.5; merged = base + 1.0 * 1.5 = 2.5
for &v in &merged {
assert!(approx_eq(v, 2.5), "expected 2.5, got {v}");
}
}
#[test]
fn test_elect_sign_single_vector() {
let tv = vec![3.0_f32, -1.0, 0.0, -5.0];
let signs = ModelMerger::elect_sign(&[&tv]);
assert!(approx_eq(signs[0], 1.0));
assert!(approx_eq(signs[1], -1.0));
// 0.0 has pos_sum=0, neg_sum=0 → pos_sum >= neg_sum → +1.0
assert!(approx_eq(signs[2], 1.0));
assert!(approx_eq(signs[3], -1.0));
}
#[test]
fn test_linear_merge_base_plus_zero_weight() {
// weight=0 for a fine-tuned model → result equals base
let base = vec![1.0_f32, 2.0, 3.0];
let ft = vec![100.0_f32, 200.0, 300.0];
let merged = ModelMerger::linear_merge(&base, &[&ft], &[0.0]);
for (m, b) in merged.iter().zip(base.iter()) {
assert!(approx_eq(*m, *b));
}
}
#[test]
fn test_ties_merge_lambda_scales_output() {
// With keep_ratio=1.0 and both models equal, lambda doubles the task vector.
let base = vec![0.0_f32; 4];
let ft = vec![1.0_f32; 4];
let cfg1 = TiesConfig { keep_ratio: 1.0, lambda: 1.0 };
let cfg2 = TiesConfig { keep_ratio: 1.0, lambda: 2.0 };
let (m1, _) = ModelMerger::ties_merge(&base, &[&ft], &cfg1);
let (m2, _) = ModelMerger::ties_merge(&base, &[&ft], &cfg2);
for (v1, v2) in m1.iter().zip(m2.iter()) {
assert!(approx_eq(*v2, *v1 * 2.0), "lambda=2 should double the result");
}
}
}