use crate::error::{InferenceError, InferenceResult};
use crate::sampling::{Sampler, SamplingConfig};
use kizzasi_model::AutoregressiveModel;
use scirs2_core::ndarray::Array1;
#[derive(Debug, Clone)]
pub struct SpeculativeConfig {
pub num_draft_tokens: usize,
pub draft_temperature: f32,
pub main_temperature: f32,
pub greedy_verification: bool,
}
impl Default for SpeculativeConfig {
fn default() -> Self {
Self {
num_draft_tokens: 4,
draft_temperature: 1.0,
main_temperature: 1.0,
greedy_verification: true,
}
}
}
impl SpeculativeConfig {
pub fn new() -> Self {
Self::default()
}
pub fn num_draft_tokens(mut self, n: usize) -> Self {
self.num_draft_tokens = n;
self
}
pub fn draft_temperature(mut self, temp: f32) -> Self {
self.draft_temperature = temp;
self
}
pub fn main_temperature(mut self, temp: f32) -> Self {
self.main_temperature = temp;
self
}
pub fn greedy_verification(mut self, greedy: bool) -> Self {
self.greedy_verification = greedy;
self
}
}
pub struct SpeculativeDecoder {
main_model: Box<dyn AutoregressiveModel>,
draft_model: Box<dyn AutoregressiveModel>,
config: SpeculativeConfig,
draft_sampler: Sampler,
main_sampler: Sampler,
total_tokens: usize,
accepted_tokens: usize,
}
impl SpeculativeDecoder {
pub fn new(
main_model: Box<dyn AutoregressiveModel>,
draft_model: Box<dyn AutoregressiveModel>,
config: SpeculativeConfig,
) -> Self {
let draft_sampler = Sampler::new(
SamplingConfig::new()
.temperature(config.draft_temperature)
.strategy(crate::sampling::SamplingStrategy::Temperature),
);
let main_sampler = Sampler::new(
SamplingConfig::new()
.temperature(config.main_temperature)
.strategy(if config.greedy_verification {
crate::sampling::SamplingStrategy::Greedy
} else {
crate::sampling::SamplingStrategy::Temperature
}),
);
Self {
main_model,
draft_model,
config,
draft_sampler,
main_sampler,
total_tokens: 0,
accepted_tokens: 0,
}
}
pub fn generate(
&mut self,
input: &Array1<f32>,
max_tokens: usize,
) -> InferenceResult<Vec<Array1<f32>>> {
let mut sequence = Vec::with_capacity(max_tokens);
let mut current = input.clone();
while sequence.len() < max_tokens {
let draft_candidates = self.generate_draft_tokens(¤t)?;
let (accepted, next_token) = self.verify_candidates(¤t, &draft_candidates)?;
for token in draft_candidates.iter().take(accepted) {
sequence.push(token.clone());
if sequence.len() >= max_tokens {
break;
}
}
self.total_tokens += self.config.num_draft_tokens;
self.accepted_tokens += accepted;
if sequence.len() < max_tokens {
sequence.push(next_token.clone());
current = next_token;
}
}
Ok(sequence)
}
fn generate_draft_tokens(&mut self, input: &Array1<f32>) -> InferenceResult<Vec<Array1<f32>>> {
let mut candidates = Vec::with_capacity(self.config.num_draft_tokens);
let mut current = input.clone();
for _ in 0..self.config.num_draft_tokens {
let logits = self
.draft_model
.step(¤t)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let sampled = self.draft_sampler.sample(&logits)?;
let token = Array1::from_elem(1, sampled);
candidates.push(token.clone());
current = token;
}
Ok(candidates)
}
fn verify_candidates(
&mut self,
input: &Array1<f32>,
candidates: &[Array1<f32>],
) -> InferenceResult<(usize, Array1<f32>)> {
let mut current = input.clone();
let mut accepted = 0;
for candidate in candidates {
let main_logits = self
.main_model
.step(¤t)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let main_prediction = self.main_sampler.sample(&main_logits)?;
let candidate_value = candidate[0];
let matches = if self.config.greedy_verification {
(candidate_value - main_prediction).abs() < 1e-6
} else {
(candidate_value - main_prediction).abs() < 0.5
};
if matches {
accepted += 1;
current = candidate.clone();
} else {
let next_token = Array1::from_elem(1, main_prediction);
return Ok((accepted, next_token));
}
}
let main_logits = self
.main_model
.step(¤t)
.map_err(|e| InferenceError::ForwardError(e.to_string()))?;
let main_prediction = self.main_sampler.sample(&main_logits)?;
let next_token = Array1::from_elem(1, main_prediction);
Ok((accepted, next_token))
}
pub fn acceptance_rate(&self) -> f32 {
if self.total_tokens == 0 {
0.0
} else {
self.accepted_tokens as f32 / self.total_tokens as f32
}
}
pub fn reset_stats(&mut self) {
self.total_tokens = 0;
self.accepted_tokens = 0;
}
pub fn config(&self) -> &SpeculativeConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
use kizzasi_model::s4::{S4Config, S4D};
#[test]
fn test_speculative_config() {
let config = SpeculativeConfig::new()
.num_draft_tokens(5)
.draft_temperature(0.8)
.main_temperature(1.2)
.greedy_verification(false);
assert_eq!(config.num_draft_tokens, 5);
assert!((config.draft_temperature - 0.8).abs() < 1e-6);
assert!((config.main_temperature - 1.2).abs() < 1e-6);
assert!(!config.greedy_verification);
}
#[test]
fn test_speculative_decoder_creation() {
let draft_config = S4Config::new()
.input_dim(1)
.hidden_dim(32)
.state_dim(8)
.num_layers(1)
.diagonal(true);
let draft_model = S4D::new(draft_config).unwrap();
let main_config = S4Config::new()
.input_dim(1)
.hidden_dim(64)
.state_dim(16)
.num_layers(2)
.diagonal(true);
let main_model = S4D::new(main_config).unwrap();
let config = SpeculativeConfig::new().num_draft_tokens(3);
let decoder = SpeculativeDecoder::new(Box::new(main_model), Box::new(draft_model), config);
assert_eq!(decoder.config().num_draft_tokens, 3);
assert_eq!(decoder.acceptance_rate(), 0.0);
}
#[test]
fn test_speculative_generation() {
let draft_config = S4Config::new()
.input_dim(1)
.hidden_dim(32)
.state_dim(8)
.num_layers(1)
.diagonal(true);
let draft_model = S4D::new(draft_config).unwrap();
let main_config = S4Config::new()
.input_dim(1)
.hidden_dim(64)
.state_dim(16)
.num_layers(2)
.diagonal(true);
let main_model = S4D::new(main_config).unwrap();
let config = SpeculativeConfig::new().num_draft_tokens(2);
let mut decoder =
SpeculativeDecoder::new(Box::new(main_model), Box::new(draft_model), config);
let input = Array1::from_vec(vec![0.5]);
let result = decoder.generate(&input, 10);
assert!(result.is_ok());
let sequence = result.unwrap();
assert_eq!(sequence.len(), 10);
let acc_rate = decoder.acceptance_rate();
assert!((0.0..=1.0).contains(&acc_rate));
}
}