diff --git a/docs/index.md b/docs/index.md index 08b6f7934..a2e5fb8a6 100644 --- a/docs/index.md +++ b/docs/index.md @@ -51,6 +51,7 @@ Organized by domain (model line / subsystem / playbook / lesson) instead of by l | `models/qwen35/kv-admission.md` | Issue #254 complete: Qwen3.5 now uses full-lifetime KV admission, deferred pressure handling, impossible-request rejection, explicit error semantics, direct rejection-event coverage, RTX 5090 e2e, and real HTTP pressure/post-pressure validation. | | `models/qwen35/optimization.md` | Hybrid 24 linear + 8 full attn optimization ledger. Decode-tuning refresh fuses MLP gate/up and tunes decode cublasLt buckets, improving direct TPOT by 2-3%; vLLM still leads 1024/256 HTTP decode. | | `models/qwen35/accuracy.md` | Qwen3.5 HF bf16 logits goldens, size-keyed (0.8b/2b/4b/9b/27b all committed), through `past_key_values`: short replay covers sequential graph, bucket-straddling batched graph, and slot-compaction; long replay covers 4097/8192-token prompts; full GSM8K 8-shot now matches the HF baseline within 0.15 percentage points. | +| `models/qwen35/speculative-verifier.md` | Extracts the target-only Qwen3.5 speculative verifier from PR #626 with transactional hybrid-state handling; DFlash drafter and serving integration stay out of this slice. | | `models/qwen35/model-crate.md` | `pegainfer-qwen35` owns Qwen3.5 model/scheduler/recurrent ops/tests/benches; feature-gated behind `qwen35` (Triton AOT is the only Python build dependency); root loads it through `EngineHandle`. Build/check/clippy, root bench sanity check, historical Qwen3.5 e2e, and scheduler e2e records live here. | | `models/qwen35/batched-step-tail.md` | Qwen3.5 issue #353 implementation record: final prefill tail is batched, decode/unified sample from batched logits, host full-vocab copies are logprobs-only, HF + scheduler e2e pass, and final serving A/B supports only the first-token/short-output TTFT claim. | | `models/qwen35/tp-design.md` | Qwen3.5 TP design: Phase 1 is eager dense TP on Qwen3's controller/worker runtime; validate TP2 first, fail closed for indivisible degrees and TP+CUDA Graph, shard dense full-attention/MLP, and leave sharded linear/GDR state to follow-up. | diff --git a/docs/models/qwen35/speculative-verifier.md b/docs/models/qwen35/speculative-verifier.md new file mode 100644 index 000000000..61c495771 --- /dev/null +++ b/docs/models/qwen35/speculative-verifier.md @@ -0,0 +1,34 @@ +# Qwen3.5 Speculative Verifier + +> **TL;DR:** PR #667 extracts the target-only verifier from #626; for the C16, prompt-1024, span-5 test, reusing existing Q/K prep and paged K/V scatter cut aggregate prep-kernel GPU time by 21.1%, while drafter and serving wiring remain follow-ups. +> +> **Last touched:** 2026-07 + +## Contract + +- Target verification only; no draft model, scheduler/server wiring, sampling, or serving-performance claim. +- Verifier reuses the existing batched Q/K prep and paged K/V scatter; normal single-request prefill keeps its fused path. +- Each request supplies a non-empty `[current token, draft tokens...]` span. A one-token span is valid when one output token remains. +- Greedy acceptance commits the matching draft prefix plus one target token. +- Full acceptance keeps verified KV and recurrent state. Partial acceptance truncates KV, restores recurrent/convolution state, and replays only the accepted span. +- Backup, verify, commit, and rollback use the context stream. Stream overrides are rejected before mutation. +- Any error after mutation restores every canonical state component; rollback failure is executor-fatal. +- Verifier logits and sampling use `selection_vocab`, so checkpoint padding rows cannot produce token ids the frontend cannot decode. + +## Verified + +- Passed: Qwen3.5 release check and Clippy; RTX 5090 verifier tests 11/11. Earlier gates also passed HF golden 2/2, scheduler E2E 1/1, page-pool 4/4, and KV-pool 6/6. + +| C16, prompt 1024, span 5 | Aggregate prep-kernel GPU time, 3 runs | Median | Whole-test median, 5 runs | +| --- | --- | --- | --- | +| Fused verifier kernels (`698ccbd`) | 52.256 / 52.448 / 52.351 us | 52.351 us | 7.1696 s | +| Shared Q/K prep + paged K/V scatter | 41.312 / 41.312 / 41.600 us | 41.312 us (-21.1%) | 7.1202 s (-0.7%) | + +The common attention kernel stayed within 0.2%, and both profiles launched the same number of kernels. Keep the shared path: the prep-kernel reduction was consistent, while five whole-test runs did not establish a meaningful improvement. This result is specific to this RTX 5090 verifier test and does not establish serving performance. + +- Claim boundary: sampling, calibrated logprob parity, and serving integration remain unverified. + +## Next + +- Add the DFlash drafter and its independent forward oracle. +- Before serving integration, move verifier scratch and recurrent backups to an executor-owned persistent workspace; then wire opt-in fallback/admission rules and collect same-host benchmark evidence. diff --git a/pegainfer-core/src/kv_pool.rs b/pegainfer-core/src/kv_pool.rs index 802a7b4b3..4f704627d 100644 --- a/pegainfer-core/src/kv_pool.rs +++ b/pegainfer-core/src/kv_pool.rs @@ -260,6 +260,21 @@ impl KvState { self.seq_len += count; } + /// Roll this request's logical KV length back to `token_count`, returning + /// any now-unused tail pages to the pool. + pub fn truncate_to(&mut self, token_count: usize) -> Result<()> { + anyhow::ensure!( + token_count <= self.seq_len, + "KvState cannot truncate from {} up to {token_count}", + self.seq_len + ); + let needed = pages_needed(token_count, self.pool.inner.layout.page_size); + self.permit.truncate(needed); + self.seq_len = token_count; + Ok(()) + } + + /// Build kernel-facing metadata for this request's KV. pub fn desc(&self) -> KvDesc<'_> { KvDesc { pages: self.permit.pages(), @@ -389,12 +404,39 @@ mod tests { assert_eq!(desc.last_page_len(), 1); assert_eq!(pool.available_pages(), 2); + // Truncate back into the first page: tail page returns immediately. + kv.truncate_to(15).unwrap(); + assert_eq!(kv.seq_len(), 15); + let desc = kv.desc(); + assert_eq!(desc.num_pages(), 1); + assert_eq!(desc.last_page_len(), 15); + assert_eq!(pool.available_pages(), 3); + + // Truncate to zero releases all request pages. + kv.truncate_to(0).unwrap(); + assert_eq!(kv.seq_len(), 0); + assert_eq!(kv.desc().num_pages(), 0); + assert_eq!(pool.available_pages(), 4); + // Reset returns all pages + kv.ensure_capacity(17).unwrap(); + kv.advance(17); kv.reset(); assert_eq!(kv.seq_len(), 0); assert_eq!(pool.available_pages(), 4); } + #[test] + fn kv_state_rejects_truncate_forward() { + let pool = test_pool(16, 3); + let mut kv = pool.alloc(); + kv.ensure_capacity(4).unwrap(); + kv.advance(4); + + let err = kv.truncate_to(5).unwrap_err().to_string(); + assert!(err.contains("cannot truncate from 4 up to 5")); + } + #[test] fn kv_state_out_of_pages() { // 3 pages total: 1 padding, 2 available → 32 tokens max diff --git a/pegainfer-core/src/ops/attention.rs b/pegainfer-core/src/ops/attention.rs index 1fecc3215..a4ab5da05 100644 --- a/pegainfer-core/src/ops/attention.rs +++ b/pegainfer-core/src/ops/attention.rs @@ -196,10 +196,8 @@ pub fn paged_attention_batch_decode_via_prefill_hd256_into( layout: &KvLayout, layer: usize, plan: &PrefillPagedPlan, - positions_d: &CudaSlice, output: &mut HiddenStates, num_qo_heads: usize, - batch_size: usize, ) -> Result<()> { pegainfer_kernels::ops::paged_attention_batch_decode_via_prefill_hd256_into( ctx, @@ -210,9 +208,7 @@ pub fn paged_attention_batch_decode_via_prefill_hd256_into( &layout.kernel_layout(), layer, plan, - positions_d, output, num_qo_heads, - batch_size, ) } diff --git a/pegainfer-core/src/page_pool.rs b/pegainfer-core/src/page_pool.rs index deb48235f..8e2d09f2b 100644 --- a/pegainfer-core/src/page_pool.rs +++ b/pegainfer-core/src/page_pool.rs @@ -115,6 +115,24 @@ impl OwnedPagePermit { } true } + + /// Return tail pages until the permit holds exactly `new_len` pages. + /// + /// Prefix page order is preserved. Pages beyond `new_len` are returned to + /// the same pool immediately, matching the drop-time LIFO reuse order. + pub(crate) fn truncate(&mut self, new_len: usize) { + assert!( + new_len <= self.pages.len(), + "cannot grow an OwnedPagePermit via truncate" + ); + if new_len == self.pages.len() { + return; + } + + let returned = self.pages.split_off(new_len); + let mut free_list = self.inner.free_list.lock(); + free_list.extend(returned.into_iter().rev()); + } } impl Drop for OwnedPagePermit { @@ -191,4 +209,27 @@ mod tests { // all 4 pages back after drop assert_eq!(pool.available_pages(), 4); } + + #[test] + fn truncate_returns_tail_pages_and_preserves_prefix() { + let pool = PagePool::new(5); + + { + let mut permit = pool.try_acquire_many(4).expect("initial acquire"); + assert_eq!( + permit.pages(), + &[PageId(0), PageId(1), PageId(2), PageId(3)] + ); + assert_eq!(pool.available_pages(), 1); + + permit.truncate(2); + assert_eq!(permit.pages(), &[PageId(0), PageId(1)]); + assert_eq!(pool.available_pages(), 3); + + let next = pool.try_acquire_many(2).expect("tail pages reusable"); + assert_eq!(next.pages(), &[PageId(2), PageId(3)]); + } + + assert_eq!(pool.available_pages(), 5); + } } diff --git a/pegainfer-kernels/csrc/qwen35/prefill_attention_hd256.cu b/pegainfer-kernels/csrc/qwen35/prefill_attention_hd256.cu index cc4850884..82a8dc5d2 100644 --- a/pegainfer-kernels/csrc/qwen35/prefill_attention_hd256.cu +++ b/pegainfer-kernels/csrc/qwen35/prefill_attention_hd256.cu @@ -291,7 +291,7 @@ __global__ void qk_norm_partial_rope_batched_decode_hd256_kernel( extern "C" { -void qk_norm_partial_rope_batched_decode_hd256_cuda( +int qk_norm_partial_rope_batched_decode_hd256_cuda( const __nv_bfloat16* q_full_batch, __nv_bfloat16* k_batch, const __nv_bfloat16* q_norm_weight, @@ -323,6 +323,7 @@ void qk_norm_partial_rope_batched_decode_hd256_cuda( rotary_dim, rms_eps ); + return static_cast(cudaGetLastError()); } void prefill_attention_hd256_prep_paged_cuda( diff --git a/pegainfer-kernels/src/ffi/qwen35.rs b/pegainfer-kernels/src/ffi/qwen35.rs index c949dc186..9b2e5d5e9 100644 --- a/pegainfer-kernels/src/ffi/qwen35.rs +++ b/pegainfer-kernels/src/ffi/qwen35.rs @@ -57,7 +57,7 @@ unsafe extern "C" { rotary_dim: i32, rms_eps: f32, stream: CUstream, - ); + ) -> i32; // Gated delta rule recurrent decode (single step) pub fn gated_delta_rule_decode_cuda( diff --git a/pegainfer-kernels/src/ops.rs b/pegainfer-kernels/src/ops.rs index 9590cdcbc..410372db5 100644 --- a/pegainfer-kernels/src/ops.rs +++ b/pegainfer-kernels/src/ops.rs @@ -138,6 +138,7 @@ pub use linear::gemm_lt_pin_tune; pub use linear::gemm_lt_pin_warmup; pub use linear::gemm_lt_tune; pub use linear::gemm_per_token; +pub use linear::gemm_per_token_into_checked; pub use linear::gemm_rows_into; pub use linear::gemm_rows_into_checked; pub use linear::gemm_strided_batched_bf16; @@ -149,6 +150,7 @@ pub use linear::per_token_served; pub use linear::pin_served; pub use linear::reset_numeric_policy_counters; pub use linear::set_numeric_policy; +pub use linear::with_gemm_lt_disabled; pub use lora::LoraDecodeGroupedProjection; pub use lora::lora_decode_fused_delta_group3_into; pub use lora::lora_decode_fused_delta_into; diff --git a/pegainfer-kernels/src/ops/attention.rs b/pegainfer-kernels/src/ops/attention.rs index a5ac890ec..7824037a8 100644 --- a/pegainfer-kernels/src/ops/attention.rs +++ b/pegainfer-kernels/src/ops/attention.rs @@ -1175,10 +1175,26 @@ pub fn qk_norm_partial_rope_batched_decode_hd256_into( num_kv_heads: usize, rotary_dim: usize, rms_eps: f32, -) { +) -> Result<()> { let batch_size = q.seq_len; - debug_assert_eq!(q_full.seq_len, batch_size); - debug_assert_eq!(k.seq_len, batch_size); + anyhow::ensure!(batch_size > 0, "Qwen3.5 QK prep requires at least one row"); + anyhow::ensure!( + q_full.seq_len == batch_size && k.seq_len == batch_size, + "Qwen3.5 QK prep row mismatch: q_full={}, q={}, k={}", + q_full.seq_len, + batch_size, + k.seq_len + ); + anyhow::ensure!( + positions_d.len() >= batch_size, + "Qwen3.5 QK prep positions too short: rows={batch_size}, positions={}", + positions_d.len() + ); + anyhow::ensure!( + u16::try_from(batch_size).is_ok(), + "Qwen3.5 QK prep row count {batch_size} exceeds CUDA grid.y limit {}", + u16::MAX + ); let (qf_ptr, _gqf) = q_full.data.device_ptr(&ctx.stream); let (q_ptr, _gq) = q.data.device_ptr_mut(&ctx.stream); @@ -1189,7 +1205,7 @@ pub fn qk_norm_partial_rope_batched_decode_hd256_into( let (sin_ptr, _gs) = sin_cache.data.device_ptr(&ctx.stream); let (pos_ptr, _gp) = positions_d.device_ptr(&ctx.stream); - unsafe { + let result = unsafe { ffi::qk_norm_partial_rope_batched_decode_hd256_cuda( qf_ptr as *const ffi::Half, k_ptr as *mut ffi::Half, @@ -1205,8 +1221,15 @@ pub fn qk_norm_partial_rope_batched_decode_hd256_into( rotary_dim as i32, rms_eps, crate::tensor::active_cu_stream(ctx), + ) + }; + if result != 0 { + anyhow::bail!( + "qk_norm_partial_rope_batched_decode_hd256_cuda failed with error {result}{}", + crate::ops::ffi_exception_message(result) ); } + Ok(()) } /// Batched paged attention decode: append K/V + FlashInfer BatchDecode for batch_size >= 1. @@ -1486,7 +1509,7 @@ pub fn paged_attention_batch_decode_split_kv_into( } #[allow(clippy::too_many_arguments)] -fn scatter_decode_kv_into_paged( +fn scatter_kv_into_paged( ctx: &DeviceContext, k: &HiddenStates, v: &HiddenStates, @@ -1498,7 +1521,7 @@ fn scatter_decode_kv_into_paged( last_page_len_d: &CudaSlice, positions_d: &CudaSlice, request_indices_d: &CudaSlice, - batch_size: usize, + total_tokens: usize, op_name: &str, ) -> Result<()> { let num_kv_heads = layout.num_kv_heads; @@ -1532,7 +1555,7 @@ fn scatter_decode_kv_into_paged( v_ptr as *const ffi::Half, ri_ptr as *const i32, pos_ptr as *const i32, - batch_size as i32, + total_tokens as i32, num_kv_heads as i32, head_dim as i32, page_size as i32, @@ -1544,7 +1567,7 @@ fn scatter_decode_kv_into_paged( }; if result != 0 { anyhow::bail!( - "paged_kv_scatter_cuda ({op_name}) failed for layer {layer}, bs={batch_size}, \ + "paged_kv_scatter_cuda ({op_name}) failed for layer {layer}, rows={total_tokens}, \ kv_heads={num_kv_heads}, head_dim={head_dim}, page_size={page_size}: {result}{}", crate::ops::ffi_exception_message(result) ); @@ -1574,7 +1597,15 @@ pub fn paged_attention_batch_decode_hd256_into( ) -> Result<()> { let num_kv_heads = layout.num_kv_heads; let head_dim = layout.head_dim; - debug_assert_eq!(head_dim, 256); + anyhow::ensure!( + head_dim == 256, + "batch HD256 decode requires head_dim=256, got {head_dim}" + ); + anyhow::ensure!( + layer < layout.num_layers, + "batch HD256 decode layer {layer} out of range for {} layers", + layout.num_layers + ); let page_size = layout.page_size; let k_offset = (layer * layout.layer_stride) as i64; @@ -1593,7 +1624,7 @@ pub fn paged_attention_batch_decode_hd256_into( let stream = crate::tensor::active_cu_stream(ctx); - scatter_decode_kv_into_paged( + scatter_kv_into_paged( ctx, k, v, @@ -1653,22 +1684,35 @@ pub fn paged_attention_batch_decode_via_prefill_hd256_into( layout: &PagedKvLayout, layer: usize, plan: &PrefillPagedPlan, - positions_d: &CudaSlice, output: &mut HiddenStates, num_qo_heads: usize, - batch_size: usize, ) -> Result<()> { let num_kv_heads = layout.num_kv_heads; let head_dim = layout.head_dim; - debug_assert_eq!(head_dim, 256); anyhow::ensure!( - batch_size == plan.total_tokens && batch_size == plan.batch_size as usize, - "decode-via-prefill plan shape mismatch: bs={batch_size}, total_tokens={}, plan_batch={}", - plan.total_tokens, - plan.batch_size + head_dim == 256, + "batch HD256 prefill requires head_dim=256, got {head_dim}" + ); + anyhow::ensure!( + layer < layout.num_layers, + "batch HD256 prefill layer {layer} out of range for {} layers", + layout.num_layers + ); + let total_tokens = q.seq_len; + anyhow::ensure!( + total_tokens > 0 + && k.seq_len == total_tokens + && v.seq_len == total_tokens + && output.seq_len == total_tokens + && total_tokens == plan.total_tokens, + "batch-prefill row mismatch: q={total_tokens}, k={}, v={}, output={}, plan={}", + k.seq_len, + v.seq_len, + output.seq_len, + plan.total_tokens ); - scatter_decode_kv_into_paged( + scatter_kv_into_paged( ctx, k, v, @@ -1678,9 +1722,9 @@ pub fn paged_attention_batch_decode_via_prefill_hd256_into( &plan.page_indices_d, &plan.page_indptr_d, &plan.last_page_len_d, - positions_d, + &plan.positions_d, &plan.batch_indices_d, - batch_size, + total_tokens, "batch hd256 decode via prefill", )?; @@ -1722,7 +1766,7 @@ pub fn paged_attention_batch_decode_via_prefill_hd256_into( num_kv_heads as i32, head_dim as i32, layout.page_size as i32, - batch_size as i32, + total_tokens as i32, plan.batch_size, plan.num_tiles, stride_page, @@ -1732,8 +1776,9 @@ pub fn paged_attention_batch_decode_via_prefill_hd256_into( }; if result != 0 { anyhow::bail!( - "batch_prefill_paged_cuda_hd256 (decode via prefill) failed for layer {layer}, \ - bs={batch_size}, tiles={}, qo_heads={num_qo_heads}, kv_heads={num_kv_heads}: {result}{}", + "batch_prefill_paged_cuda_hd256 failed for layer {layer}, rows={total_tokens}, \ + batch={}, tiles={}, qo_heads={num_qo_heads}, kv_heads={num_kv_heads}: {result}{}", + plan.batch_size, plan.num_tiles, crate::ops::ffi_exception_message(result) ); @@ -2108,7 +2153,7 @@ pub fn paged_attention_batch_decode_via_prefill_hd512_into( kernel_cta_tile_q ); - scatter_decode_kv_into_paged( + scatter_kv_into_paged( ctx, k, v, diff --git a/pegainfer-kernels/src/ops/linear.rs b/pegainfer-kernels/src/ops/linear.rs index c3013909b..d6aa92af4 100644 --- a/pegainfer-kernels/src/ops/linear.rs +++ b/pegainfer-kernels/src/ops/linear.rs @@ -1,3 +1,4 @@ +use std::cell::Cell; use std::sync::atomic::AtomicU8; use std::sync::atomic::AtomicU64; use std::sync::atomic::Ordering; @@ -521,7 +522,7 @@ pub fn gemm_graphsafe_ref_into_checked( gemm_ref_into_with_policy(ctx, weight, x, out, true) } -fn gemm_per_token_into_checked( +pub fn gemm_per_token_into_checked( ctx: &DeviceContext, weight: &DeviceMatrix, x: &HiddenStates, @@ -642,6 +643,46 @@ pub enum NumericPolicy { PerToken = 2, } +thread_local! { + static GEMM_LT_DISABLE_DEPTH: Cell = const { Cell::new(0) }; +} + +struct GemmLtDisableGuard; + +impl Drop for GemmLtDisableGuard { + fn drop(&mut self) { + GEMM_LT_DISABLE_DEPTH.with(|depth| { + depth.set( + depth + .get() + .checked_sub(1) + .expect("unbalanced GEMM Lt guard"), + ); + }); + } +} + +/// Run `f` with timing-tuned cublasLt plans disabled on the current thread. +/// +/// Transaction replay uses the stable cuBLAS fallback so its canonical state +/// does not depend on which near-equal small-N algorithm won startup timing. +pub fn with_gemm_lt_disabled(f: impl FnOnce() -> T) -> T { + GEMM_LT_DISABLE_DEPTH.with(|depth| { + depth.set( + depth + .get() + .checked_add(1) + .expect("GEMM Lt disable depth overflow"), + ); + }); + let _guard = GemmLtDisableGuard; + f() +} + +fn gemm_lt_disabled() -> bool { + GEMM_LT_DISABLE_DEPTH.with(|depth| depth.get() != 0) +} + static NUMERIC_POLICY: AtomicU8 = AtomicU8::new(NumericPolicy::Tuned as u8); static PIN_SERVED: AtomicU64 = AtomicU64::new(0); static PER_TOKEN_SERVED: AtomicU64 = AtomicU64::new(0); @@ -811,19 +852,20 @@ fn launch_gemm( // NOTE: gemm_lt is disabled when a stream override is active (SM-partition // concurrent mode). cuBLASLt has device-global state that conflicts when // two green-ctx streams run cublasLtMatmul concurrently, causing Xid 31. - let mut status = if n <= GEMM_LT_MAX_N && !crate::tensor::has_stream_override() { - ffi::gemm_lt_cuda( - w_ptr, - x_ptr, - y_ptr, - m as i32, - n as i32, - k as i32, - crate::tensor::active_cu_stream(ctx), - ) - } else { - GEMM_LT_UNTUNED - }; + let mut status = + if n <= GEMM_LT_MAX_N && !crate::tensor::has_stream_override() && !gemm_lt_disabled() { + ffi::gemm_lt_cuda( + w_ptr, + x_ptr, + y_ptr, + m as i32, + n as i32, + k as i32, + crate::tensor::active_cu_stream(ctx), + ) + } else { + GEMM_LT_UNTUNED + }; if status == GEMM_LT_UNTUNED { status = if graphsafe { ffi::gemm_graphsafe_cuda( diff --git a/pegainfer-qwen35/src/batch_decode.rs b/pegainfer-qwen35/src/batch_decode.rs index a778d3781..18c9e8f2a 100644 --- a/pegainfer-qwen35/src/batch_decode.rs +++ b/pegainfer-qwen35/src/batch_decode.rs @@ -124,7 +124,7 @@ impl Qwen35Model { num_key_value_heads, self.config.rotary_dim, eps, - ); + )?; ops::paged_attention_batch_decode_hd256_into( &self.ctx, @@ -198,7 +198,7 @@ impl Qwen35Model { self.config.num_key_value_heads, self.config.rotary_dim, eps, - ); + )?; ops::paged_attention_batch_decode_via_prefill_hd256_into( &self.ctx, @@ -209,10 +209,8 @@ impl Qwen35Model { layout, layer_idx, plan, - &bufs.positions_d, &mut bufs.attn_out_full, self.config.num_attention_heads, - bs, )?; unsafe { diff --git a/pegainfer-qwen35/src/batch_decode_graph.rs b/pegainfer-qwen35/src/batch_decode_graph.rs index 27bcbdcab..268594189 100644 --- a/pegainfer-qwen35/src/batch_decode_graph.rs +++ b/pegainfer-qwen35/src/batch_decode_graph.rs @@ -116,16 +116,74 @@ impl BatchDecodeGraphState { src: &RecurrentState, slot_idx: usize, ) -> Result<()> { + anyhow::ensure!( + slot_idx < self.slot_states.len(), + "Qwen3.5 graph slot {slot_idx} exceeds capacity {}", + self.slot_states.len() + ); let dst = &mut self.slot_states[slot_idx]; - for (dst_layer, src_layer) in dst.layers.iter_mut().zip(src.layers.iter()) { - ctx.stream - .memcpy_dtod(&src_layer.state, &mut dst_layer.state) - .map_err(|e| anyhow::anyhow!("copy recurrent state to slot {slot_idx}: {e}"))?; - ctx.stream - .memcpy_dtod(&src_layer.conv_state.data, &mut dst_layer.conv_state.data) - .map_err(|e| anyhow::anyhow!("copy conv state to slot {slot_idx}: {e}"))?; - } - dst.seq_len = src.seq_len; + dst.copy_from(ctx, src) + .map_err(|e| anyhow::anyhow!("copy recurrent state to slot {slot_idx}: {e}"))?; + Ok(()) + } + + /// D2D copy slot `slot_idx` recurrent state into a standalone state. + #[cfg_attr( + not(test), + expect( + dead_code, + reason = "crate-private verifier substrate uses this before serving wiring" + ) + )] + pub(crate) fn copy_slot_to_state( + &self, + ctx: &DeviceContext, + slot_idx: usize, + dst: &mut RecurrentState, + ) -> Result<()> { + anyhow::ensure!( + slot_idx < self.slot_states.len(), + "Qwen3.5 graph slot {slot_idx} exceeds capacity {}", + self.slot_states.len() + ); + dst.copy_from(ctx, &self.slot_states[slot_idx]) + .map_err(|e| anyhow::anyhow!("copy recurrent slot {slot_idx} to state: {e}"))?; Ok(()) } + + /// D2D copy one graph slot's recurrent/conv state into another slot. + pub(crate) fn copy_slot_to_slot( + &mut self, + ctx: &DeviceContext, + src_slot_idx: usize, + dst_slot_idx: usize, + ) -> Result<()> { + anyhow::ensure!( + src_slot_idx < self.slot_states.len(), + "Qwen3.5 recurrent source slot {src_slot_idx} out of range {}", + self.slot_states.len() + ); + anyhow::ensure!( + dst_slot_idx < self.slot_states.len(), + "Qwen3.5 recurrent destination slot {dst_slot_idx} out of range {}", + self.slot_states.len() + ); + if src_slot_idx == dst_slot_idx { + return Ok(()); + } + if src_slot_idx < dst_slot_idx { + let (left, right) = self.slot_states.split_at_mut(dst_slot_idx); + let src = &left[src_slot_idx]; + let dst = &mut right[0]; + dst.copy_from(ctx, src) + } else { + let (left, right) = self.slot_states.split_at_mut(src_slot_idx); + let dst = &mut left[dst_slot_idx]; + let src = &right[0]; + dst.copy_from(ctx, src) + } + .map_err(|e| { + anyhow::anyhow!("copy Qwen3.5 recurrent slot {src_slot_idx} to {dst_slot_idx}: {e}") + }) + } } diff --git a/pegainfer-qwen35/src/executor.rs b/pegainfer-qwen35/src/executor.rs index 60e021330..f21d3655a 100644 --- a/pegainfer-qwen35/src/executor.rs +++ b/pegainfer-qwen35/src/executor.rs @@ -99,16 +99,16 @@ pub struct DecodeResult { pub requests: Vec, } -struct ActiveRequest { - request_id: RequestId, - kv: KvState, - graph_slot_idx: usize, +pub(crate) struct ActiveRequest { + pub(crate) request_id: RequestId, + pub(crate) kv: KvState, + pub(crate) graph_slot_idx: usize, } pub struct Qwen35Executor { - model: Qwen35Model, - graph_state: BatchDecodeGraphState, - active: Vec, + pub(crate) model: Qwen35Model, + pub(crate) graph_state: BatchDecodeGraphState, + pub(crate) active: Vec, } impl Qwen35Executor { @@ -281,40 +281,36 @@ impl Qwen35Executor { self.active[idx].graph_slot_idx, last ); - for layer_idx in 0..self.graph_state.slot_states[last].layers.len() { - let (src_part, dst_part) = if idx < last { - let (left, right) = self.graph_state.slot_states.split_at_mut(last); - ( - &right[0].layers[layer_idx], - &mut left[idx].layers[layer_idx], - ) - } else { - unreachable!("idx < active.len() <= last"); - }; - self.model - .device_ctx() - .stream - .memcpy_dtod(&src_part.state, &mut dst_part.state) - .map_err(|e| { - anyhow::anyhow!("compact Qwen3.5 logits executor state copy failed: {e}") - })?; - self.model - .device_ctx() - .stream - .memcpy_dtod(&src_part.conv_state.data, &mut dst_part.conv_state.data) - .map_err(|e| { - anyhow::anyhow!( - "compact Qwen3.5 logits executor conv_state copy failed: {e}" - ) - })?; - } - self.graph_state.slot_states[idx].seq_len = self.graph_state.slot_states[last].seq_len; + self.graph_state + .copy_slot_to_slot(self.model.device_ctx(), last, idx)?; self.active[idx].graph_slot_idx = idx; } Ok(()) } } +#[cfg(test)] +#[derive(Clone, Debug, Eq, PartialEq)] +pub(crate) struct ExecutorStateSummary { + pub(crate) request_id: RequestId, + pub(crate) kv_seq_len: usize, + pub(crate) recurrent_seq_len: usize, +} + +#[cfg(test)] +impl Qwen35Executor { + pub(crate) fn debug_state_summary(&self) -> Vec { + self.active + .iter() + .map(|active| ExecutorStateSummary { + request_id: active.request_id, + kv_seq_len: active.kv.seq_len(), + recurrent_seq_len: self.graph_state.slot_states[active.graph_slot_idx].seq_len, + }) + .collect() + } +} + fn select_default_tokens_from_logits( model: &Qwen35Model, logits: &HiddenStates, diff --git a/pegainfer-qwen35/src/lib.rs b/pegainfer-qwen35/src/lib.rs index 29ea7bdbe..fb718d90d 100644 --- a/pegainfer-qwen35/src/lib.rs +++ b/pegainfer-qwen35/src/lib.rs @@ -15,11 +15,37 @@ pub mod model_line; mod ops; mod prefill; pub mod prefill_buffers; +#[cfg_attr( + not(test), + expect( + dead_code, + reason = "crate-private verifier substrate is exercised by GPU tests before serving wiring" + ) +)] +mod prefill_verify; pub(crate) mod recurrent; pub(crate) mod recurrent_state; mod scheduler; +#[cfg_attr( + not(test), + expect( + dead_code, + reason = "crate-private verifier substrate is exercised by GPU tests before serving wiring" + ) +)] +mod speculative; +#[cfg(test)] +mod speculative_tests; mod tp_executor; mod unified_forward; +#[cfg_attr( + not(test), + expect( + dead_code, + reason = "crate-private verifier substrate is exercised by GPU tests before serving wiring" + ) +)] +mod verify_buffers; mod weights; use std::path::Path; diff --git a/pegainfer-qwen35/src/ops.rs b/pegainfer-qwen35/src/ops.rs index 3d288b074..cdf1f9502 100644 --- a/pegainfer-qwen35/src/ops.rs +++ b/pegainfer-qwen35/src/ops.rs @@ -4,10 +4,12 @@ pub(crate) use pegainfer_core::ops::GEMM_LT_MAX_N; pub(crate) use pegainfer_core::ops::PrefillPagedPlan; pub(crate) use pegainfer_core::ops::add_batch; pub(crate) use pegainfer_core::ops::add_batch_into; +pub(crate) use pegainfer_core::ops::copy_hidden_token_range_into; pub(crate) use pegainfer_core::ops::embedding_batch; pub(crate) use pegainfer_core::ops::extract_vec; pub(crate) use pegainfer_core::ops::gemm; pub(crate) use pegainfer_core::ops::gemm_into; +pub(crate) use pegainfer_core::ops::gemm_into_checked; pub(crate) use pegainfer_core::ops::gemm_lt_tune; pub(crate) use pegainfer_core::ops::gemm_rows_into_checked; pub(crate) use pegainfer_core::ops::paged_attention_batch_decode_hd256_into; diff --git a/pegainfer-qwen35/src/prefill.rs b/pegainfer-qwen35/src/prefill.rs index f641983a9..8402ffc9a 100644 --- a/pegainfer-qwen35/src/prefill.rs +++ b/pegainfer-qwen35/src/prefill.rs @@ -320,17 +320,17 @@ impl Qwen35Model { // Step 1: QK norm + partial RoPE + direct paged K/V write. unsafe { - let (qf_ptr, _) = q_full_batch.data.device_ptr(&self.ctx.stream); - let (k_ptr, _) = k_batch.data.device_ptr(&self.ctx.stream); - let (v_ptr, _) = v_batch.data.device_ptr(&self.ctx.stream); - let (qn_ptr, _) = attn.q_norm.data.device_ptr(&self.ctx.stream); - let (kn_ptr, _) = attn.k_norm.data.device_ptr(&self.ctx.stream); - let (cos_ptr, _) = self.cos_cache.data.device_ptr(&self.ctx.stream); - let (sin_ptr, _) = self.sin_cache.data.device_ptr(&self.ctx.stream); - let (qp_ptr, _) = q_prepped.data.device_ptr_mut(&self.ctx.stream); - let (buf_ptr, _) = kv_state.buffer().device_ptr(&self.ctx.stream); - let (pi_ptr, _) = prefill_plan.page_indices_d().device_ptr(&self.ctx.stream); - let (sp_ptr, _) = start_pos_cpu.device_ptr(&self.ctx.stream); + let (qf_ptr, _gqf) = q_full_batch.data.device_ptr(&self.ctx.stream); + let (k_ptr, _gk) = k_batch.data.device_ptr(&self.ctx.stream); + let (v_ptr, _gv) = v_batch.data.device_ptr(&self.ctx.stream); + let (qn_ptr, _gqn) = attn.q_norm.data.device_ptr(&self.ctx.stream); + let (kn_ptr, _gkn) = attn.k_norm.data.device_ptr(&self.ctx.stream); + let (cos_ptr, _gcos) = self.cos_cache.data.device_ptr(&self.ctx.stream); + let (sin_ptr, _gsin) = self.sin_cache.data.device_ptr(&self.ctx.stream); + let (qp_ptr, _gqp) = q_prepped.data.device_ptr_mut(&self.ctx.stream); + let (buf_ptr, _gkv) = kv_state.buffer().device_ptr(&self.ctx.stream); + let (pi_ptr, _gpi) = prefill_plan.page_indices_d().device_ptr(&self.ctx.stream); + let (sp_ptr, _gsp) = start_pos_cpu.device_ptr(&self.ctx.stream); ffi::prefill_attention_hd256_prep_paged_cuda( qf_ptr as *const ffi::Half, k_ptr as *const ffi::Half, diff --git a/pegainfer-qwen35/src/prefill_buffers.rs b/pegainfer-qwen35/src/prefill_buffers.rs index 92b38eb8a..ccdaeb0cb 100644 --- a/pegainfer-qwen35/src/prefill_buffers.rs +++ b/pegainfer-qwen35/src/prefill_buffers.rs @@ -45,6 +45,14 @@ pub struct GdrChunkwiseScratch35 { /// Per-chunk recurrent state snapshots, fp32: [num_chunks, num_value_heads, key_dim, value_dim] pub(crate) chunk_state: CudaSlice, + #[cfg_attr( + not(test), + expect( + dead_code, + reason = "crate-private verifier reuses this scratch before serving wiring" + ) + )] + max_seq_len: usize, } impl GdrChunkwiseScratch35 { @@ -104,9 +112,31 @@ impl GdrChunkwiseScratch35 { u: HiddenStates::zeros(ctx, vv_hidden_dim, seq_len)?, v_new: HiddenStates::zeros(ctx, vv_hidden_dim, seq_len)?, chunk_state, + max_seq_len: seq_len, }) } + #[cfg_attr( + not(test), + expect( + dead_code, + reason = "crate-private verifier reuses this scratch before serving wiring" + ) + )] + pub(crate) fn set_rows(&mut self, seq_len: usize) { + assert!( + seq_len <= self.max_seq_len, + "Qwen3.5 GDR scratch rows {seq_len} exceeds capacity {}", + self.max_seq_len + ); + self.q_expanded.seq_len = seq_len; + self.k_expanded.seq_len = seq_len; + self.v_raw.seq_len = seq_len; + self.w.seq_len = seq_len; + self.u.seq_len = seq_len; + self.v_new.seq_len = seq_len; + } + pub(crate) fn num_chunks(seq_len: usize) -> usize { seq_len.div_ceil(Self::CHUNK_SIZE) } diff --git a/pegainfer-qwen35/src/prefill_verify.rs b/pegainfer-qwen35/src/prefill_verify.rs new file mode 100644 index 000000000..1bf960b29 --- /dev/null +++ b/pegainfer-qwen35/src/prefill_verify.rs @@ -0,0 +1,367 @@ +use anyhow::Result; +use cudarc::driver::DevicePtr; +use cudarc::driver::DevicePtrMut; +use pegainfer_core::kv_pool::KvState; + +use crate::ffi; +use crate::ops; +use crate::prefill::PREFILL_CHUNK_LEN; +use crate::recurrent_state::RecurrentState; +use crate::verify_buffers::VerifyBuffers35; +use crate::weights::FullAttentionLayer; +use crate::weights::LayerKind; +use crate::weights::LinearAttentionLayer; +use crate::weights::Qwen35Model; +use crate::weights::TransformerBlock35; + +impl Qwen35Model { + pub(crate) fn prefill_verify_into( + &self, + spans: &[&[u32]], + kv_states: &mut [&mut KvState], + recurrent_states: &mut [&mut RecurrentState], + bufs: &mut VerifyBuffers35, + ) -> Result<()> { + anyhow::ensure!( + !pegainfer_kernels::tensor::has_stream_override(), + "Qwen3.5 verify prefill does not support a CUDA stream override" + ); + anyhow::ensure!(!spans.is_empty(), "Qwen3.5 verify needs at least one span"); + anyhow::ensure!( + spans.len() == kv_states.len() && spans.len() == recurrent_states.len(), + "Qwen3.5 verify spans/KV/recurrent mismatch: spans={}, kv={}, recurrent={}", + spans.len(), + kv_states.len(), + recurrent_states.len() + ); + let kv_buffer = kv_states[0].buffer(); + anyhow::ensure!( + kv_states + .iter() + .all(|kv| std::ptr::eq(kv.buffer(), kv_buffer)), + "Qwen3.5 verify KV states must share one pool" + ); + anyhow::ensure!( + spans.len() <= bufs.max_batch(), + "Qwen3.5 verify batch {} exceeds buffer capacity {}", + spans.len(), + bufs.max_batch() + ); + for span in spans { + anyhow::ensure!( + !span.is_empty() && span.len() <= PREFILL_CHUNK_LEN, + "Qwen3.5 verify span len {} out of range", + span.len() + ); + } + + let total_rows = bufs.stage_tokens(&self.ctx, spans)?; + let seq_lens: Vec = spans.iter().map(|span| span.len()).collect(); + let start_positions: Vec = kv_states.iter().map(|kv| kv.seq_len()).collect(); + for (kv, (&base_pos, &seq_len)) in kv_states + .iter_mut() + .zip(start_positions.iter().zip(seq_lens.iter())) + { + let end_pos = base_pos.checked_add(seq_len).ok_or_else(|| { + anyhow::anyhow!( + "Qwen3.5 verify position overflow: base_pos={base_pos}, seq_len={seq_len}" + ) + })?; + anyhow::ensure!( + end_pos <= self.config.max_position_embeddings, + "Qwen3.5 verify requested end_pos={end_pos}, beyond max_position_embeddings={}", + self.config.max_position_embeddings + ); + self.ensure_rope_cache_covers(end_pos)?; + kv.ensure_capacity(end_pos)?; + kv.advance(seq_len); + } + + let page_indices: Vec> = + kv_states.iter().map(|kv| kv.page_indices_i32()).collect(); + let last_page_lens: Vec = kv_states.iter().map(|kv| kv.last_page_len()).collect(); + bufs.plan.update_batch_with_cta_tile_q( + &self.ctx, + &page_indices, + &last_page_lens, + &start_positions, + &seq_lens, + self.config.num_attention_heads, + self.config.num_key_value_heads, + self.config.head_dim, + 0, + )?; + + ops::embedding_batch( + &self.ctx, + &self.embed_tokens, + &bufs.token_ids_d, + &mut bufs.hidden, + )?; + + let mut linear_idx = 0usize; + let mut full_idx = 0usize; + for layer in &self.layers { + self.prefill_verify_layer_into( + layer, + &seq_lens, + kv_states, + recurrent_states, + &mut linear_idx, + &mut full_idx, + bufs, + )?; + } + + for (recurrent, &seq_len) in recurrent_states.iter_mut().zip(seq_lens.iter()) { + recurrent.seq_len += seq_len; + } + + ops::rms_norm_batch_offset_into( + &self.ctx, + &bufs.hidden, + &self.norm, + self.config.rms_norm_eps, + &mut bufs.logits_normed, + )?; + ops::gemm_rows_into_checked( + &self.ctx, + self.output_projection(), + 0, + self.config.selection_vocab, + &bufs.logits_normed, + &mut bufs.logits, + )?; + debug_assert_eq!(bufs.logits.seq_len, total_rows); + Ok(()) + } + + #[allow(clippy::too_many_arguments)] + fn prefill_verify_layer_into( + &self, + layer: &TransformerBlock35, + seq_lens: &[usize], + kv_states: &[&mut KvState], + recurrent_states: &mut [&mut RecurrentState], + linear_idx: &mut usize, + full_idx: &mut usize, + bufs: &mut VerifyBuffers35, + ) -> Result<()> { + ops::rms_norm_batch_offset_into( + &self.ctx, + &bufs.hidden, + &layer.input_layernorm, + self.config.rms_norm_eps, + &mut bufs.normed, + )?; + + match &layer.attn { + LayerKind::FullAttention(attn) => { + self.prefill_verify_full_attention_into(attn, kv_states, *full_idx, bufs)?; + *full_idx += 1; + } + LayerKind::LinearAttention(attn) => { + self.prefill_verify_linear_attention_into( + attn, + seq_lens, + recurrent_states, + *linear_idx, + bufs, + )?; + *linear_idx += 1; + } + } + + ops::add_batch_into( + &self.ctx, + &bufs.hidden, + &bufs.attn_results, + &mut bufs.hidden_mid, + )?; + ops::rms_norm_batch_offset_into( + &self.ctx, + &bufs.hidden_mid, + &layer.post_attention_layernorm, + self.config.rms_norm_eps, + &mut bufs.normed, + )?; + ops::gemm_into_checked( + &self.ctx, + &layer.mlp.gate_up_proj, + &bufs.normed, + &mut bufs.gate_up_out, + )?; + ops::silu_mul_fused_batch_into(&self.ctx, &bufs.gate_up_out, &mut bufs.act_out)?; + ops::gemm_into_checked( + &self.ctx, + &layer.mlp.down_proj, + &bufs.act_out, + &mut bufs.mlp_out, + )?; + ops::add_batch_into( + &self.ctx, + &bufs.hidden_mid, + &bufs.mlp_out, + &mut bufs.hidden_next, + )?; + std::mem::swap(&mut bufs.hidden, &mut bufs.hidden_next); + Ok(()) + } + + fn prefill_verify_full_attention_into( + &self, + attn: &FullAttentionLayer, + kv_states: &[&mut KvState], + full_idx: usize, + bufs: &mut VerifyBuffers35, + ) -> Result<()> { + let c = &self.config; + ops::gemm_into_checked(&self.ctx, &attn.q_proj, &bufs.normed, &mut bufs.q_full)?; + ops::gemm_into_checked(&self.ctx, &attn.k_proj, &bufs.normed, &mut bufs.k_full)?; + ops::gemm_into_checked(&self.ctx, &attn.v_proj, &bufs.normed, &mut bufs.v_full)?; + + ops::qk_norm_partial_rope_batched_decode_hd256_into( + &self.ctx, + &bufs.q_full, + &mut bufs.q_prepped, + &mut bufs.k_full, + &attn.q_norm, + &attn.k_norm, + &self.cos_cache, + &self.sin_cache, + bufs.plan.positions_d(), + c.num_attention_heads, + c.num_key_value_heads, + c.rotary_dim, + c.rms_norm_eps, + )?; + + ops::paged_attention_batch_decode_via_prefill_hd256_into( + &self.ctx, + &bufs.q_prepped, + &bufs.k_full, + &bufs.v_full, + kv_states[0].buffer(), + kv_states[0].layout(), + full_idx, + &bufs.plan, + &mut bufs.attn_out_full, + c.num_attention_heads, + )?; + + unsafe { + let (qf_ptr, _gqf) = bufs.q_full.data.device_ptr(&self.ctx.stream); + let (out_ptr, _go) = bufs.attn_out_full.data.device_ptr_mut(&self.ctx.stream); + ffi::attention_gate_batch_hd256_cuda( + qf_ptr as *const ffi::Half, + out_ptr as *mut ffi::Half, + c.num_attention_heads as i32, + bufs.q_prepped.seq_len as i32, + self.ctx.stream.cu_stream(), + ); + } + ops::gemm_into_checked( + &self.ctx, + &attn.o_proj, + &bufs.attn_out_full, + &mut bufs.attn_results, + )?; + Ok(()) + } + + fn prefill_verify_linear_attention_into( + &self, + attn: &LinearAttentionLayer, + seq_lens: &[usize], + recurrent_states: &mut [&mut RecurrentState], + linear_idx: usize, + bufs: &mut VerifyBuffers35, + ) -> Result<()> { + let c = &self.config; + ops::gemm_into_checked(&self.ctx, &attn.in_proj_qkv, &bufs.normed, &mut bufs.qkv)?; + ops::gemm_into_checked(&self.ctx, &attn.in_proj_z, &bufs.normed, &mut bufs.z)?; + ops::gemm_into_checked(&self.ctx, &attn.in_proj_b, &bufs.normed, &mut bufs.b_proj)?; + ops::gemm_into_checked(&self.ctx, &attn.in_proj_a, &bufs.normed, &mut bufs.a_proj)?; + + let mut row_offset = 0usize; + for (recurrent, &seq_len) in recurrent_states.iter_mut().zip(seq_lens.iter()) { + let layer_state = &mut recurrent.layers[linear_idx]; + bufs.set_compact_rows(seq_len); + ops::copy_hidden_token_range_into( + &self.ctx, + &bufs.qkv, + row_offset, + &mut bufs.compact_qkv, + 0, + seq_len, + )?; + ops::conv1d_prefill_batch_into( + &self.ctx, + &bufs.compact_qkv, + &attn.conv1d_weight, + &mut layer_state.conv_state, + &mut bufs.compact_qkv_conv, + c.linear_conv_kernel_dim, + ); + ops::copy_hidden_token_range_into( + &self.ctx, + &bufs.b_proj, + row_offset, + &mut bufs.compact_b, + 0, + seq_len, + )?; + ops::copy_hidden_token_range_into( + &self.ctx, + &bufs.a_proj, + row_offset, + &mut bufs.compact_a, + 0, + seq_len, + )?; + ops::gated_delta_rule_prefill_chunkwise_into( + &self.ctx, + &bufs.compact_qkv_conv, + &bufs.compact_b, + &bufs.compact_a, + &attn.dt_bias, + &attn.a_log, + &mut layer_state.state, + &mut bufs.gdr_scratch, + &mut bufs.compact_gdr, + c.linear_num_key_heads, + c.linear_num_value_heads, + c.linear_key_head_dim, + c.linear_value_head_dim, + )?; + ops::copy_hidden_token_range_into( + &self.ctx, + &bufs.compact_gdr, + 0, + &mut bufs.gdr_out, + row_offset, + seq_len, + )?; + row_offset += seq_len; + } + bufs.gdr_scratch.set_rows(bufs.qkv.seq_len); + + ops::rms_norm_gated_batch_into( + &self.ctx, + &bufs.gdr_out, + &attn.norm_weight, + &bufs.z, + &mut bufs.normed_gated, + c.linear_num_value_heads, + c.linear_value_head_dim, + c.rms_norm_eps, + ); + ops::gemm_into_checked( + &self.ctx, + &attn.out_proj, + &bufs.normed_gated, + &mut bufs.attn_results, + )?; + Ok(()) + } +} diff --git a/pegainfer-qwen35/src/recurrent.rs b/pegainfer-qwen35/src/recurrent.rs index 3fc76f10b..79a4965ab 100644 --- a/pegainfer-qwen35/src/recurrent.rs +++ b/pegainfer-qwen35/src/recurrent.rs @@ -463,11 +463,11 @@ pub fn gated_delta_rule_prefill_chunkwise_into( let expected_chunk_ai_len = expected_chunk_a_len; let expected_chunk_state_len = GdrChunkwiseScratch35::num_chunks(qkv.seq_len) * num_value_heads * val_dim * key_dim; - assert_eq!(scratch.g_cumsum.len(), expected_gate_len); - assert_eq!(scratch.beta.len(), expected_gate_len); - assert_eq!(scratch.a_tril.len(), expected_chunk_a_len); - assert_eq!(scratch.a_inv.len(), expected_chunk_ai_len); - assert_eq!(scratch.chunk_state.len(), expected_chunk_state_len); + assert!(scratch.g_cumsum.len() >= expected_gate_len); + assert!(scratch.beta.len() >= expected_gate_len); + assert!(scratch.a_tril.len() >= expected_chunk_a_len); + assert!(scratch.a_inv.len() >= expected_chunk_ai_len); + assert!(scratch.chunk_state.len() >= expected_chunk_state_len); gated_delta_rule_prefill_chunk_prepare_into( ctx, diff --git a/pegainfer-qwen35/src/recurrent_state.rs b/pegainfer-qwen35/src/recurrent_state.rs index f7279ec12..d6edafe3f 100644 --- a/pegainfer-qwen35/src/recurrent_state.rs +++ b/pegainfer-qwen35/src/recurrent_state.rs @@ -70,6 +70,40 @@ impl RecurrentState { Ok(Self { layers, seq_len: 0 }) } + + /// D2D copy all recurrent and convolution state from `src`. + pub(crate) fn copy_from(&mut self, ctx: &DeviceContext, src: &RecurrentState) -> Result<()> { + anyhow::ensure!( + self.layers.len() == src.layers.len(), + "Qwen3.5 recurrent copy layer mismatch: dst={}, src={}", + self.layers.len(), + src.layers.len() + ); + for (layer_idx, (dst_layer, src_layer)) in + self.layers.iter_mut().zip(src.layers.iter()).enumerate() + { + anyhow::ensure!( + dst_layer.state.len() == src_layer.state.len(), + "Qwen3.5 recurrent state length mismatch at layer {layer_idx}: dst={}, src={}", + dst_layer.state.len(), + src_layer.state.len() + ); + anyhow::ensure!( + dst_layer.conv_state.len == src_layer.conv_state.len, + "Qwen3.5 conv state length mismatch at layer {layer_idx}: dst={}, src={}", + dst_layer.conv_state.len, + src_layer.conv_state.len + ); + ctx.stream + .memcpy_dtod(&src_layer.state, &mut dst_layer.state) + .map_err(|e| anyhow::anyhow!("copy recurrent state layer {layer_idx}: {e}"))?; + ctx.stream + .memcpy_dtod(&src_layer.conv_state.data, &mut dst_layer.conv_state.data) + .map_err(|e| anyhow::anyhow!("copy conv state layer {layer_idx}: {e}"))?; + } + self.seq_len = src.seq_len; + Ok(()) + } } impl LinearStatePointerTables { diff --git a/pegainfer-qwen35/src/speculative.rs b/pegainfer-qwen35/src/speculative.rs new file mode 100644 index 000000000..238a039a0 --- /dev/null +++ b/pegainfer-qwen35/src/speculative.rs @@ -0,0 +1,578 @@ +#[cfg(test)] +use std::cell::Cell; + +use anyhow::Result; +#[cfg(test)] +use pegainfer_core::engine::TokenLogprob; +use pegainfer_core::kv_pool::KvState; +use pegainfer_core::sampler::SamplingParams; + +use crate::batch_decode_graph::BatchDecodeGraphState; +use crate::executor::Qwen35Executor; +use crate::executor::RequestId; +#[cfg(test)] +use crate::logprobs::snapshot_requested_logprobs; +use crate::prefill::PREFILL_CHUNK_LEN; +use crate::recurrent_state::RecurrentState; +use crate::verify_buffers::VerifyBuffers35; +use crate::weights::Qwen35Model; + +#[cfg(test)] +thread_local! { + static FAIL_AFTER_REPLAY_SYNC: Cell = const { Cell::new(false) }; + static FAIL_AFTER_GRAPH_COMMIT_SYNC: Cell = const { Cell::new(false) }; +} + +#[cfg(test)] +pub(crate) fn set_fail_after_replay_sync(enabled: bool) { + FAIL_AFTER_REPLAY_SYNC.set(enabled); +} + +#[cfg(test)] +pub(crate) fn set_fail_after_graph_commit_sync(enabled: bool) { + FAIL_AFTER_GRAPH_COMMIT_SYNC.set(enabled); +} + +#[derive(Clone, Debug)] +pub(crate) struct VerifyStepItem { + pub(crate) request_id: RequestId, + pub(crate) token_ids: Vec, + #[cfg(test)] + pub(crate) diagnostic_logprobs: usize, +} + +impl VerifyStepItem { + #[cfg(not(test))] + pub(crate) fn new(request_id: RequestId, token_ids: Vec) -> Self { + Self { + request_id, + token_ids, + } + } + + #[cfg(test)] + pub(crate) fn new( + request_id: RequestId, + token_ids: Vec, + diagnostic_logprobs: usize, + ) -> Self { + Self { + request_id, + token_ids, + diagnostic_logprobs, + } + } +} + +#[derive(Clone, Copy)] +pub(crate) struct VerifyPlan<'a> { + pub(crate) requests: &'a [VerifyStepItem], +} + +#[cfg(test)] +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct VerifyDiagnostic { + pub(crate) token: u32, + pub(crate) logprob: Option, +} + +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct VerifyRequestResult { + pub(crate) request_id: RequestId, + pub(crate) matched_draft_tokens: usize, + pub(crate) accepted_tokens: Vec, + #[cfg(test)] + pub(crate) diagnostic_posteriors: Vec, +} + +pub(crate) struct VerifyResult { + pub(crate) requests: Vec, +} + +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct VerifySpanResult { + pub(crate) matched_draft_tokens: usize, + pub(crate) accepted_tokens: Vec, + #[cfg(test)] + pub(crate) diagnostic_posteriors: Vec, +} + +#[must_use] +pub(crate) fn accept_greedy(proposed: &[u32], target_argmax: &[u32]) -> (usize, Vec) { + debug_assert_eq!( + target_argmax.len(), + proposed.len() + 1, + "verify must produce one posterior token per draft plus one bonus" + ); + let mut matched = 0usize; + while matched < proposed.len() && proposed[matched] == target_argmax[matched] { + matched += 1; + } + let mut accepted = Vec::with_capacity(matched + 1); + accepted.extend_from_slice(&proposed[..matched]); + accepted.push(target_argmax[matched]); + (matched, accepted) +} + +pub(crate) fn capture_hybrid_states( + model: &Qwen35Model, + graph_state: &BatchDecodeGraphState, + graph_slot_indices: &[usize], + backup_states: &mut [RecurrentState], +) -> Result<()> { + anyhow::ensure!( + graph_slot_indices.len() == backup_states.len(), + "Qwen3.5 speculative backup batch size mismatch" + ); + for (&graph_slot_idx, backup_state) in graph_slot_indices.iter().zip(backup_states.iter_mut()) { + graph_state.copy_slot_to_state(model.device_ctx(), graph_slot_idx, backup_state)?; + } + Ok(()) +} + +#[allow(clippy::too_many_arguments)] +pub(crate) fn run_hybrid_verify( + model: &Qwen35Model, + kv_states: &mut [&mut KvState], + spans: &[&[u32]], + verify_bufs: &mut VerifyBuffers35, + backup_states: &[RecurrentState], + verify_states: &mut [RecurrentState], +) -> Result<()> { + let batch = spans.len(); + anyhow::ensure!(batch > 0, "Qwen3.5 speculative verify batch is empty"); + anyhow::ensure!( + kv_states.len() == batch && backup_states.len() == batch && verify_states.len() == batch, + "Qwen3.5 speculative verify batch size mismatch" + ); + for span in spans { + anyhow::ensure!( + !span.is_empty(), + "Qwen3.5 speculative verify requires a non-empty token span" + ); + } + + for (verify_state, backup_state) in verify_states.iter_mut().zip(backup_states.iter()) { + verify_state.copy_from(model.device_ctx(), backup_state)?; + } + let mut verify_refs: Vec<&mut RecurrentState> = verify_states.iter_mut().collect(); + model.prefill_verify_into(spans, kv_states, &mut verify_refs, verify_bufs) +} + +#[allow(clippy::too_many_arguments)] +pub(crate) fn verify_hybrid_spans( + model: &Qwen35Model, + kv_states: &mut [&mut KvState], + spans: &[&[u32]], + #[cfg(test)] diagnostic_logprobs: &[usize], + verify_bufs: &mut VerifyBuffers35, + backup_states: &[RecurrentState], + verify_states: &mut [RecurrentState], +) -> Result> { + let batch = spans.len(); + #[cfg(test)] + anyhow::ensure!( + batch == diagnostic_logprobs.len(), + "Qwen3.5 speculative verify diagnostic batch size mismatch" + ); + for span in spans { + anyhow::ensure!( + !span.is_empty(), + "Qwen3.5 speculative verify needs at least the current token" + ); + } + run_hybrid_verify( + model, + kv_states, + spans, + verify_bufs, + backup_states, + verify_states, + )?; + + #[cfg(test)] + let row_diagnostics: Vec = spans + .iter() + .zip(diagnostic_logprobs.iter()) + .flat_map(|(span, &logprobs)| std::iter::repeat_n(logprobs, span.len())) + .collect(); + #[cfg(test)] + let diagnostic_logits = + snapshot_requested_logprobs(model.device_ctx(), &verify_bufs.logits, &row_diagnostics)?; + let greedy = SamplingParams::default(); + let params = vec![&greedy; verify_bufs.logits.seq_len]; + let steps = vec![0_u64; verify_bufs.logits.seq_len]; + let target_tokens = pegainfer_sample::select_batch( + model.device_ctx(), + &verify_bufs.logits, + ¶ms, + &steps, + 0, + &mut verify_bufs.sample, + )?; + + let mut outputs = Vec::with_capacity(batch); + let mut row_offset = 0usize; + for span in spans { + #[cfg(test)] + let slot_idx = outputs.len(); + let row_end = row_offset + span.len(); + let target_slice = &target_tokens[row_offset..row_end]; + let (matched, accepted_ids) = accept_greedy(&span[1..], target_slice); + #[cfg(test)] + let diagnostic_posteriors = target_slice + .iter() + .copied() + .enumerate() + .map(|(i, token)| VerifyDiagnostic { + token, + logprob: diagnostic_logits[row_offset + i].as_ref().and_then(|row| { + pegainfer_sample::token_logprob_from_row( + row, + token, + diagnostic_logprobs[slot_idx], + ) + }), + }) + .collect(); + outputs.push(VerifySpanResult { + matched_draft_tokens: matched, + accepted_tokens: accepted_ids, + #[cfg(test)] + diagnostic_posteriors, + }); + row_offset = row_end; + } + Ok(outputs) +} + +#[allow(clippy::too_many_arguments)] +pub(crate) fn commit_hybrid_states( + model: &Qwen35Model, + kv_states: &mut [&mut KvState], + graph_state: &mut BatchDecodeGraphState, + graph_slot_indices: &[usize], + spans: &[&[u32]], + backup_states: &[RecurrentState], + verify_states: &mut [RecurrentState], + verify_bufs: &mut VerifyBuffers35, + original_seq_lens: &[usize], + results: &[VerifySpanResult], +) -> Result<()> { + let batch = spans.len(); + anyhow::ensure!( + kv_states.len() == batch + && graph_slot_indices.len() == batch + && backup_states.len() == batch + && verify_states.len() == batch + && original_seq_lens.len() == batch + && results.len() == batch, + "Qwen3.5 speculative commit batch size mismatch" + ); + + for slot_idx in 0..batch { + let accepted_len = results[slot_idx].accepted_tokens.len(); + anyhow::ensure!( + accepted_len > 0 && accepted_len <= spans[slot_idx].len(), + "Qwen3.5 speculative accepted span {} is invalid for verify span {}", + accepted_len, + spans[slot_idx].len() + ); + if accepted_len == spans[slot_idx].len() { + continue; + } + + kv_states[slot_idx].truncate_to(original_seq_lens[slot_idx])?; + verify_states[slot_idx].copy_from(model.device_ctx(), &backup_states[slot_idx])?; + let mut replay_tokens = Vec::with_capacity(accepted_len); + replay_tokens.push(spans[slot_idx][0]); + replay_tokens.extend( + results[slot_idx] + .accepted_tokens + .iter() + .take(accepted_len - 1) + .copied(), + ); + let replay_spans = [replay_tokens.as_slice()]; + let mut one_kv = [&mut *kv_states[slot_idx]]; + let mut one_state = [&mut verify_states[slot_idx]]; + pegainfer_kernels::ops::with_gemm_lt_disabled(|| { + model.prefill_verify_into(&replay_spans, &mut one_kv, &mut one_state, verify_bufs) + })?; + } + + model + .device_ctx() + .stream + .synchronize() + .map_err(|e| anyhow::anyhow!("Qwen3.5 speculative replay synchronization failed: {e}"))?; + #[cfg(test)] + anyhow::ensure!( + !FAIL_AFTER_REPLAY_SYNC.replace(false), + "injected Qwen3.5 speculative replay synchronization failure" + ); + + for slot_idx in 0..batch { + graph_state.copy_state_to_slot( + model.device_ctx(), + &verify_states[slot_idx], + graph_slot_indices[slot_idx], + )?; + } + model + .device_ctx() + .stream + .synchronize() + .map_err(|e| anyhow::anyhow!("Qwen3.5 speculative commit synchronization failed: {e}"))?; + #[cfg(test)] + anyhow::ensure!( + !FAIL_AFTER_GRAPH_COMMIT_SYNC.replace(false), + "injected Qwen3.5 speculative graph commit synchronization failure" + ); + Ok(()) +} + +pub(crate) fn restore_hybrid_states( + model: &Qwen35Model, + kv_states: &mut [&mut KvState], + graph_state: &mut BatchDecodeGraphState, + graph_slot_indices: &[usize], + backup_states: &[RecurrentState], + original_seq_lens: &[usize], +) -> Result<()> { + anyhow::ensure!( + kv_states.len() == graph_slot_indices.len() + && kv_states.len() == backup_states.len() + && kv_states.len() == original_seq_lens.len(), + "Qwen3.5 speculative restore batch size mismatch" + ); + let mut errors = Vec::new(); + for slot_idx in 0..kv_states.len() { + if let Err(err) = kv_states[slot_idx].truncate_to(original_seq_lens[slot_idx]) { + errors.push(format!( + "truncate slot {slot_idx} to {} failed: {err}", + original_seq_lens[slot_idx] + )); + } + if let Err(err) = graph_state.copy_state_to_slot( + model.device_ctx(), + &backup_states[slot_idx], + graph_slot_indices[slot_idx], + ) { + errors.push(format!("restore recurrent slot {slot_idx} failed: {err}")); + } + } + if let Err(err) = model.device_ctx().stream.synchronize() { + errors.push(format!( + "synchronize restored recurrent states failed: {err}" + )); + } + anyhow::ensure!( + errors.is_empty(), + "Qwen3.5 speculative rollback failed: {}", + errors.join("; ") + ); + Ok(()) +} + +impl Qwen35Executor { + pub(crate) fn execute_speculative_verify( + &mut self, + plan: VerifyPlan<'_>, + ) -> Result { + anyhow::ensure!( + !pegainfer_kernels::tensor::has_stream_override(), + "Qwen3.5 speculative verify does not support a CUDA stream override" + ); + self.validate_speculative_verify(plan)?; + let batch = self.active.len(); + let graph_slot_indices: Vec = self + .active + .iter() + .map(|active| active.graph_slot_idx) + .collect(); + let original_seq_lens: Vec = self + .active + .iter() + .map(|active| active.kv.seq_len()) + .collect(); + let mut backup_states = Vec::with_capacity(batch); + let mut verify_states = Vec::with_capacity(batch); + for _ in 0..batch { + backup_states.push(RecurrentState::new( + self.model.device_ctx(), + self.model.config(), + )?); + verify_states.push(RecurrentState::new( + self.model.device_ctx(), + self.model.config(), + )?); + } + capture_hybrid_states( + &self.model, + &self.graph_state, + &graph_slot_indices, + &mut backup_states, + )?; + self.model.device_ctx().stream.synchronize().map_err(|e| { + anyhow::anyhow!("Qwen3.5 speculative backup synchronization failed: {e}") + })?; + + let max_span = plan + .requests + .iter() + .map(|req| req.token_ids.len()) + .max() + .unwrap_or(1); + let mut verify_bufs = VerifyBuffers35::new( + self.model.device_ctx(), + self.model.config(), + batch, + max_span, + self.model.kv_pool().capacity_pages(), + )?; + let spans: Vec<&[u32]> = plan + .requests + .iter() + .map(|req| req.token_ids.as_slice()) + .collect(); + #[cfg(test)] + let diagnostic_logprobs: Vec = plan + .requests + .iter() + .map(|req| req.diagnostic_logprobs) + .collect(); + + let transaction = (|| -> Result> { + let mut kv_states: Vec<&mut KvState> = self + .active + .iter_mut() + .map(|active| &mut active.kv) + .collect(); + #[cfg(test)] + let results = verify_hybrid_spans( + &self.model, + &mut kv_states, + &spans, + &diagnostic_logprobs, + &mut verify_bufs, + &backup_states, + &mut verify_states, + )?; + #[cfg(not(test))] + let results = verify_hybrid_spans( + &self.model, + &mut kv_states, + &spans, + &mut verify_bufs, + &backup_states, + &mut verify_states, + )?; + commit_hybrid_states( + &self.model, + &mut kv_states, + &mut self.graph_state, + &graph_slot_indices, + &spans, + &backup_states, + &mut verify_states, + &mut verify_bufs, + &original_seq_lens, + &results, + )?; + Ok(results) + })(); + + let results = match transaction { + Ok(results) => results, + Err(err) => { + let mut kv_states: Vec<&mut KvState> = self + .active + .iter_mut() + .map(|active| &mut active.kv) + .collect(); + if let Err(rollback_err) = restore_hybrid_states( + &self.model, + &mut kv_states, + &mut self.graph_state, + &graph_slot_indices, + &backup_states, + &original_seq_lens, + ) { + anyhow::bail!("{err}; additionally failed to roll back: {rollback_err}"); + } + return Err(err); + } + }; + + Ok(VerifyResult { + requests: plan + .requests + .iter() + .zip(results) + .map(|(request, result)| VerifyRequestResult { + request_id: request.request_id, + matched_draft_tokens: result.matched_draft_tokens, + accepted_tokens: result.accepted_tokens, + #[cfg(test)] + diagnostic_posteriors: result.diagnostic_posteriors, + }) + .collect(), + }) + } + + fn validate_speculative_verify(&self, plan: VerifyPlan<'_>) -> Result<()> { + anyhow::ensure!( + !plan.requests.is_empty(), + "Qwen3.5 speculative verify plan requires at least one request" + ); + anyhow::ensure!( + plan.requests.len() == self.active.len(), + "Qwen3.5 speculative verify must include all active requests in slot order" + ); + for (slot_idx, req) in plan.requests.iter().enumerate() { + anyhow::ensure!( + self.active[slot_idx].request_id == req.request_id, + "Qwen3.5 speculative verify request order differs from active slot order" + ); + anyhow::ensure!( + !req.token_ids.is_empty(), + "Qwen3.5 speculative verify request {} needs at least the current token", + req.request_id.get() + ); + anyhow::ensure!( + req.token_ids.len() <= PREFILL_CHUNK_LEN, + "Qwen3.5 speculative verify request {} span len {} exceeds max chunk {PREFILL_CHUNK_LEN}", + req.request_id.get(), + req.token_ids.len() + ); + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn accepts_full_run_plus_bonus() { + let (matched, accepted) = accept_greedy(&[10, 11, 12], &[10, 11, 12, 13]); + assert_eq!(matched, 3); + assert_eq!(accepted, vec![10, 11, 12, 13]); + } + + #[test] + fn accepts_prefix_then_correction() { + let (matched, accepted) = accept_greedy(&[10, 11, 99], &[10, 11, 22, 33]); + assert_eq!(matched, 2); + assert_eq!(accepted, vec![10, 11, 22]); + } + + #[test] + fn rejects_first_candidate_commits_one() { + let (matched, accepted) = accept_greedy(&[10, 11, 12], &[7, 8, 9, 10]); + assert_eq!(matched, 0); + assert_eq!(accepted, vec![7]); + } +} diff --git a/pegainfer-qwen35/src/speculative_tests.rs b/pegainfer-qwen35/src/speculative_tests.rs new file mode 100644 index 000000000..66e26f077 --- /dev/null +++ b/pegainfer-qwen35/src/speculative_tests.rs @@ -0,0 +1,1017 @@ +use std::path::Path; + +use pegainfer_kernels::ops::NumericPolicy; +use pegainfer_kernels::ops::numeric_policy; +use pegainfer_kernels::ops::set_numeric_policy; +use pegainfer_kernels::tensor::StreamOverrideGuard; + +use crate::executor::DecodePlan; +use crate::executor::DecodeStepItem; +use crate::executor::PrefillPlan; +use crate::executor::PrefillStepItem; +use crate::executor::Qwen35Executor; +use crate::executor::RequestId; +use crate::speculative::VerifyPlan; +use crate::speculative::VerifyStepItem; +use crate::speculative::set_fail_after_graph_commit_sync; +use crate::speculative::set_fail_after_replay_sync; + +const MODEL_PATH: &str = concat!(env!("CARGO_MANIFEST_DIR"), "/../models/Qwen3.5-4B"); +const LOGPROBS: usize = 1; +const DIAG_LOGPROBS: usize = 20; +const MARGIN_TOL: f32 = 0.20; + +struct NumericPolicyGuard(NumericPolicy); + +impl NumericPolicyGuard { + fn set(policy: NumericPolicy) -> Self { + let previous = numeric_policy(); + set_numeric_policy(policy); + Self(previous) + } +} + +impl Drop for NumericPolicyGuard { + fn drop(&mut self) { + set_numeric_policy(self.0); + } +} + +struct TestFailpointGuard { + setter: fn(bool), +} + +impl TestFailpointGuard { + fn new(setter: fn(bool)) -> Self { + setter(true); + Self { setter } + } +} + +impl Drop for TestFailpointGuard { + fn drop(&mut self) { + (self.setter)(false); + } +} + +#[derive(Clone, Debug)] +struct TokenDiag { + token: u32, + top_logprobs: Vec<(u32, f32)>, +} + +#[derive(Clone)] +struct CaseSpec { + request_id: RequestId, + prompt_tokens: Vec, + draft_len: usize, + reject_at: Option, +} + +#[derive(Clone)] +struct CaseExpectation { + request_id: RequestId, + first_token: u32, + draft_tokens: Vec, + accepted_tokens: Vec, +} + +fn model_path() -> String { + let path = + std::env::var("OPENINFER_TEST_MODEL_PATH").unwrap_or_else(|_| MODEL_PATH.to_string()); + assert!( + Path::new(&path).join("config.json").exists(), + "Qwen3.5 model is missing at {path}; set OPENINFER_TEST_MODEL_PATH" + ); + path +} + +fn build_executor(model_path: &str, capacity: usize) -> Qwen35Executor { + let capacity = [1usize, 2, 4, 8, 16, 32, 64] + .into_iter() + .find(|bucket| *bucket >= capacity) + .expect("test batch exceeds Qwen3.5 decode bucket capacity"); + Qwen35Executor::from_runtime_with_capacity(model_path, false, &[0], capacity) + .expect("load Qwen3.5 executor") +} + +fn prefill(exec: &mut Qwen35Executor, cases: &[CaseSpec]) -> Vec { + let reqs: Vec<_> = cases + .iter() + .map(|case| PrefillStepItem::new(case.request_id, case.prompt_tokens.clone(), LOGPROBS)) + .collect(); + exec.execute_prefill(PrefillPlan { requests: &reqs }) + .expect("prefill") + .requests + .into_iter() + .map(|result| result.first_token) + .collect() +} + +fn prefill_with_logprobs( + exec: &mut Qwen35Executor, + cases: &[CaseSpec], + logprobs: usize, +) -> Vec { + let reqs: Vec<_> = cases + .iter() + .map(|case| PrefillStepItem::new(case.request_id, case.prompt_tokens.clone(), logprobs)) + .collect(); + exec.execute_prefill(PrefillPlan { requests: &reqs }) + .expect("prefill") + .requests + .into_iter() + .map(|result| TokenDiag { + token: result.first_token, + top_logprobs: result + .first_token_logprob + .map(|lp| lp.top_logprobs) + .unwrap_or_default(), + }) + .collect() +} + +fn decode_once(exec: &mut Qwen35Executor, tokens: &[u32], cases: &[CaseSpec]) -> Vec { + let reqs: Vec<_> = cases + .iter() + .zip(tokens.iter()) + .map(|(case, &token)| DecodeStepItem::new(case.request_id, token, LOGPROBS)) + .collect(); + exec.execute_decode(DecodePlan { requests: &reqs }) + .expect("decode") + .requests + .into_iter() + .map(|result| result.token) + .collect() +} + +fn decode_once_with_logprobs( + exec: &mut Qwen35Executor, + tokens: &[u32], + cases: &[CaseSpec], + logprobs: usize, +) -> Vec { + let reqs: Vec<_> = cases + .iter() + .zip(tokens.iter()) + .map(|(case, &token)| DecodeStepItem::new(case.request_id, token, logprobs)) + .collect(); + exec.execute_decode(DecodePlan { requests: &reqs }) + .expect("decode") + .requests + .into_iter() + .map(|result| TokenDiag { + token: result.token, + top_logprobs: result.logprob.map(|lp| lp.top_logprobs).unwrap_or_default(), + }) + .collect() +} + +fn regret(top_logprobs: &[(u32, f32)], token: u32) -> Option { + top_logprobs.first().and_then(|(_, top_lp)| { + top_logprobs + .iter() + .find(|(id, _)| *id == token) + .map(|(_, lp)| top_lp - lp) + }) +} + +fn top_ids(top_logprobs: &[(u32, f32)]) -> Vec { + top_logprobs.iter().take(8).map(|(id, _)| *id).collect() +} + +fn assert_first_posterior_matches_decode(model_path: &str, cases: &[CaseSpec], context: &str) { + let mut decode_exec = build_executor(model_path, cases.len()); + let first_diag = prefill_with_logprobs(&mut decode_exec, cases, DIAG_LOGPROBS); + let first_tokens: Vec = first_diag.iter().map(|diag| diag.token).collect(); + let decode_next = + decode_once_with_logprobs(&mut decode_exec, &first_tokens, cases, DIAG_LOGPROBS); + drop(decode_exec); + + let mut verify_exec = build_executor(model_path, cases.len()); + let verify_first_tokens = prefill(&mut verify_exec, cases); + assert_eq!(verify_first_tokens, first_tokens); + let verify_items: Vec<_> = cases + .iter() + .zip(first_tokens.iter()) + .map(|(case, &first)| { + VerifyStepItem::new( + case.request_id, + vec![first, first.wrapping_add(17)], + DIAG_LOGPROBS, + ) + }) + .collect(); + let verify = verify_exec + .execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) + .expect("speculative verify"); + let verify_next: Vec = verify + .requests + .iter() + .map(|row| { + let posterior = &row.diagnostic_posteriors[0]; + assert_eq!(row.accepted_tokens[0], posterior.token); + TokenDiag { + token: posterior.token, + top_logprobs: posterior + .logprob + .as_ref() + .map(|lp| lp.top_logprobs.clone()) + .unwrap_or_default(), + } + }) + .collect(); + + let hard_mismatches: Vec = decode_next + .iter() + .zip(verify_next.iter()) + .enumerate() + .filter_map(|(idx, (decode, verify))| { + if decode.token == verify.token { + return None; + } + let decode_regret_for_verify = regret(&decode.top_logprobs, verify.token); + let verify_regret_for_decode = regret(&verify.top_logprobs, decode.token); + let within_decode = decode_regret_for_verify.is_some_and(|r| r <= MARGIN_TOL); + let within_verify = verify_regret_for_decode.is_some_and(|r| r <= MARGIN_TOL); + (!within_decode || !within_verify).then(|| { + format!( + "idx={idx} first={} decode={} verify={} decode_regret_for_verify={:?} verify_regret_for_decode={:?} decode_top={:?} verify_top={:?}", + first_tokens[idx], + decode.token, + verify.token, + decode_regret_for_verify, + verify_regret_for_decode, + top_ids(&decode.top_logprobs), + top_ids(&verify.top_logprobs), + ) + }) + }) + .collect(); + assert!( + hard_mismatches.is_empty(), + "Qwen3.5 speculative verifier {context} posterior has non-tie divergence from decode:\n{}", + hard_mismatches.join("\n") + ); +} + +fn deterministic_long_prompt(len: usize, request_idx: usize) -> Vec { + (0..len) + .map(|i| 100 + ((i * 7919 + request_idx * 104_729) % 99_000) as u32) + .collect() +} + +fn stable_text_like_prompt(len: usize, request_idx: usize) -> Vec { + let segment = [ + 2387, 220, 16, 321, 9707, 374, 3565, 3838, 374, 220, 17, 10, 17, 785, 9282, 374, 3565, 198, + 15123, 839, 13, 220, 1024, 11, 256, 11, 4096, 13, 220, 2301, 374, 690, 1012, 13, 220, + ]; + let mut tokens = Vec::with_capacity(len); + while tokens.len() < len { + tokens.extend(segment.iter().map(|token| token + request_idx as u32)); + } + tokens.truncate(len); + tokens +} + +fn decode_oracle_tokens( + model_path: &str, + prompts: &[Vec], + token_count: usize, +) -> Vec> { + assert!(token_count > 0, "oracle must request at least one token"); + let batch = prompts.len(); + let mut exec = build_executor(model_path, batch); + let cases: Vec<_> = prompts + .iter() + .enumerate() + .map(|(idx, prompt_tokens)| CaseSpec { + request_id: RequestId::new((idx + 1) as u64), + prompt_tokens: prompt_tokens.clone(), + draft_len: 0, + reject_at: None, + }) + .collect(); + let first = prefill_with_logprobs(&mut exec, &cases, DIAG_LOGPROBS); + let mut generated: Vec> = first.into_iter().map(|diag| vec![diag]).collect(); + + while generated[0].len() < token_count { + let fed: Vec = generated + .iter() + .map(|tokens| tokens.last().expect("generated token").token) + .collect(); + let next = decode_once_with_logprobs(&mut exec, &fed, &cases, DIAG_LOGPROBS); + for (row, token) in generated.iter_mut().zip(next) { + row.push(token); + } + } + generated +} + +fn build_expectations(model_path: &str, cases: &[CaseSpec]) -> Vec { + let mut exec = build_executor(model_path, cases.len()); + let first_tokens = prefill(&mut exec, cases); + let max_len = cases + .iter() + .map(|case| case.draft_len + 3) + .max() + .expect("at least one case"); + let mut generated: Vec> = first_tokens.into_iter().map(|token| vec![token]).collect(); + while generated.iter().any(|tokens| tokens.len() < max_len) { + let fed: Vec = generated + .iter() + .map(|tokens| *tokens.last().expect("prefill token")) + .collect(); + for (tokens, next) in generated + .iter_mut() + .zip(decode_once(&mut exec, &fed, cases)) + { + tokens.push(next); + } + } + + cases + .iter() + .zip(generated.iter()) + .map(|(case, generated)| { + let first = generated[0]; + + let mut draft_tokens = generated[1..case.draft_len + 1].to_vec(); + if let Some(reject_at) = case.reject_at { + draft_tokens[reject_at] = draft_tokens[reject_at].wrapping_add(17); + } + let matched = case.reject_at.unwrap_or(case.draft_len); + let accepted_tokens = generated[1..=matched + 1].to_vec(); + + CaseExpectation { + request_id: case.request_id, + first_token: first, + draft_tokens, + accepted_tokens, + } + }) + .collect() +} + +fn run_speculative_case(model_path: &str, cases: &[CaseSpec]) { + let expectations = build_expectations(model_path, cases); + let mut exec = build_executor(model_path, cases.len()); + let first_tokens = prefill(&mut exec, cases); + assert_eq!( + first_tokens, + expectations + .iter() + .map(|expect| expect.first_token) + .collect::>() + ); + let before_state = exec.debug_state_summary(); + + let verify_items: Vec<_> = expectations + .iter() + .map(|expect| { + let mut token_ids = Vec::with_capacity(expect.draft_tokens.len() + 1); + token_ids.push(expect.first_token); + token_ids.extend_from_slice(&expect.draft_tokens); + VerifyStepItem::new(expect.request_id, token_ids, LOGPROBS) + }) + .collect(); + let result = exec + .execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) + .expect("speculative verify"); + let after_state = exec.debug_state_summary(); + + assert_eq!(result.requests.len(), expectations.len()); + for ((row, expect), case) in result + .requests + .iter() + .zip(expectations.iter()) + .zip(cases.iter()) + { + assert_eq!(row.request_id, expect.request_id); + let accepted_ids = &row.accepted_tokens; + assert!( + !accepted_ids.is_empty(), + "speculative verify must commit at least one token for request {:?}", + expect.request_id + ); + let expected_matched = case.reject_at.unwrap_or(expect.draft_tokens.len()); + assert_eq!( + row.matched_draft_tokens, expected_matched, + "speculative fixture did not execute its intended acceptance branch for request {:?}: expected matched={}, actual matched={}, drafts={:?}", + expect.request_id, expected_matched, row.matched_draft_tokens, expect.draft_tokens + ); + assert_eq!( + accepted_ids.len(), + row.matched_draft_tokens + 1, + "accepted token count must equal matched drafts plus bonus for request {:?}: first={}, drafts={:?}, actual_accepted={:?}", + expect.request_id, + expect.first_token, + expect.draft_tokens, + accepted_ids + ); + assert_eq!( + accepted_ids, &expect.accepted_tokens, + "spec verify accepted tokens differ from the independent sequential decode oracle for request {:?}", + expect.request_id + ); + assert_eq!( + accepted_ids + .iter() + .copied() + .take(row.matched_draft_tokens) + .collect::>(), + expect.draft_tokens[..row.matched_draft_tokens], + "spec verify accepted prefix mismatch for request {:?}: first={}, drafts={:?}", + expect.request_id, + expect.first_token, + expect.draft_tokens + ); + } + for ((before, after), row) in before_state + .iter() + .zip(after_state.iter()) + .zip(result.requests.iter()) + { + assert_eq!(before.request_id, after.request_id); + assert_eq!( + before.kv_seq_len + row.accepted_tokens.len(), + after.kv_seq_len + ); + assert_eq!( + before.recurrent_seq_len + row.accepted_tokens.len(), + after.recurrent_seq_len + ); + } + + let accepted_rows: Vec> = result + .requests + .iter() + .map(|row| row.accepted_tokens.clone()) + .collect(); + let last_tokens: Vec = accepted_rows + .iter() + .map(|row| *row.last().expect("accepted token")) + .collect(); + let followup_reqs: Vec<_> = cases + .iter() + .zip(last_tokens.iter()) + .map(|(case, &token)| DecodeStepItem::new(case.request_id, token, DIAG_LOGPROBS)) + .collect(); + let followup = exec + .execute_decode(DecodePlan { + requests: &followup_reqs, + }) + .expect("post-spec followup") + .requests; + let actual_followup: Vec<_> = followup + .iter() + .zip(expectations.iter()) + .map(|(actual, expect)| { + assert_eq!(actual.request_id, expect.request_id); + TokenDiag { + token: actual.token, + top_logprobs: actual + .logprob + .as_ref() + .map(|logprob| logprob.top_logprobs.clone()) + .unwrap_or_default(), + } + }) + .collect(); + assert_eq!(actual_followup.len(), cases.len()); + drop(exec); + + let oracle_cases: Vec<_> = cases + .iter() + .zip(expectations.iter()) + .zip(accepted_rows.iter()) + .map(|((case, expect), accepted)| { + let mut prompt_tokens = case.prompt_tokens.clone(); + prompt_tokens.push(expect.first_token); + prompt_tokens.extend_from_slice(accepted); + CaseSpec { + request_id: case.request_id, + prompt_tokens, + draft_len: 0, + reject_at: None, + } + }) + .collect(); + let mut oracle_exec = build_executor(model_path, cases.len()); + let oracle_followup = prefill_with_logprobs(&mut oracle_exec, &oracle_cases, DIAG_LOGPROBS); + + let mut hard_mismatches = Vec::new(); + for (idx, (actual, oracle)) in actual_followup + .iter() + .zip(oracle_followup.iter()) + .enumerate() + { + if actual.token != oracle.token { + let actual_regret_for_oracle = regret(&actual.top_logprobs, oracle.token); + let oracle_regret_for_actual = regret(&oracle.top_logprobs, actual.token); + let within_actual = actual_regret_for_oracle.is_some_and(|regret| regret <= MARGIN_TOL); + let within_oracle = oracle_regret_for_actual.is_some_and(|regret| regret <= MARGIN_TOL); + if !within_actual || !within_oracle { + hard_mismatches.push(format!( + "idx={idx} actual={} oracle={} actual_regret_for_oracle={actual_regret_for_oracle:?} oracle_regret_for_actual={oracle_regret_for_actual:?}", + actual.token, oracle.token + )); + } + } + } + assert!( + hard_mismatches.is_empty(), + "Qwen3.5 speculative commit state diverged from a full-context replay:\n{}", + hard_mismatches.join("\n") + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_verify_first_token_matches_decode_long_batch() { + let model_path = model_path(); + let batch = 8usize; + let cases: Vec<_> = (0..batch) + .map(|idx| CaseSpec { + request_id: RequestId::new((idx + 1) as u64), + prompt_tokens: deterministic_long_prompt(4096, idx), + draft_len: 1, + reject_at: Some(0), + }) + .collect(); + assert_first_posterior_matches_decode(&model_path, &cases, "long-batch first"); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_verify_first_token_matches_decode_benchmark_c16() { + let model_path = model_path(); + let batch = 16usize; + let cases: Vec<_> = (0..batch) + .map(|idx| CaseSpec { + request_id: RequestId::new((idx + 1) as u64), + prompt_tokens: stable_text_like_prompt(1024, idx), + draft_len: 1, + reject_at: Some(0), + }) + .collect(); + assert_first_posterior_matches_decode(&model_path, &cases, "benchmark c16 first"); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_verify_multitoken_span_matches_decode_benchmark_c16() { + let model_path = model_path(); + let batch = 16usize; + let verify_span = 5usize; + let prompts: Vec<_> = (0..batch) + .map(|idx| stable_text_like_prompt(1024, idx)) + .collect(); + let oracle = decode_oracle_tokens(&model_path, &prompts, verify_span + 1); + + let cases: Vec<_> = prompts + .into_iter() + .enumerate() + .map(|(idx, prompt_tokens)| CaseSpec { + request_id: RequestId::new((idx + 1) as u64), + prompt_tokens, + draft_len: verify_span - 1, + reject_at: None, + }) + .collect(); + let mut verify_exec = build_executor(&model_path, batch); + let first_tokens = prefill(&mut verify_exec, &cases); + let expected_first: Vec = oracle.iter().map(|row| row[0].token).collect(); + assert_eq!(first_tokens, expected_first); + + let verify_items: Vec<_> = cases + .iter() + .zip(oracle.iter()) + .map(|(case, row)| { + let token_ids = row[..verify_span] + .iter() + .map(|diag| diag.token) + .collect::>(); + VerifyStepItem::new(case.request_id, token_ids, DIAG_LOGPROBS) + }) + .collect(); + let verify = verify_exec + .execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) + .expect("speculative verify"); + + let mut hard_mismatches = Vec::new(); + for (idx, (result, oracle_row)) in verify.requests.iter().zip(oracle.iter()).enumerate() { + let expected = &oracle_row[1..=verify_span]; + let actual = result + .diagnostic_posteriors + .iter() + .map(|posterior| TokenDiag { + token: posterior.token, + top_logprobs: posterior + .logprob + .as_ref() + .map(|lp| lp.top_logprobs.clone()) + .unwrap_or_default(), + }) + .collect::>(); + assert_eq!( + actual.len(), + expected.len(), + "every verifier posterior row must have a diagnostic oracle row for idx={idx}" + ); + let submitted_drafts = &verify_items[idx].token_ids[1..]; + assert_eq!( + result.matched_draft_tokens, + submitted_drafts.len(), + "all sequential-oracle drafts must be accepted for idx={idx}" + ); + let expected_accepted = expected.iter().map(|diag| diag.token).collect::>(); + assert_eq!( + result.accepted_tokens, expected_accepted, + "accepted tokens must match the independent sequential decode oracle for idx={idx}" + ); + for (row_idx, (actual_diag, expected_diag)) in + actual.iter().zip(expected.iter()).enumerate() + { + if actual_diag.token == expected_diag.token { + continue; + } + let oracle_regret = regret(&expected_diag.top_logprobs, actual_diag.token); + let verify_regret = regret(&actual_diag.top_logprobs, expected_diag.token); + let within_oracle = oracle_regret.is_some_and(|regret| regret <= MARGIN_TOL); + let within_verify = verify_regret.is_some_and(|regret| regret <= MARGIN_TOL); + if !within_oracle || !within_verify { + hard_mismatches.push(format!( + "idx={idx} row={row_idx} accepted_len={} matched_drafts={} expected={} actual={} oracle_regret={oracle_regret:?} verify_regret={verify_regret:?}", + actual.len(), + result.matched_draft_tokens, + expected_diag.token, + actual_diag.token, + )); + } + } + } + + assert!( + hard_mismatches.is_empty(), + "Qwen3.5 speculative verifier benchmark c16 multitoken posterior diverged from decode oracle:\n{}", + hard_mismatches.join("\n") + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_single_request_commits_transaction_state() { + let model_path = model_path(); + run_speculative_case( + &model_path, + &[CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: vec![9707], + draft_len: 3, + reject_at: None, + }], + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_accept_prefix_and_reject_first_transaction_state() { + let model_path = model_path(); + run_speculative_case( + &model_path, + &[ + CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: vec![3838, 374, 220, 17, 10, 17], + draft_len: 4, + reject_at: Some(2), + }, + CaseSpec { + request_id: RequestId::new(2), + prompt_tokens: vec![9707], + draft_len: 4, + reject_at: Some(0), + }, + ], + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_mixed_batch_commits_transaction_state() { + let model_path = model_path(); + run_speculative_case( + &model_path, + &[ + CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: vec![9707], + draft_len: 0, + reject_at: None, + }, + CaseSpec { + request_id: RequestId::new(2), + prompt_tokens: vec![3838, 374, 220, 17, 10, 17], + draft_len: 4, + reject_at: Some(1), + }, + CaseSpec { + request_id: RequestId::new(3), + prompt_tokens: vec![785, 9282, 374, 3565], + draft_len: 3, + reject_at: Some(0), + }, + ], + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_rolls_back_after_projection_error() { + let model_path = model_path(); + let cases = [CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: stable_text_like_prompt(15, 0), + draft_len: 2, + reject_at: None, + }]; + let expectations = build_expectations(&model_path, &cases); + let mut exec = build_executor(&model_path, cases.len()); + let first_tokens = prefill(&mut exec, &cases); + let before_state = exec.debug_state_summary(); + let verify_items = [VerifyStepItem::new( + cases[0].request_id, + vec![ + first_tokens[0], + expectations[0].draft_tokens[0], + expectations[0].draft_tokens[1], + ], + LOGPROBS, + )]; + + let override_error = { + let stream = exec.model.device_ctx().stream.cu_stream(); + let _override = unsafe { StreamOverrideGuard::activate(stream) }; + match exec.execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) { + Err(error) => error, + Ok(_) => panic!("a speculative verify stream override must be rejected"), + } + }; + assert!( + override_error + .to_string() + .contains("does not support a CUDA stream override"), + "unexpected stream override failure: {override_error:#}" + ); + assert_eq!(exec.debug_state_summary(), before_state); + + let error = { + let _policy = NumericPolicyGuard::set(NumericPolicy::Pin); + match exec.execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) { + Err(error) => error, + Ok(_) => panic!("an unwarmed Pin projection must fail inside speculative verify"), + } + }; + assert!( + error.to_string().contains("Pin GEMM cannot serve"), + "unexpected speculative verify failure: {error:#}" + ); + assert_eq!(exec.debug_state_summary(), before_state); + + let followup = decode_once(&mut exec, &first_tokens, &cases); + assert_eq!( + followup, + vec![expectations[0].draft_tokens[0]], + "normal decode after rollback must match an independent executor" + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_rolls_back_after_replay_sync_error() { + let model_path = model_path(); + let cases = [CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: stable_text_like_prompt(15, 0), + draft_len: 2, + reject_at: Some(1), + }]; + let expectations = build_expectations(&model_path, &cases); + let mut exec = build_executor(&model_path, cases.len()); + let first_tokens = prefill(&mut exec, &cases); + let before_state = exec.debug_state_summary(); + let verify_items = [VerifyStepItem::new( + cases[0].request_id, + vec![ + first_tokens[0], + expectations[0].draft_tokens[0], + expectations[0].draft_tokens[1], + ], + LOGPROBS, + )]; + + let error = { + let _failure = TestFailpointGuard::new(set_fail_after_replay_sync); + match exec.execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) { + Err(error) => error, + Ok(_) => panic!("injected replay synchronization failure must roll back"), + } + }; + assert!( + error + .to_string() + .contains("injected Qwen3.5 speculative replay synchronization failure"), + "unexpected speculative replay failure: {error:#}" + ); + assert_eq!(exec.debug_state_summary(), before_state); + + let followup = decode_once(&mut exec, &first_tokens, &cases); + assert_eq!( + followup, + vec![expectations[0].draft_tokens[0]], + "normal decode after replay rollback must match an independent executor" + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_rolls_back_after_graph_commit_sync_error() { + let model_path = model_path(); + let cases = [CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: stable_text_like_prompt(15, 0), + draft_len: 2, + reject_at: None, + }]; + let expectations = build_expectations(&model_path, &cases); + let mut exec = build_executor(&model_path, cases.len()); + let first_tokens = prefill(&mut exec, &cases); + let before_state = exec.debug_state_summary(); + let verify_items = [VerifyStepItem::new( + cases[0].request_id, + vec![ + first_tokens[0], + expectations[0].draft_tokens[0], + expectations[0].draft_tokens[1], + ], + LOGPROBS, + )]; + + let error = { + let _failure = TestFailpointGuard::new(set_fail_after_graph_commit_sync); + match exec.execute_speculative_verify(VerifyPlan { + requests: &verify_items, + }) { + Err(error) => error, + Ok(_) => panic!("injected graph commit synchronization failure must roll back"), + } + }; + assert!( + error + .to_string() + .contains("injected Qwen3.5 speculative graph commit synchronization failure"), + "unexpected speculative graph commit failure: {error:#}" + ); + assert_eq!(exec.debug_state_summary(), before_state); + + let followup = decode_once(&mut exec, &first_tokens, &cases); + assert_eq!( + followup, + vec![expectations[0].draft_tokens[0]], + "normal decode after graph commit rollback must match an independent executor" + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_page_boundary_transactions_match_full_context() { + let model_path = model_path(); + run_speculative_case( + &model_path, + &[ + CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: stable_text_like_prompt(15, 0), + draft_len: 3, + reject_at: None, + }, + CaseSpec { + request_id: RequestId::new(2), + prompt_tokens: stable_text_like_prompt(15, 1), + draft_len: 3, + reject_at: Some(0), + }, + CaseSpec { + request_id: RequestId::new(3), + prompt_tokens: stable_text_like_prompt(16, 2), + draft_len: 4, + reject_at: Some(2), + }, + ], + ); +} + +#[test] +#[ignore = "requires Qwen3.5 weights on a CUDA GPU"] +fn qwen35_speculative_replays_captured_graph_after_commit_and_rollback() { + let model_path = model_path(); + let prompt = stable_text_like_prompt(15, 0); + let oracle = decode_oracle_tokens(&model_path, std::slice::from_ref(&prompt), 11); + let oracle = &oracle[0]; + let cases = [CaseSpec { + request_id: RequestId::new(1), + prompt_tokens: prompt, + draft_len: 0, + reject_at: None, + }]; + + let mut exec = build_executor(&model_path, 1); + let first_tokens = prefill(&mut exec, &cases); + assert_eq!(first_tokens, vec![oracle[0].token]); + assert_eq!( + decode_once(&mut exec, &first_tokens, &cases), + vec![oracle[1].token] + ); + assert!(exec.graph_state.graphs[0].is_captured()); + + let full = [VerifyStepItem::new( + cases[0].request_id, + vec![oracle[1].token, oracle[2].token, oracle[3].token], + LOGPROBS, + )]; + let full_result = exec + .execute_speculative_verify(VerifyPlan { requests: &full }) + .expect("full accept after graph capture"); + assert_eq!(full_result.requests[0].matched_draft_tokens, 2); + assert_eq!( + full_result.requests[0].accepted_tokens, + vec![oracle[2].token, oracle[3].token, oracle[4].token] + ); + assert_eq!( + decode_once(&mut exec, &[oracle[4].token], &cases), + vec![oracle[5].token] + ); + + let partial = [VerifyStepItem::new( + cases[0].request_id, + vec![ + oracle[5].token, + oracle[6].token, + oracle[7].token.wrapping_add(17), + ], + LOGPROBS, + )]; + let partial_result = exec + .execute_speculative_verify(VerifyPlan { requests: &partial }) + .expect("partial accept after graph replay"); + assert_eq!(partial_result.requests[0].matched_draft_tokens, 1); + assert_eq!( + partial_result.requests[0].accepted_tokens, + vec![oracle[6].token, oracle[7].token] + ); + assert_eq!( + decode_once(&mut exec, &[oracle[7].token], &cases), + vec![oracle[8].token] + ); + + let before_rollback = exec.debug_state_summary(); + let rollback = [VerifyStepItem::new( + cases[0].request_id, + vec![oracle[8].token, oracle[9].token, oracle[10].token], + LOGPROBS, + )]; + let error = { + let _failure = TestFailpointGuard::new(set_fail_after_graph_commit_sync); + match exec.execute_speculative_verify(VerifyPlan { + requests: &rollback, + }) { + Err(error) => error, + Ok(_) => panic!("graph commit failpoint must roll back"), + } + }; + assert!( + error + .to_string() + .contains("injected Qwen3.5 speculative graph commit synchronization failure"), + "unexpected graph replay rollback failure: {error:#}" + ); + assert_eq!(exec.debug_state_summary(), before_rollback); + assert!(exec.graph_state.graphs[0].is_captured()); + assert_eq!( + decode_once(&mut exec, &[oracle[8].token], &cases), + vec![oracle[9].token] + ); +} diff --git a/pegainfer-qwen35/src/verify_buffers.rs b/pegainfer-qwen35/src/verify_buffers.rs new file mode 100644 index 000000000..ecedab66f --- /dev/null +++ b/pegainfer-qwen35/src/verify_buffers.rs @@ -0,0 +1,193 @@ +//! Fixed scratch for Qwen3.5 target-side speculative verification. + +use anyhow::Result; +use cudarc::driver::CudaSlice; +use pegainfer_core::tensor::DeviceContext; +use pegainfer_core::tensor::HiddenStates; + +use crate::config::Config35; +use crate::ops::PrefillPagedPlan; +use crate::prefill_buffers::GdrChunkwiseScratch35; + +pub(crate) struct VerifyBuffers35 { + max_batch: usize, + span: usize, + max_rows: usize, + token_ids_h: Vec, + pub(crate) token_ids_d: CudaSlice, + + pub(crate) hidden: HiddenStates, + pub(crate) hidden_next: HiddenStates, + pub(crate) normed: HiddenStates, + pub(crate) attn_results: HiddenStates, + pub(crate) hidden_mid: HiddenStates, + pub(crate) gate_up_out: HiddenStates, + pub(crate) act_out: HiddenStates, + pub(crate) mlp_out: HiddenStates, + pub(crate) logits_normed: HiddenStates, + pub(crate) logits: HiddenStates, + pub(crate) q_full: HiddenStates, + pub(crate) k_full: HiddenStates, + pub(crate) v_full: HiddenStates, + pub(crate) q_prepped: HiddenStates, + pub(crate) attn_out_full: HiddenStates, + + pub(crate) qkv: HiddenStates, + pub(crate) z: HiddenStates, + pub(crate) b_proj: HiddenStates, + pub(crate) a_proj: HiddenStates, + pub(crate) gdr_out: HiddenStates, + pub(crate) normed_gated: HiddenStates, + pub(crate) compact_qkv: HiddenStates, + pub(crate) compact_b: HiddenStates, + pub(crate) compact_a: HiddenStates, + pub(crate) compact_qkv_conv: HiddenStates, + pub(crate) compact_gdr: HiddenStates, + pub(crate) gdr_scratch: GdrChunkwiseScratch35, + + pub(crate) plan: PrefillPagedPlan, + pub(crate) sample: pegainfer_sample::SampleScratch, +} + +impl VerifyBuffers35 { + pub(crate) fn new( + ctx: &DeviceContext, + config: &Config35, + max_batch: usize, + span: usize, + max_total_pages: usize, + ) -> Result { + anyhow::ensure!(max_batch > 0, "Qwen3.5 verify buffers need max_batch > 0"); + anyhow::ensure!(span > 0, "Qwen3.5 verify buffers need span > 0"); + let max_rows = max_batch * span; + let hidden = config.hidden_size; + let q_proj_dim = config.full_attn_q_proj_dim(); + let q_dim = config.full_attn_q_dim(); + let kv_dim = config.full_attn_kv_dim(); + let qkv_dim = config.linear_attn_qkv_dim(); + let z_dim = config.linear_attn_z_dim(); + let group_size = config.num_attention_heads / config.num_key_value_heads; + let max_tiles = max_batch * span * group_size.max(1); + + Ok(Self { + max_batch, + span, + max_rows, + token_ids_h: vec![0; max_rows], + token_ids_d: ctx.stream.alloc_zeros(max_rows)?, + + hidden: HiddenStates::zeros(ctx, hidden, max_rows)?, + hidden_next: HiddenStates::zeros(ctx, hidden, max_rows)?, + normed: HiddenStates::zeros(ctx, hidden, max_rows)?, + attn_results: HiddenStates::zeros(ctx, hidden, max_rows)?, + hidden_mid: HiddenStates::zeros(ctx, hidden, max_rows)?, + gate_up_out: HiddenStates::zeros(ctx, 2 * config.intermediate_size, max_rows)?, + act_out: HiddenStates::zeros(ctx, config.intermediate_size, max_rows)?, + mlp_out: HiddenStates::zeros(ctx, hidden, max_rows)?, + logits_normed: HiddenStates::zeros(ctx, hidden, max_rows)?, + logits: HiddenStates::zeros(ctx, config.selection_vocab, max_rows)?, + q_full: HiddenStates::zeros(ctx, q_proj_dim, max_rows)?, + k_full: HiddenStates::zeros(ctx, kv_dim, max_rows)?, + v_full: HiddenStates::zeros(ctx, kv_dim, max_rows)?, + q_prepped: HiddenStates::zeros(ctx, q_dim, max_rows)?, + attn_out_full: HiddenStates::zeros(ctx, q_dim, max_rows)?, + + qkv: HiddenStates::zeros(ctx, qkv_dim, max_rows)?, + z: HiddenStates::zeros(ctx, z_dim, max_rows)?, + b_proj: HiddenStates::zeros(ctx, config.linear_num_value_heads, max_rows)?, + a_proj: HiddenStates::zeros(ctx, config.linear_num_value_heads, max_rows)?, + gdr_out: HiddenStates::zeros(ctx, z_dim, max_rows)?, + normed_gated: HiddenStates::zeros(ctx, z_dim, max_rows)?, + compact_qkv: HiddenStates::zeros(ctx, qkv_dim, span)?, + compact_b: HiddenStates::zeros(ctx, config.linear_num_value_heads, span)?, + compact_a: HiddenStates::zeros(ctx, config.linear_num_value_heads, span)?, + compact_qkv_conv: HiddenStates::zeros(ctx, qkv_dim, span)?, + compact_gdr: HiddenStates::zeros(ctx, z_dim, span)?, + gdr_scratch: GdrChunkwiseScratch35::new(ctx, config, max_rows)?, + + plan: PrefillPagedPlan::new_preallocated( + ctx, + max_rows, + max_total_pages, + max_batch, + max_tiles, + )?, + sample: pegainfer_sample::SampleScratch::new(ctx, config.selection_vocab, max_rows)?, + }) + } + + pub(crate) fn max_batch(&self) -> usize { + self.max_batch + } + + pub(crate) fn set_rows(&mut self, rows: usize) { + assert!( + rows <= self.max_rows, + "Qwen3.5 verify rows {rows} exceeds capacity {}", + self.max_rows + ); + self.hidden.seq_len = rows; + self.hidden_next.seq_len = rows; + self.normed.seq_len = rows; + self.attn_results.seq_len = rows; + self.hidden_mid.seq_len = rows; + self.gate_up_out.seq_len = rows; + self.act_out.seq_len = rows; + self.mlp_out.seq_len = rows; + self.logits_normed.seq_len = rows; + self.logits.seq_len = rows; + self.q_full.seq_len = rows; + self.k_full.seq_len = rows; + self.v_full.seq_len = rows; + self.q_prepped.seq_len = rows; + self.attn_out_full.seq_len = rows; + self.qkv.seq_len = rows; + self.z.seq_len = rows; + self.b_proj.seq_len = rows; + self.a_proj.seq_len = rows; + self.gdr_out.seq_len = rows; + self.normed_gated.seq_len = rows; + self.gdr_scratch.set_rows(rows); + } + + pub(crate) fn set_compact_rows(&mut self, rows: usize) { + assert!( + rows <= self.span, + "Qwen3.5 compact verify rows {rows} exceeds span {}", + self.span + ); + self.compact_qkv.seq_len = rows; + self.compact_b.seq_len = rows; + self.compact_a.seq_len = rows; + self.compact_qkv_conv.seq_len = rows; + self.compact_gdr.seq_len = rows; + self.gdr_scratch.set_rows(rows); + } + + pub(crate) fn stage_tokens(&mut self, ctx: &DeviceContext, spans: &[&[u32]]) -> Result { + anyhow::ensure!( + spans.len() <= self.max_batch, + "Qwen3.5 verify batch {} exceeds capacity {}", + spans.len(), + self.max_batch + ); + let total_rows: usize = spans.iter().map(|span| span.len()).sum(); + anyhow::ensure!( + total_rows <= self.max_rows, + "Qwen3.5 verify rows {total_rows} exceeds capacity {}", + self.max_rows + ); + self.token_ids_h.clear(); + self.token_ids_h.reserve(total_rows); + for span in spans { + self.token_ids_h.extend_from_slice(span); + } + self.set_rows(total_rows); + if total_rows > 0 { + let mut token_ids_d = self.token_ids_d.slice_mut(..total_rows); + ctx.stream + .memcpy_htod(&self.token_ids_h, &mut token_ids_d)?; + } + Ok(total_rows) + } +}