use crate::config::{FinishReason, GenerateOptions, GenerateResult, GenerateTokenCallback};
use crate::decode_loop::{
DecodeLoopState, commit_selected_token, finish_result, reached_context_limit,
};
use crate::logits::{ProcessorChain, TokenId};
use crate::native_decode::{NativeDecodeSession, NativeProposerSession};
use crate::processors::ensure_constrained_finish;
use crate::sampling::sample_greedy;
use crate::speculative::{
LinearEmbedder, NgramProposer, SpeculativeProposer, SpeculativeProposerContext,
SpeculativeStats, TokenEmbedder, argmax,
};
use anyhow::Context;
use onnx_genai_ort::Tokenizer;
pub(crate) struct NativeSpeculativeDriver<'a> {
session: &'a mut NativeDecodeSession,
proposer: NativeProposer<'a>,
draft_width: usize,
}
enum NativeProposer<'a> {
PromptLookup(NgramProposer),
SharedKv {
session: &'a mut NativeProposerSession,
embedder: &'a LinearEmbedder,
groups: &'a [onnx_genai_metadata::SharedKvGroup],
hidden_size: usize,
},
}
impl<'a> NativeSpeculativeDriver<'a> {
pub(crate) fn new_prompt_lookup(
session: &'a mut NativeDecodeSession,
ngram: usize,
max_tokens: usize,
draft_width: usize,
) -> anyhow::Result<Self> {
Ok(Self {
session,
proposer: NativeProposer::PromptLookup(NgramProposer::new(ngram, max_tokens)?),
draft_width: draft_width.max(1),
})
}
pub(crate) fn new_shared_kv(
session: &'a mut NativeDecodeSession,
proposer_session: &'a mut NativeProposerSession,
embedder: &'a LinearEmbedder,
groups: &'a [onnx_genai_metadata::SharedKvGroup],
hidden_size: usize,
draft_width: usize,
) -> anyhow::Result<Self> {
if hidden_size == 0 || embedder.hidden_size() != hidden_size {
anyhow::bail!(
"native shared-KV proposer hidden size {hidden_size} does not match embedding width {}",
embedder.hidden_size()
);
}
proposer_session.reset();
Ok(Self {
session,
proposer: NativeProposer::SharedKv {
session: proposer_session,
embedder,
groups,
hidden_size,
},
draft_width: draft_width.max(1),
})
}
pub(crate) fn generate(
&mut self,
prompt_tokens: &[TokenId],
options: &GenerateOptions,
chain: &ProcessorChain,
tokenizer: &Tokenizer,
stats: &mut SpeculativeStats,
mut callback: Option<&mut GenerateTokenCallback<'_>>,
) -> anyhow::Result<GenerateResult> {
if prompt_tokens.is_empty() {
anyhow::bail!("native speculative generation requires at least one prompt token");
}
self.session.reset()?;
let prompt_len = prompt_tokens.len();
let mut state = DecodeLoopState::new(0, options.seed, options.top_logprobs);
let mut pending: Vec<TokenId> = prompt_tokens.to_vec();
loop {
if state.generated_tokens.len() >= options.max_new_tokens {
ensure_constrained_finish(options, &state.generated_text, FinishReason::MaxTokens)?;
return finish_result(
tokenizer,
&state.generated_tokens,
FinishReason::MaxTokens,
0,
state.logprobs.as_deref(),
);
}
let context_len = prompt_len + state.generated_tokens.len();
if reached_context_limit(context_len, options.max_context) {
ensure_constrained_finish(options, &state.generated_text, FinishReason::Length)?;
return finish_result(
tokenizer,
&state.generated_tokens,
FinishReason::Length,
0,
state.logprobs.as_deref(),
);
}
let past = self.session.current_len();
let base_logits = self
.session
.decode(&pending, past)?
.pop()
.context("native speculative decode produced no base logits")?;
pending.clear();
let base = self.session.current_len();
debug_assert_eq!(base, context_len);
let remaining_tokens = options.max_new_tokens - state.generated_tokens.len();
let remaining_context = options
.max_context
.map(|limit| limit.saturating_sub(context_len))
.unwrap_or(remaining_tokens);
let width = self
.draft_width
.min(remaining_tokens)
.min(remaining_context)
.max(1);
let context_tokens: Vec<TokenId> = prompt_tokens
.iter()
.copied()
.chain(state.generated_tokens.iter().copied())
.collect();
let mut draft = match &mut self.proposer {
NativeProposer::PromptLookup(proposer) => {
let proposer_context = SpeculativeProposerContext {
width,
context_tokens: &context_tokens,
generated_tokens: &state.generated_tokens,
generated_text: &state.generated_text,
first_step: state.step,
options,
chain,
target_hidden: None,
target_hidden_layers: None,
guaranteed_token: None,
shared_kv_slices: None,
};
proposer.propose(&proposer_context)?.tokens
}
NativeProposer::SharedKv {
session: proposer,
embedder,
groups,
hidden_size,
} => {
let target_hidden = self.session.last_hidden().with_context(|| {
"native shared-KV proposer requires the target decoder's declared io.hidden_output; the target forward produced no hidden state"
})?;
if target_hidden.len() != *hidden_size {
anyhow::bail!(
"native target hidden output has width {}, but shared-KV metadata declares backbone_hidden_size {}; fix model.io.hidden_output or speculative.backbone_hidden_size",
target_hidden.len(),
hidden_size
);
}
let guaranteed = TokenId::try_from(
argmax(&base_logits).context("native target logits were empty")?,
)
.context("native target token id exceeds u32 range")?;
let shared_inputs = self.session.shared_kv_inputs(groups)?;
let seed = *context_tokens
.last()
.context("native shared-KV proposer requires at least one context token")?;
let mut hidden = target_hidden.to_vec();
let mut token = seed;
let mut embeddings = vec![0.0; hidden_size.saturating_mul(2)];
let mut tokens = Vec::with_capacity(width);
tokens.push(guaranteed);
let position = context_tokens.len().saturating_sub(1);
for step in 0..width {
embedder.embed(token, &mut embeddings[..*hidden_size])?;
embeddings[*hidden_size..].copy_from_slice(&hidden);
let output =
proposer.step_inputs_embeds(&embeddings, position, &shared_inputs)?;
hidden = output.projected_state.with_context(|| {
"native shared-KV proposer metadata must assign io.hidden_output to its projected recurrent state"
})?;
let logits = output.logits.with_context(
|| "native shared-KV proposer metadata must assign io.logits_output",
)?;
let drafted = TokenId::try_from(
argmax(
logits
.last()
.context("native shared-KV proposer emitted no logits rows")?,
)
.context("native shared-KV proposer logits row was empty")?,
)
.context("native shared-KV proposer token id exceeds u32 range")?;
if step == 0 {
token = guaranteed;
} else {
tokens.push(drafted);
token = drafted;
}
}
tokens
}
};
draft.truncate(width);
if draft.is_empty() {
let token = sample_greedy(&base_logits);
if let Some(reason) = commit_selected_token(
&mut state,
prompt_tokens,
token,
options,
chain,
tokenizer,
callback.as_deref_mut(),
)? {
return finish_result(
tokenizer,
&state.generated_tokens,
reason,
0,
state.logprobs.as_deref(),
);
}
pending.push(token);
continue;
}
stats.verification_steps += 1;
stats.proposed_tokens += draft.len();
let rows = self.session.decode_verify(&draft, base)?;
let mut target_tokens = Vec::with_capacity(rows.len() + 1);
target_tokens.push(sample_greedy(&base_logits));
for row in &rows {
target_tokens.push(sample_greedy(row));
}
let mut accepted = 0usize;
while accepted < draft.len() && target_tokens[accepted] == draft[accepted] {
accepted += 1;
}
let bonus = target_tokens[accepted];
stats.accepted_tokens += accepted;
if accepted >= 2 {
stats.multi_token_accepts += 1;
}
self.session.rewind(base + accepted)?;
let mut commit_iter = draft[..accepted]
.iter()
.copied()
.chain(std::iter::once(bonus))
.enumerate();
for (idx, token) in commit_iter.by_ref() {
let is_bonus = idx == accepted;
if state.generated_tokens.len() >= options.max_new_tokens {
break;
}
if is_bonus {
let context_now = prompt_len + state.generated_tokens.len();
if reached_context_limit(context_now, options.max_context) {
break;
}
}
if let Some(reason) = commit_selected_token(
&mut state,
prompt_tokens,
token,
options,
chain,
tokenizer,
callback.as_deref_mut(),
)? {
return finish_result(
tokenizer,
&state.generated_tokens,
reason,
0,
state.logprobs.as_deref(),
);
}
if is_bonus {
pending.push(token);
}
}
}
}
}