use rand::Rng;
pub enum SampleStep {
Accept,
Reject { replacement_token: u32 },
}
pub fn softmax_with_temp(logprobs: &[f32], temp: f32) -> Vec<f32> {
assert!(temp > 0.0, "softmax_with_temp: temp must be > 0");
let scale = 1.0 / temp;
let max_scaled = logprobs
.iter()
.map(|&x| x * scale)
.fold(f32::NEG_INFINITY, f32::max);
let mut probs: Vec<f32> = logprobs
.iter()
.map(|&x| (x * scale - max_scaled).exp())
.collect();
let z: f32 = probs.iter().sum();
let inv_z = 1.0 / z;
for p in probs.iter_mut() {
*p *= inv_z;
}
probs
}
pub fn leviathan_step(
draft_token: u32,
target_probs: &[f32],
drafter_probs: &[f32],
rng: &mut impl Rng,
) -> SampleStep {
let vocab = target_probs.len();
assert_eq!(
drafter_probs.len(),
vocab,
"leviathan_step: target/drafter probs must have same vocab"
);
let v = draft_token as usize;
assert!(
v < vocab,
"leviathan_step: draft_token {} out of vocab {}",
v,
vocab
);
let p = target_probs[v];
let q = drafter_probs[v];
let accept_prob = if q <= 0.0 { 0.0 } else { (p / q).min(1.0) };
let u: f32 = rng.gen();
if u < accept_prob {
return SampleStep::Accept;
}
let mut residual: Vec<f32> = target_probs
.iter()
.zip(drafter_probs.iter())
.map(|(&pv, &qv)| (pv - qv).max(0.0))
.collect();
let z: f32 = residual.iter().sum();
if z <= 0.0 {
let argmax_t = target_probs
.iter()
.enumerate()
.fold((0usize, f32::NEG_INFINITY), |(i_max, v_max), (i, &v)| {
if v > v_max {
(i, v)
} else {
(i_max, v_max)
}
})
.0;
return SampleStep::Reject {
replacement_token: argmax_t as u32,
};
}
let inv_z = 1.0 / z;
for r in residual.iter_mut() {
*r *= inv_z;
}
let u_resample: f32 = rng.gen();
let mut acc = 0.0f32;
for (i, &r) in residual.iter().enumerate() {
acc += r;
if u_resample < acc {
return SampleStep::Reject {
replacement_token: i as u32,
};
}
}
let last = residual.len() - 1;
SampleStep::Reject {
replacement_token: last as u32,
}
}
pub fn leviathan_accept_prefix(
drafts: &[u32],
target_probs_per_pos: &[Vec<f32>],
drafter_probs_per_pos: &[Vec<f32>],
rng: &mut impl Rng,
) -> (usize, u32) {
assert_eq!(
drafter_probs_per_pos.len(),
drafts.len(),
"leviathan_accept_prefix: drafter_probs must have len = drafts.len()"
);
assert_eq!(
target_probs_per_pos.len(),
drafts.len() + 1,
"leviathan_accept_prefix: target_probs must have len = drafts.len() + 1"
);
for (i, &draft) in drafts.iter().enumerate() {
let step = leviathan_step(
draft,
&target_probs_per_pos[i],
&drafter_probs_per_pos[i],
rng,
);
if let SampleStep::Reject { replacement_token } = step {
return (i, replacement_token);
}
}
let target_last = &target_probs_per_pos[drafts.len()];
let u: f32 = rng.gen();
let mut acc = 0.0f32;
for (i, &p) in target_last.iter().enumerate() {
acc += p;
if u < acc {
return (drafts.len(), i as u32);
}
}
let last = target_last.len() - 1;
(drafts.len(), last as u32)
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::SeedableRng;
#[test]
fn softmax_with_temp_normalizes() {
let logprobs = vec![1.0, 2.0, 3.0, 4.0];
let probs = softmax_with_temp(&logprobs, 1.0);
let sum: f32 = probs.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "softmax must sum to 1; got {sum}");
assert!(probs[3] > probs[0]);
}
#[test]
fn softmax_temp_zero_panics_via_assert() {
let result = std::panic::catch_unwind(|| softmax_with_temp(&[1.0, 2.0, 3.0], 0.0));
assert!(result.is_err(), "temp=0 should panic via assert");
}
#[test]
fn leviathan_step_accepts_when_target_dominates() {
let target_probs = vec![0.05, 0.9, 0.05];
let drafter_probs = vec![0.45, 0.1, 0.45];
let mut rng = StdRng::seed_from_u64(42);
for _ in 0..100 {
match leviathan_step(1, &target_probs, &drafter_probs, &mut rng) {
SampleStep::Accept => {}
SampleStep::Reject { .. } => panic!("should always accept when p >> q"),
}
}
}
#[test]
fn leviathan_step_rejects_when_drafter_dominates() {
let target_probs = vec![0.475, 0.05, 0.475];
let drafter_probs = vec![0.025, 0.95, 0.025];
let mut rng = StdRng::seed_from_u64(42);
let mut rejects = 0;
for _ in 0..1000 {
match leviathan_step(1, &target_probs, &drafter_probs, &mut rng) {
SampleStep::Accept => {}
SampleStep::Reject { .. } => rejects += 1,
}
}
assert!(
(800..=990).contains(&rejects),
"expected ~950 rejects, got {rejects}"
);
}
#[test]
fn leviathan_step_residual_replacement_is_correct_token() {
let target_probs = vec![0.4, 0.1, 0.5];
let drafter_probs = vec![0.1, 0.8, 0.1];
let mut rng = StdRng::seed_from_u64(42);
let mut sample_counts = [0u32; 3];
let mut rejects = 0;
for _ in 0..2000 {
match leviathan_step(1, &target_probs, &drafter_probs, &mut rng) {
SampleStep::Accept => {}
SampleStep::Reject { replacement_token } => {
rejects += 1;
sample_counts[replacement_token as usize] += 1;
}
}
}
assert_eq!(sample_counts[1], 0, "residual at token 1 was 0");
let total_resampled = sample_counts[0] + sample_counts[2];
assert!(total_resampled > 100, "should have many rejects: {rejects}");
let p0 = sample_counts[0] as f32 / total_resampled as f32;
assert!(
(0.35..=0.50).contains(&p0),
"expected token 0 fraction ≈ 0.43, got {p0}"
);
}
#[test]
fn leviathan_accept_prefix_full_accept_returns_continuation() {
let drafts = vec![1u32, 2u32];
let target_probs = vec![
vec![0.1, 0.8, 0.1],
vec![0.1, 0.1, 0.8],
vec![0.5, 0.3, 0.2],
];
let drafter_probs = vec![vec![0.45, 0.1, 0.45], vec![0.45, 0.45, 0.1]];
let mut rng = StdRng::seed_from_u64(42);
let (accept_count, continuation) =
leviathan_accept_prefix(&drafts, &target_probs, &drafter_probs, &mut rng);
assert_eq!(accept_count, 2, "all drafts accepted");
assert!(continuation < 3, "continuation in vocab");
}
#[test]
fn leviathan_accept_prefix_partial_reject_truncates() {
let drafts = vec![1u32, 1u32];
let target_probs = vec![
vec![0.1, 0.8, 0.1],
vec![0.475, 0.05, 0.475],
vec![0.5, 0.3, 0.2],
];
let drafter_probs = vec![vec![0.1, 0.8, 0.1], vec![0.025, 0.95, 0.025]];
let mut saw_partial = false;
for seed in 0..100 {
let mut rng = StdRng::seed_from_u64(seed);
let (accept_count, _) =
leviathan_accept_prefix(&drafts, &target_probs, &drafter_probs, &mut rng);
if accept_count == 1 {
saw_partial = true;
break;
}
}
assert!(saw_partial, "should have seen partial-accept across seeds");
}
}