use std::fmt;
use super::request::Token;
#[derive(Debug, Clone)]
pub struct ExponentialMovingAverage {
value: f64,
alpha: f64,
count: u64,
}
impl ExponentialMovingAverage {
pub fn new(alpha: f64) -> Self {
Self {
value: 0.0,
alpha: alpha.clamp(0.0, 1.0),
count: 0,
}
}
pub fn update(&mut self, sample: f64) {
if self.count == 0 {
self.value = sample;
} else {
self.value = self.alpha * sample + (1.0 - self.alpha) * self.value;
}
self.count += 1;
}
pub fn value(&self) -> f64 {
self.value
}
pub fn reset(&mut self) {
self.value = 0.0;
self.count = 0;
}
}
impl Default for ExponentialMovingAverage {
fn default() -> Self {
Self::new(0.1)
}
}
#[derive(Debug, Clone)]
pub struct SpeculativeOutput {
pub accepted: Vec<Token>,
pub rejection_idx: Option<usize>,
pub target_token: Token,
pub draft_count: usize,
}
impl SpeculativeOutput {
pub fn acceptance_rate(&self) -> f64 {
if self.draft_count == 0 {
return 0.0;
}
self.accepted.len() as f64 / self.draft_count as f64
}
pub fn total_tokens(&self) -> usize {
self.accepted.len() + 1
}
}
#[derive(Debug)]
pub struct SpeculativeDecoder {
k: usize,
acceptance_rate: ExponentialMovingAverage,
total_steps: u64,
total_accepted: u64,
total_draft: u64,
}
impl SpeculativeDecoder {
pub fn new(k: usize) -> Self {
Self {
k,
acceptance_rate: ExponentialMovingAverage::new(0.1),
total_steps: 0,
total_accepted: 0,
total_draft: 0,
}
}
pub fn k(&self) -> usize {
self.k
}
pub fn set_k(&mut self, k: usize) {
self.k = k;
}
pub fn simulate_step(
&mut self,
draft_tokens: &[Token],
target_probs: &[(Token, f64)],
) -> SpeculativeOutput {
let draft_count = draft_tokens.len().min(self.k);
let mut accepted = Vec::new();
let mut rejection_idx = None;
for (i, &draft_token) in draft_tokens.iter().take(draft_count).enumerate() {
if let Some((target_token, _)) = target_probs.get(i) {
if *target_token == draft_token {
accepted.push(draft_token);
} else {
rejection_idx = Some(i);
break;
}
} else {
rejection_idx = Some(i);
break;
}
}
let target_token = if let Some(idx) = rejection_idx {
target_probs.get(idx).map(|(t, _)| *t).unwrap_or(0)
} else {
target_probs.get(draft_count).map(|(t, _)| *t).unwrap_or(0)
};
let output = SpeculativeOutput {
accepted: accepted.clone(),
rejection_idx,
target_token,
draft_count,
};
self.total_steps += 1;
self.total_accepted += accepted.len() as u64;
self.total_draft += draft_count as u64;
self.acceptance_rate.update(output.acceptance_rate());
output
}
pub fn acceptance_rate(&self) -> f64 {
self.acceptance_rate.value()
}
pub fn overall_acceptance_rate(&self) -> f64 {
if self.total_draft == 0 {
return 0.0;
}
self.total_accepted as f64 / self.total_draft as f64
}
pub fn speedup(&self) -> f64 {
let rate = self.acceptance_rate();
1.0 + (self.k as f64) * rate
}
pub fn stats(&self) -> (u64, u64, u64) {
(self.total_steps, self.total_accepted, self.total_draft)
}
}
impl fmt::Display for SpeculativeDecoder {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"SpeculativeDecoder(k={}, acceptance={:.1}%, speedup={:.2}x)",
self.k,
self.acceptance_rate() * 100.0,
self.speedup()
)
}
}