diff --git a/openinfer-qwen3/src/eagle3.rs b/openinfer-qwen3/src/eagle3.rs index ce85bc133..1a09a8139 100644 --- a/openinfer-qwen3/src/eagle3.rs +++ b/openinfer-qwen3/src/eagle3.rs @@ -6,10 +6,15 @@ use anyhow::Result; use crate::config::Eagle3Config; use openinfer_core::tensor::{DeviceMatrix, DeviceVec}; +mod forward; mod loading; mod reservation; -// Wired into the KV budget by the forward-pass PR; kept here as the skeleton lands. +// Consumed by the scheduler PR that drives the drafter; exported as the forward +// path lands. +#[allow(unused_imports)] +pub(crate) use forward::{Eagle3RequestState, Eagle3Scratch}; +// Wired into the KV budget by the scheduler PR; kept here as the skeleton lands. #[allow(unused_imports)] pub(crate) use reservation::Eagle3MemoryReservation; diff --git a/openinfer-qwen3/src/eagle3/forward.rs b/openinfer-qwen3/src/eagle3/forward.rs new file mode 100644 index 000000000..527acfed4 --- /dev/null +++ b/openinfer-qwen3/src/eagle3/forward.rs @@ -0,0 +1,883 @@ +use anyhow::{Context, Result}; +use cudarc::driver::CudaSlice; + +use crate::weights::Qwen3Model; +use openinfer_core::ops; +use openinfer_core::tensor::{DeviceContext, HiddenStates}; + +use super::Eagle3DraftModel; + +/// Draft slots available at `cached_len`, capped to `requested_k`. +/// +/// The constrain is `cached_len + k + 1 <= max_cache_len`, where +1 is the bonus token +/// generated by speculative decoding +fn chain_capacity(cached_len: usize, max_cache_len: usize, requested_k: usize) -> usize { + let draftable = max_cache_len.saturating_sub(cached_len).saturating_sub(1); + requested_k.min(draftable) +} + +/// Per-request state the EAGLE-3 draft carries *between* forwards. +/// To record kv cache size for draft memory allocation +pub(crate) struct Eagle3RequestState { + k: HiddenStates, + v: HiddenStates, + cached_len: usize, + max_cache_len: usize, + /// Boundary target feature `[3 * hidden, 1]` at the last committed position + last_aux_hidden_states: Option, +} + +impl Eagle3RequestState { + pub(crate) fn cached_len(&self) -> usize { + self.cached_len + } + + /// The boundary target feature `[3 * hidden, 1]` (`None` until captured). + pub(crate) fn last_aux_hidden_states(&self) -> Option<&HiddenStates> { + self.last_aux_hidden_states.as_ref() + } +} + +/// Single-token draft scratch (`seq_len == 1` everywhere). v1 runs one request, +/// one token at a time — batching the draft chain is a follow-up +pub(crate) struct Eagle3Scratch { + token_id_d: CudaSlice, + embed: HiddenStates, // [hidden, 1] + hidden: HiddenStates, // [hidden, 1] residual stream + normed_embed: HiddenStates, // [hidden, 1] + normed_hidden: HiddenStates, // [hidden, 1] + attn_input: HiddenStates, // [2 * hidden, 1] + q: HiddenStates, // [q_dim, 1] + k: HiddenStates, // [kv_dim, 1] + v: HiddenStates, // [kv_dim, 1] + attn_out: HiddenStates, // [q_dim, 1] + o: HiddenStates, // [hidden, 1] + normed_post: HiddenStates, // [hidden, 1] + gate: HiddenStates, // [inter, 1] + up: HiddenStates, // [inter, 1] + act: HiddenStates, // [inter, 1] + mlp_out: HiddenStates, // [hidden, 1] + normed_final: HiddenStates, // [hidden, 1] + logits: HiddenStates, // [draft_vocab, 1] +} + +impl Eagle3DraftModel { + fn q_dim(&self) -> usize { + self.midlayer.q_dim + } + + fn kv_dim(&self) -> usize { + self.midlayer.kv_dim + } + + /// Assert a request state was allocated for this drafter's geometry. + fn ensure_state_geometry(&self, state: &Eagle3RequestState) -> Result<()> { + let kv_dim = self.kv_dim(); + anyhow::ensure!( + state.k.hidden_dim == kv_dim && state.v.hidden_dim == kv_dim, + "EAGLE-3 request-state K/V dim [{}, {}] does not match drafter kv_dim {}", + state.k.hidden_dim, + state.v.hidden_dim, + kv_dim + ); + Ok(()) + } + + /// Assert a scratch buffer set was allocated for this drafter's geometry. + fn ensure_scratch_geometry(&self, scratch: &Eagle3Scratch) -> Result<()> { + let hidden = self.config.hidden_size; + let q_dim = self.q_dim(); + let kv_dim = self.kv_dim(); + let inter = self.config.intermediate_size; + let vocab = self.config.draft_vocab_size; + anyhow::ensure!( + scratch.q.hidden_dim == q_dim + && scratch.k.hidden_dim == kv_dim + && scratch.v.hidden_dim == kv_dim + && scratch.attn_input.hidden_dim == 2 * hidden + && scratch.hidden.hidden_dim == hidden + && scratch.gate.hidden_dim == inter + && scratch.logits.hidden_dim == vocab, + "EAGLE-3 scratch geometry does not match drafter \ + (hidden {hidden}, q_dim {q_dim}, kv_dim {kv_dim}, inter {inter}, vocab {vocab})" + ); + Ok(()) + } + + /// Allocate the single-layer K/V cache for one request. `max_cache_len` bounds + /// the total drafted+committed positions and must fit the rope cache. + pub(crate) fn new_request_state( + &self, + ctx: &DeviceContext, + max_cache_len: usize, + ) -> Result { + anyhow::ensure!( + max_cache_len > 0 && max_cache_len <= self.config.max_position_embeddings, + "EAGLE-3 request cache length {} must be in 1..={}", + max_cache_len, + self.config.max_position_embeddings + ); + let kv_dim = self.kv_dim(); + let k = HiddenStates::zeros(ctx, kv_dim, max_cache_len)?; + let v = HiddenStates::zeros(ctx, kv_dim, max_cache_len)?; + Ok(Eagle3RequestState { + k, + v, + cached_len: 0, + max_cache_len, + last_aux_hidden_states: None, + }) + } + + /// Allocate the single-token draft scratch. + pub(crate) fn new_scratch(&self, ctx: &DeviceContext) -> Result { + let hidden = self.config.hidden_size; + let q_dim = self.q_dim(); + let kv_dim = self.kv_dim(); + let inter = self.config.intermediate_size; + Ok(Eagle3Scratch { + token_id_d: ctx.stream.alloc_zeros(1)?, + embed: HiddenStates::zeros(ctx, hidden, 1)?, + hidden: HiddenStates::zeros(ctx, hidden, 1)?, + normed_embed: HiddenStates::zeros(ctx, hidden, 1)?, + normed_hidden: HiddenStates::zeros(ctx, hidden, 1)?, + attn_input: HiddenStates::zeros(ctx, 2 * hidden, 1)?, + q: HiddenStates::zeros(ctx, q_dim, 1)?, + k: HiddenStates::zeros(ctx, kv_dim, 1)?, + v: HiddenStates::zeros(ctx, kv_dim, 1)?, + attn_out: HiddenStates::zeros(ctx, q_dim, 1)?, + o: HiddenStates::zeros(ctx, hidden, 1)?, + normed_post: HiddenStates::zeros(ctx, hidden, 1)?, + gate: HiddenStates::zeros(ctx, inter, 1)?, + up: HiddenStates::zeros(ctx, inter, 1)?, + act: HiddenStates::zeros(ctx, inter, 1)?, + mlp_out: HiddenStates::zeros(ctx, hidden, 1)?, + normed_final: HiddenStates::zeros(ctx, hidden, 1)?, + logits: HiddenStates::zeros(ctx, self.config.draft_vocab_size, 1)?, + }) + } + + /// Fuse the draft's input hidden from target context features (test helper). + pub(crate) fn fuse_input_hidden_from_context( + &self, + ctx: &DeviceContext, + context_features: &HiddenStates, + scratch: &mut Eagle3Scratch, + ) -> Result<()> { + anyhow::ensure!( + context_features.hidden_dim == self.fc_input_dim(), + "EAGLE-3 context feature dim {} does not match fc input {}", + context_features.hidden_dim, + self.fc_input_dim() + ); + anyhow::ensure!( + context_features.seq_len == 1, + "EAGLE-3 input-hidden fuse expects one position, got {}", + context_features.seq_len + ); + self.ensure_scratch_geometry(scratch)?; + ops::gemm_into(ctx, &self.fc, context_features, &mut scratch.hidden); + Ok(()) + } + + fn fc_input_dim(&self) -> usize { + // fc: [hidden, 3 * hidden] , 3 * hidden for fuser input + self.fc.cols + } + + /// One EAGLE-3 draft step: consume `token_id` (the current token) and the + /// residual stream in `scratch.hidden`, update `scratch.hidden` as output + pub(crate) fn draft_step<'s>( + &self, + target: &Qwen3Model, + state: &mut Eagle3RequestState, + scratch: &'s mut Eagle3Scratch, + token_id: u32, + position: usize, + ) -> Result<&'s HiddenStates> { + let ctx = target.device_ctx(); + let hidden = self.config.hidden_size; + let q_dim = self.q_dim(); + let kv_dim = self.kv_dim(); + let inter = self.config.intermediate_size; + let eps = self.config.rms_norm_eps; + let num_q = self.config.num_attention_heads; + let num_kv = self.config.num_key_value_heads; + let head_dim = self.config.head_dim; + + anyhow::ensure!( + position < state.max_cache_len, + "EAGLE-3 draft position {} exceeds cache {}", + position, + state.max_cache_len + ); + anyhow::ensure!( + position == state.cached_len, + "EAGLE-3 draft step expects position {} == cached_len {}", + position, + state.cached_len + ); + self.ensure_state_geometry(state)?; + self.ensure_scratch_geometry(scratch)?; + + // 1. Embed the current token (reuses the target's embed_tokens). + { + let mut dst = scratch.token_id_d.slice_mut(..1); + ctx.stream.memcpy_htod(&[token_id], &mut dst)?; + } + target.get_embeddings_batch_into(&scratch.token_id_d, &mut scratch.embed)?; + + // 2. Norm the embedding and the fused hidden separately. + ops::rms_norm_batch_into( + ctx, + &scratch.embed, + &self.midlayer.input_layernorm, + eps, + &mut scratch.normed_embed, + ); + ops::rms_norm_batch_into( + ctx, + &scratch.hidden, + &self.midlayer.hidden_norm, + eps, + &mut scratch.normed_hidden, + ); + + // 3. attn_input = [normed_embed (rows 0..hidden) | normed_hidden (rows hidden..2h)]. + ops::copy_hidden_rows_into(ctx, &scratch.normed_embed, &mut scratch.attn_input, 0)?; + ops::copy_hidden_rows_into(ctx, &scratch.normed_hidden, &mut scratch.attn_input, hidden)?; + + // 4. q/k/v projections (qkv_proj input is 2 * hidden). + ops::gemm_rows_into( + ctx, + &self.midlayer.qkv_proj, + 0, + q_dim, + &scratch.attn_input, + &mut scratch.q, + ); + ops::gemm_rows_into( + ctx, + &self.midlayer.qkv_proj, + q_dim, + kv_dim, + &scratch.attn_input, + &mut scratch.k, + ); + ops::gemm_rows_into( + ctx, + &self.midlayer.qkv_proj, + q_dim + kv_dim, + kv_dim, + &scratch.attn_input, + &mut scratch.v, + ); + + // 5. Plain RoPE (no QK-norm) on the single q/k token at `position`. + ops::eagle3_rope_into( + ctx, + &mut scratch.q, + 0, + 1, + &mut scratch.k, + &self.cos_cache, + &self.sin_cache, + num_q, + num_kv, + head_dim, + position, + position, + )?; + + // 6. Append the rotated K/V into the cache at `position`. + ops::copy_hidden_token_range_into(ctx, &scratch.k, 0, &mut state.k, position, 1)?; + ops::copy_hidden_token_range_into(ctx, &scratch.v, 0, &mut state.v, position, 1)?; + let kv_len = position + 1; + + // 7. Attention: single-query decode — the one draft query attends the whole + // [0, kv_len) prefix of the draft's contiguous KV. + ops::single_decode_nhd_into( + ctx, + &scratch.q, + &state.k, + &state.v, + &mut scratch.attn_out, + num_q, + num_kv, + head_dim, + kv_len, + )?; + + // 8. Output projection. + ops::gemm_into( + ctx, + &self.midlayer.o_proj, + &scratch.attn_out, + &mut scratch.o, + ); + + // 9. Residual + post-attention norm (hidden += o in place; normed_post = norm(hidden)). + openinfer_kernels::ops::fused_add_rms_norm_round_batch_into( + ctx, + &mut scratch.hidden, + &scratch.o, + &self.midlayer.post_attention_layernorm, + eps, + &mut scratch.normed_post, + )?; + + // 10. MLP (SwiGLU). + ops::gemm_rows_into( + ctx, + &self.midlayer.gate_up_proj, + 0, + inter, + &scratch.normed_post, + &mut scratch.gate, + ); + ops::gemm_rows_into( + ctx, + &self.midlayer.gate_up_proj, + inter, + inter, + &scratch.normed_post, + &mut scratch.up, + ); + ops::silu_mul_batch_into(ctx, &scratch.gate, &scratch.up, &mut scratch.act)?; + ops::gemm_into( + ctx, + &self.midlayer.down_proj, + &scratch.act, + &mut scratch.mlp_out, + ); + + // 11. Residual + final norm. + openinfer_kernels::ops::fused_add_rms_norm_round_batch_into( + ctx, + &mut scratch.hidden, + &scratch.mlp_out, + &self.norm, + eps, + &mut scratch.normed_final, + )?; + + // 12. Draft head over the reduced vocabulary. + ops::gemm_into( + ctx, + &self.lm_head, + &scratch.normed_final, + &mut scratch.logits, + ); + + state.cached_len = kv_len; + Ok(&scratch.logits) + } + + /// Teacher-forced prefil for EAGLE Draft + /// Returns `(logits [draft_vocab, N], last_hidden [hidden, 1])` + /// Buffers are allocated inline (prefill is one-shot, `N` varies). Does not use + /// the single-token `Eagle3Scratch`. `tokens[i]` sits at `start_position + i`. + /// + /// NOTE: the returned `(logits, last_hidden)` is not consumed yet — the current + /// callers only need the KV/boundary side effects. It is reserved for the + /// scheduler PR that drives prompt prefill (first-token logits / chain input hidden). + pub(crate) fn prefill_batched( + &self, + target: &Qwen3Model, + state: &mut Eagle3RequestState, + features: &HiddenStates, + tokens: &[u32], + start_position: usize, + ) -> Result<(HiddenStates, HiddenStates)> { + let num_tokens = tokens.len(); + anyhow::ensure!(num_tokens > 0, "EAGLE-3 prefill needs tokens"); + anyhow::ensure!( + features.hidden_dim == self.fc_input_dim() && features.seq_len == num_tokens, + "EAGLE-3 batched prefill needs features [{}, {}], got [{}, {}]", + self.fc_input_dim(), + num_tokens, + features.hidden_dim, + features.seq_len + ); + anyhow::ensure!( + start_position == state.cached_len, + "EAGLE-3 batched prefill expects start {} == cached_len {}", + start_position, + state.cached_len + ); + anyhow::ensure!( + start_position + num_tokens <= state.max_cache_len, + "EAGLE-3 batched prefill overflows cache: {} + {} > {}", + start_position, + num_tokens, + state.max_cache_len + ); + self.ensure_state_geometry(state)?; + + let ctx = target.device_ctx(); + let hidden = self.config.hidden_size; + let q_dim = self.q_dim(); + let kv_dim = self.kv_dim(); + let inter = self.config.intermediate_size; + let eps = self.config.rms_norm_eps; + let num_q = self.config.num_attention_heads; + let num_kv = self.config.num_key_value_heads; + let head_dim = self.config.head_dim; + + // 1. Embed all N tokens (reuses the target's embed_tokens). + let mut token_ids_d = ctx.stream.alloc_zeros::(num_tokens)?; + ctx.stream.memcpy_htod(tokens, &mut token_ids_d)?; + let mut embed = HiddenStates::zeros(ctx, hidden, num_tokens)?; + target.get_embeddings_batch_into(&token_ids_d, &mut embed)?; + + // 2. Residual stream = fc(per-position target features) — teacher forcing. + // (draft_step gets this as the input hidden in `scratch.hidden` instead.) + let mut residual = HiddenStates::zeros(ctx, hidden, num_tokens)?; + ops::gemm_into(ctx, &self.fc, features, &mut residual); + + // 3. Norm the embedding and the teacher-forced hidden separately. + let mut normed_embed = HiddenStates::zeros(ctx, hidden, num_tokens)?; + let mut normed_hidden = HiddenStates::zeros(ctx, hidden, num_tokens)?; + ops::rms_norm_batch_into( + ctx, + &embed, + &self.midlayer.input_layernorm, + eps, + &mut normed_embed, + ); + ops::rms_norm_batch_into( + ctx, + &residual, + &self.midlayer.hidden_norm, + eps, + &mut normed_hidden, + ); + + // 4. attn_input = [normed_embed (rows 0..h) | normed_hidden (rows h..2h)]. + let mut attn_input = HiddenStates::zeros(ctx, 2 * hidden, num_tokens)?; + ops::copy_hidden_rows_into(ctx, &normed_embed, &mut attn_input, 0)?; + ops::copy_hidden_rows_into(ctx, &normed_hidden, &mut attn_input, hidden)?; + + // 5. q/k/v projections (qkv_proj input is 2 * hidden). + let mut query = HiddenStates::zeros(ctx, q_dim, num_tokens)?; + let mut key = HiddenStates::zeros(ctx, kv_dim, num_tokens)?; + let mut value = HiddenStates::zeros(ctx, kv_dim, num_tokens)?; + ops::gemm_rows_into( + ctx, + &self.midlayer.qkv_proj, + 0, + q_dim, + &attn_input, + &mut query, + ); + ops::gemm_rows_into( + ctx, + &self.midlayer.qkv_proj, + q_dim, + kv_dim, + &attn_input, + &mut key, + ); + ops::gemm_rows_into( + ctx, + &self.midlayer.qkv_proj, + q_dim + kv_dim, + kv_dim, + &attn_input, + &mut value, + ); + + // 6. Plain RoPE (no QK-norm) on all N q/k at positions [start, start+N). + ops::eagle3_rope_into( + ctx, + &mut query, + 0, + num_tokens, + &mut key, + &self.cos_cache, + &self.sin_cache, + num_q, + num_kv, + head_dim, + start_position, + start_position, + )?; + + // 7. Append all N rotated k/v into the cache at [start, start+N). + ops::copy_hidden_token_range_into(ctx, &key, 0, &mut state.k, start_position, num_tokens)?; + ops::copy_hidden_token_range_into( + ctx, + &value, + 0, + &mut state.v, + start_position, + num_tokens, + )?; + let kv_len = start_position + num_tokens; + + // 8. One causal attention over the [0, kv_len) prefix. + let mut attn_out = HiddenStates::zeros(ctx, q_dim, num_tokens)?; + ops::single_prefill_nhd_causal_into( + ctx, + &query, + 0, + num_tokens, + &state.k, + &state.v, + &mut attn_out, + num_q, + num_kv, + head_dim, + kv_len, + )?; + + // 9. Output projection. + let mut attn_proj = HiddenStates::zeros(ctx, hidden, num_tokens)?; + ops::gemm_into(ctx, &self.midlayer.o_proj, &attn_out, &mut attn_proj); + + // 10. residual += attn_proj; normed_post = norm(residual). + let mut normed_post = HiddenStates::zeros(ctx, hidden, num_tokens)?; + openinfer_kernels::ops::fused_add_rms_norm_round_batch_into( + ctx, + &mut residual, + &attn_proj, + &self.midlayer.post_attention_layernorm, + eps, + &mut normed_post, + )?; + + // 11. MLP (SwiGLU). + let mut gate = HiddenStates::zeros(ctx, inter, num_tokens)?; + let mut up = HiddenStates::zeros(ctx, inter, num_tokens)?; + let mut act = HiddenStates::zeros(ctx, inter, num_tokens)?; + ops::gemm_rows_into( + ctx, + &self.midlayer.gate_up_proj, + 0, + inter, + &normed_post, + &mut gate, + ); + ops::gemm_rows_into( + ctx, + &self.midlayer.gate_up_proj, + inter, + inter, + &normed_post, + &mut up, + ); + ops::silu_mul_batch_into(ctx, &gate, &up, &mut act)?; + let mut mlp_out = HiddenStates::zeros(ctx, hidden, num_tokens)?; + ops::gemm_into(ctx, &self.midlayer.down_proj, &act, &mut mlp_out); + + // 12. residual += mlp_out; normed_final = norm(residual). + let mut normed_final = HiddenStates::zeros(ctx, hidden, num_tokens)?; + openinfer_kernels::ops::fused_add_rms_norm_round_batch_into( + ctx, + &mut residual, + &mlp_out, + &self.norm, + eps, + &mut normed_final, + )?; + + // 13. Draft head over the reduced vocabulary. + let mut logits = HiddenStates::zeros(ctx, self.config.draft_vocab_size, num_tokens)?; + ops::gemm_into(ctx, &self.lm_head, &normed_final, &mut logits); + + // 14. The last position's decoder output (post-mlp residual) is the chain's input hidden. + let mut last_hidden = HiddenStates::zeros(ctx, hidden, 1)?; + ops::copy_hidden_token_range_into(ctx, &residual, num_tokens - 1, &mut last_hidden, 0, 1)?; + + state.cached_len = kv_len; + Ok((logits, last_hidden)) + } + + /// Build the draft KV for a freshly-prefilled prompt and record the boundary + /// feature, applying the EAGLE feature↔token shift: pair feature `f_j` with the + /// next token's embedding `e_{j+1}` (predicting `t_{j+2}`). Teacher-forces the + /// `P-1` pairs (`j = 0..P-2`) into the draft KV; the last feature `f_{P-1}` is + /// kept as the boundary (`last_aux_hidden_states`) — it pairs with the first + /// *generated* token in the next round. + /// + /// `captured_all` is the batch-wide capture, `token_offset` this request's first + /// column in it. v1 needs the whole prompt in one chunk (`cached_len == 0`); + /// the caller skips longer prompts and falls back to plain decode. + pub(crate) fn prefill_prompt( + &self, + target: &Qwen3Model, + state: &mut Eagle3RequestState, + captured_all: &HiddenStates, + token_offset: usize, + prompt_tokens: &[u32], + ) -> Result<()> { + let p = prompt_tokens.len(); + anyhow::ensure!(p >= 2, "EAGLE-3 prefill needs >= 2 prompt tokens"); + anyhow::ensure!( + state.cached_len == 0, + "EAGLE-3 prefill_prompt expects a fresh state (cached_len {})", + state.cached_len + ); + anyhow::ensure!( + captured_all.hidden_dim == self.fc_input_dim(), + "EAGLE-3 capture features have dim {} but fc expects {}", + captured_all.hidden_dim, + self.fc_input_dim() + ); + anyhow::ensure!( + token_offset + p <= captured_all.seq_len, + "EAGLE-3 capture slice [{}, {}) overflows {} captured rows", + token_offset, + token_offset + p, + captured_all.seq_len + ); + let ctx = target.device_ctx(); + let dim = self.fc_input_dim(); + // Teacher-force the shifted pairs (f_j, e_{j+1}) for j = 0..P-2. + let mut feat = HiddenStates::zeros(ctx, dim, p - 1)?; + ops::copy_hidden_token_range_into(ctx, captured_all, token_offset, &mut feat, 0, p - 1)?; + self.prefill_batched(target, state, &feat, &prompt_tokens[1..p], 0)?; + // Boundary = the last target feature f_{P-1}, kept pre-fc. + let mut boundary = HiddenStates::zeros(ctx, dim, 1)?; + ops::copy_hidden_token_range_into( + ctx, + captured_all, + token_offset + p - 1, + &mut boundary, + 0, + 1, + )?; + state.last_aux_hidden_states = Some(boundary); + Ok(()) + } + + /// Autoregressive **chain** draft (v1): from the fused target boundary + /// hidden (`fused_target_hidden` = `fc(last_aux_hidden_states)`) and the last + /// committed token, draft `k` tokens one at a time — each `draft_step` produces + /// logits, greedy-argmax picks a draft id, `d2t` maps it to the target vocab. + /// + /// The drafted token feeds the next step while the drafter's own hidden state + /// (the residual stream) is updated in place and carried into the next step. + /// + /// Next Step: v1 syncs logits to host per step for the argmax; device-side sampling (like + /// DFlash `select_step_tokens`) is a perf follow-up. + pub(crate) fn draft_chain( + &self, + target: &Qwen3Model, + state: &mut Eagle3RequestState, + scratch: &mut Eagle3Scratch, + fused_target_hidden: &HiddenStates, + last_token: u32, + start_position: usize, + k: usize, + ) -> Result> { + anyhow::ensure!(k > 0, "EAGLE-3 draft chain needs k > 0"); + anyhow::ensure!( + fused_target_hidden.hidden_dim == self.config.hidden_size + && fused_target_hidden.seq_len == 1, + "EAGLE-3 fused_target_hidden must be [hidden, 1]" + ); + anyhow::ensure!( + chain_capacity(start_position, state.max_cache_len, k) == k, + "EAGLE-3 chain of {} drafts at {} overflows cache {} — an all-accepted \ + round commits one slot more than it drafts", + k, + start_position, + state.max_cache_len + ); + let ctx = target.device_ctx(); + + // Initialize the residual stream with the fused target boundary hidden. + ops::copy_hidden_token_range_into(ctx, fused_target_hidden, 0, &mut scratch.hidden, 0, 1)?; + + let mut span = Vec::with_capacity(k); + let mut token = last_token; + for i in 0..k { + // Scope the logits borrow so the next step can re-borrow `scratch`. + let host = { + let logits = self.draft_step(target, state, scratch, token, start_position + i)?; + logits.to_host(ctx)? + }; + let draft_id = host + .iter() + .enumerate() + .max_by(|a, b| a.1.partial_cmp(b.1).unwrap_or(std::cmp::Ordering::Equal)) + .map(|(idx, _)| idx) + .expect("non-empty draft logits"); + let target_id = self + .draft_to_target_id(draft_id) + .context("EAGLE-3 draft id maps outside target vocab")?; + span.push(target_id); + token = target_id; + } + Ok(span) + } + + /// One speculative draft round for a single request: Initialize the residual stream + /// with the fused target boundary hidden (`fc(last_aux_hidden_states)`), draft + /// `k` tokens via `draft_chain`, then rewind the draft KV to the round-start + /// slot so the round is side-effect-free except for the returned tokens. + /// + /// The chain's KV writes (slots `[C, C+k)`) are speculative; `rebuild_after_verify` + /// rebuilds the accepted prefix teacher-forced from the verify's captured target + /// hidden, so we discard them here by resetting `cached_len` to `C`. + pub(crate) fn chain_round( + &self, + target: &Qwen3Model, + state: &mut Eagle3RequestState, + scratch: &mut Eagle3Scratch, + current_token: u32, + k: usize, + ) -> Result> { + let start = state.cached_len; + let k = chain_capacity(start, state.max_cache_len, k); + if k == 0 { + log::debug!( + "EAGLE-3 draft cache is full at {} of {}; skipping the chain round", + start, + state.max_cache_len + ); + return Ok(Vec::new()); + } + let ctx = target.device_ctx(); + let mut fused_target_hidden = HiddenStates::zeros(ctx, self.config.hidden_size, 1)?; + { + let feature = state + .last_aux_hidden_states + .as_ref() + .context("EAGLE-3 draft chain has no boundary feature (prompt not captured?)")?; + debug_assert_eq!( + feature.hidden_dim, + self.fc_input_dim(), + "last_aux_hidden_states must be [fc_input_dim, 1]" + ); + ops::gemm_into(ctx, &self.fc, feature, &mut fused_target_hidden); + } + let result = self.draft_chain( + target, + state, + scratch, + &fused_target_hidden, + current_token, + start, + k, + ); + // Discard the speculative chain KV; `rebuild_after_verify` rebuilds the accepted + // prefix teacher-forced. On a draft error the request is dropped. + state.cached_len = start; + result + } + + /// Rebuild the Draft state(KV Cache and Target boundary feature) after a verify step. + pub(crate) fn rebuild_after_verify( + &self, + target: &Qwen3Model, + state: &mut Eagle3RequestState, + captured: &HiddenStates, + token_offset: usize, + span_tokens: &[u32], + matched_draft_tokens: usize, + ) -> Result<()> { + let n = matched_draft_tokens; + let len = n + 1; + anyhow::ensure!( + len <= span_tokens.len(), + "EAGLE-3 Draft State rebuild needs {} span tokens(1 bonus+k drafts), got {}", + len, + span_tokens.len() + ); + anyhow::ensure!( + captured.hidden_dim == self.fc_input_dim(), + "EAGLE-3 verify capture has dim {} but fc expects {}", + captured.hidden_dim, + self.fc_input_dim() + ); + // Need captured columns [pos .. pos + n - 1] for the matched-feature pairs + // plus pos + n for KV Cache rebuild for the new boundary target feature. + anyhow::ensure!( + token_offset + n < captured.seq_len, + "EAGLE-3 KV Cache rebuild slice [{}, {}] overflows {} captured rows", + token_offset, + token_offset + n, + captured.seq_len + ); + let ctx = target.device_ctx(); + let dim = self.fc_input_dim(); + + // Shifted feature row: [boundary f_{pos-1} ‖ f_pos..f_{pos+n-1}] = n+1 cols + let mut feat = HiddenStates::zeros(ctx, dim, len)?; + { + let boundary = state + .last_aux_hidden_states + .as_ref() + .context("EAGLE-3 rebuild has no boundary feature")?; + debug_assert_eq!( + boundary.hidden_dim, dim, + "last_aux_hidden_states must be [fc_input_dim, 1]" + ); + ops::copy_hidden_token_range_into(ctx, boundary, 0, &mut feat, 0, 1)?; + } + if n > 0 { + ops::copy_hidden_token_range_into(ctx, captured, token_offset, &mut feat, 1, n)?; + } + // Teacher-force the committed prefix into draft slots [C, C+n+1). + let start_position = state.cached_len; + self.prefill_batched(target, state, &feat, &span_tokens[..len], start_position)?; + + // New boundary feature = f_{pos+n} (the target feature at the last committed position. + let mut new_boundary = HiddenStates::zeros(ctx, dim, 1)?; + ops::copy_hidden_token_range_into( + ctx, + captured, + token_offset + n, + &mut new_boundary, + 0, + 1, + )?; + state.last_aux_hidden_states = Some(new_boundary); + Ok(()) + } + + /// Map a draft-vocabulary id back to the target vocabulary: + /// `target_id = draft_id + d2t[draft_id]`. + pub(crate) fn draft_to_target_id(&self, draft_id: usize) -> Option { + let offset = *self.d2t.get(draft_id)?; + let target_id = draft_id as i64 + offset; + (0..self.config.vocab_size as i64) + .contains(&target_id) + .then_some(target_id as u32) + } +} + +#[cfg(test)] +mod tests { + use super::chain_capacity; + use crate::eagle3::EAGLE3_CHAIN_LENGTH; + + #[test] + fn chain_shrinks_at_cache_edge() { + let (max, k) = (100usize, 3usize); + + // when max_cache=100 and k=3 and current kv length is 97, when all 3 token pass verify + // speculative decoding actually generate 4 token, which causing cache overflow + let granted = chain_capacity(97, max, k); + + // chain_capacity will return granted=2, which means we only draft 2 token to avoid overflow + assert_eq!(granted, 2, "one draft given up so the commit fits"); + assert_eq!(97 + granted + 1, max, "commits exactly to the cache end"); + + // when max_cache=100 and k=3 and current kv length is 96, when all 3 token pass verify + // speculative decoding actually generate 4 token, in this case will not causing overflow + assert_eq!(chain_capacity(96, max, k), 3); + assert_eq!( + chain_capacity(50, max, k), + 3, + "unshrunk away from the boundary" + ); + + assert_eq!(chain_capacity(98, max, k), 1); + assert_eq!(chain_capacity(99, max, k), 0); + assert_eq!(chain_capacity(100, max, k), 0, "cache already full"); + assert_eq!(chain_capacity(0, 0, k), 0, "zero-length cache never drafts"); + } +}