use crate::generate::{argmax, multinomial, softmax};
use crate::runtime::serve::Runtime;
pub trait DraftModel: Send + Sync {
fn draft(&self, context: &[u32], gamma: usize) -> Vec<u32>;
}
pub struct MlpDraftModel {
vocab_size: usize,
hidden: usize,
draft_hidden: usize,
embed: Vec<f32>,
w1: Vec<f32>,
b1: Vec<f32>,
w2: Vec<f32>,
b2: Vec<f32>,
temperature: f32,
top_p: f32,
}
impl MlpDraftModel {
pub fn from_runtime(
runtime: &Runtime,
vocab_size: usize,
hidden: usize,
draft_hidden: usize,
temperature: f32,
top_p: f32,
) -> Option<Self> {
let embed = runtime
.get("draft.embed.weight")
.map(|t| t.data.clone())
.or_else(|| {
runtime
.get("transformer.wte.weight")
.or_else(|| runtime.get("model.embed_tokens.weight"))
.map(|t| t.data.clone())
})?;
let w1 = runtime.get("draft.fc1.weight").map(|t| t.data.clone())?;
let b1 = runtime.get("draft.fc1.bias").map(|t| t.data.clone())?;
let w2 = runtime.get("draft.fc2.weight").map(|t| t.data.clone())?;
let b2 = runtime.get("draft.fc2.bias").map(|t| t.data.clone())?;
Some(Self {
vocab_size,
hidden,
draft_hidden,
embed,
w1,
b1,
w2,
b2,
temperature,
top_p,
})
}
#[allow(clippy::needless_range_loop)]
fn forward(&self, token_id: u32) -> Vec<f32> {
let idx = (token_id as usize) * self.hidden;
let embed = if idx + self.hidden <= self.embed.len() {
&self.embed[idx..idx + self.hidden]
} else {
return vec![0.0f32; self.vocab_size];
};
let mut hidden = vec![0.0f32; self.draft_hidden];
for i in 0..self.draft_hidden {
let mut sum = self.b1[i];
let row_start = i * self.hidden;
for j in 0..self.hidden {
sum += self.w1[row_start + j] * embed[j];
}
hidden[i] = sum.max(0.0); }
let mut logits = vec![0.0f32; self.vocab_size];
for i in 0..self.vocab_size {
let mut sum = self.b2[i];
let row_start = i * self.draft_hidden;
for j in 0..self.draft_hidden {
sum += self.w2[row_start + j] * hidden[j];
}
logits[i] = sum;
}
logits
}
}
impl DraftModel for MlpDraftModel {
fn draft(&self, context: &[u32], gamma: usize) -> Vec<u32> {
let mut draft = Vec::with_capacity(gamma);
let mut current = context.last().copied().unwrap_or(0);
for _ in 0..gamma {
let logits = self.forward(current);
let next = if self.temperature <= 0.0 {
argmax(&logits) as u32
} else {
let scaled: Vec<f32> = logits.iter().map(|&l| l / self.temperature).collect();
let mut probs = softmax(&scaled);
if self.top_p > 0.0 && self.top_p < 1.0 {
crate::generate::apply_top_p(&mut probs, self.top_p);
}
multinomial(&probs, 0, None)
};
draft.push(next);
current = next;
}
draft
}
}
pub struct PromptLookupDraftModel {
num_tokens: usize,
}
impl PromptLookupDraftModel {
pub fn new(num_tokens: usize) -> Self {
Self {
num_tokens: num_tokens.max(1),
}
}
}
impl DraftModel for PromptLookupDraftModel {
fn draft(&self, context: &[u32], gamma: usize) -> Vec<u32> {
if context.len() < self.num_tokens + 1 || gamma == 0 {
return Vec::new();
}
let needle = &context[context.len() - self.num_tokens..];
let search_end = context.len() - self.num_tokens;
for start in (0..search_end).rev() {
if context[start..start + self.num_tokens] == *needle {
let remaining = search_end - (start + self.num_tokens);
let take = gamma.min(remaining);
return context[start + self.num_tokens..start + self.num_tokens + take].to_vec();
}
}
Vec::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mlp_draft_model_from_runtime_requires_all_weights() {
let runtime = Runtime::from_raw(&std::collections::HashMap::new());
let model = MlpDraftModel::from_runtime(&runtime, 10, 4, 2, 0.0, 0.0);
assert!(model.is_none(), "empty runtime should not yield draft model");
}
#[test]
fn mlp_draft_model_produces_deterministic_tokens_when_greedy() {
fn f32s(xs: &[f32]) -> Vec<u8> {
xs.iter().flat_map(|f| f.to_le_bytes()).collect()
}
let mut tensors = std::collections::HashMap::new();
tensors.insert(
"transformer.wte.weight".to_string(),
crate::model::TensorData {
data: f32s(&[1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 0.0, 0.0]),
shape: vec![4, 2],
dtype: crate::model::DataType::F32,
},
);
tensors.insert(
"draft.fc1.weight".to_string(),
crate::model::TensorData {
data: f32s(&[1.0, 0.0, 0.0, 1.0]),
shape: vec![2, 2],
dtype: crate::model::DataType::F32,
},
);
tensors.insert(
"draft.fc1.bias".to_string(),
crate::model::TensorData {
data: f32s(&[0.0, 0.0]),
shape: vec![2],
dtype: crate::model::DataType::F32,
},
);
tensors.insert(
"draft.fc2.weight".to_string(),
crate::model::TensorData {
data: f32s(&[0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0]),
shape: vec![4, 2],
dtype: crate::model::DataType::F32,
},
);
tensors.insert(
"draft.fc2.bias".to_string(),
crate::model::TensorData {
data: f32s(&[0.0, 0.0, 0.0, 0.0]),
shape: vec![4],
dtype: crate::model::DataType::F32,
},
);
let runtime = Runtime::from_raw(&tensors);
let model = MlpDraftModel::from_runtime(&runtime, 4, 2, 2, 0.0, 0.0)
.expect("should construct with all weights present");
let draft = model.draft(&[0], 3);
assert_eq!(draft.len(), 3, "should produce exactly gamma tokens");
assert_eq!(draft, vec![2, 2, 2]);
}
#[test]
fn prompt_lookup_finds_continuation() {
let model = PromptLookupDraftModel::new(3);
let draft = model.draft(&[1, 2, 3, 4, 1, 2, 3], 4);
assert_eq!(draft, vec![4]);
}
#[test]
fn prompt_lookup_returns_empty_when_no_match() {
let model = PromptLookupDraftModel::new(3);
let draft = model.draft(&[1, 2, 3, 4, 5, 6, 7], 4);
assert!(draft.is_empty(), "no match should yield empty draft");
}
#[test]
fn prompt_lookup_returns_empty_for_short_context() {
let model = PromptLookupDraftModel::new(3);
let draft = model.draft(&[1, 2], 4);
assert!(draft.is_empty(), "short context should yield empty draft");
}
#[test]
fn prompt_lookup_caps_at_gamma() {
let model = PromptLookupDraftModel::new(2);
let draft = model.draft(&[1, 2, 3, 4, 5, 1, 2], 2);
assert_eq!(draft, vec![3, 4]);
}
#[test]
fn prompt_lookup_finds_most_recent_match() {
let model = PromptLookupDraftModel::new(2);
let draft = model.draft(&[1, 2, 9, 1, 2, 8, 1, 2], 3);
assert_eq!(draft, vec![8]);
}
#[test]
fn prompt_lookup_zero_gamma_returns_empty() {
let model = PromptLookupDraftModel::new(2);
let draft = model.draft(&[1, 2, 3, 1, 2], 0);
assert!(draft.is_empty());
}
}