[XE2] Support MTP of QWEN model - #368
Conversation
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
There was a problem hiding this comment.
Pull request overview
This PR extends the XPU GDN attention implementation to support QWEN MTP / speculative decoding by splitting “non-spec” and “spec” token streams, adding optional token-index remapping to write/read interleaved global buffers, and introducing spec-decoding kernels/paths for both causal_conv1d and gated_delta_rule (including XE2 chunk kernels).
Changes:
- Extend the
gdn_attentionTorch op schema + C++ API withnum_spec_decodesand new optional tensors for spec/non-spec routing and acceptance metadata. - Add token-index remapping (
token_indx) support so kernels can read from / write to interleaved global token slots without host-side gather/scatter. - Add spec-decoding execution paths/kernels for causal_conv1d and gated_delta_rule; update XE2 chunk kernels to support remapped output writes.
Reviewed changes
Copilot reviewed 9 out of 9 changed files in this pull request and generated 6 comments.
Show a summary per file
| File | Description |
|---|---|
| csrc/xpu/torch_bindings.cpp | Updates Torch op schema for gdn_attention to accept spec-decoding inputs. |
| csrc/xpu/ops.h | Updates gdn_attention C++ signature to include spec-decoding parameters and optionals. |
| csrc/xpu/gdn_attn/gdn_attn_interface.cpp | Implements spec/non-spec splitting, validation, and dispatch to spec/non-spec kernel paths (including XE2). |
| csrc/xpu/gdn_attn/causal_conv1d.hpp | Adds token remap support and a spec-decoding kernel path for causal conv1d. |
| csrc/xpu/gdn_attn/gated_delta_rule.hpp | Adds token remap support and a spec-decoding kernel path for gated delta rule. |
| csrc/xpu/gdn_attn/xe_2/chunk_causal_conv1d_xe2.hpp | Adds optional token remap support and actual-token override for XE2 chunk causal conv1d. |
| csrc/xpu/gdn_attn/xe_2/chunk_gated_delta_rule_xe2.h | Extends XE2 chunk gated delta rule API to accept optional token remap pointer. |
| csrc/xpu/gdn_attn/xe_2/chunk_gated_delta_rule_xe2.cpp | Plumbs optional token remap pointer into the XE2 chunk gated delta rule implementation. |
| csrc/xpu/gdn_attn/xe_2/chunk_gated_delta_rule_kernels_xe2.hpp | Remaps output writes using token_indx so chunk outputs land in interleaved global slots. |
Comments suppressed due to low confidence (5)
csrc/xpu/gdn_attn/gdn_attn_interface.cpp:135
- non_spec_state_indices_tensor is a std::optional but is dereferenced without checking has_value(). If the caller passes None while num_prefills + num_decodes > 0, this will crash/UB instead of producing a helpful error.
TORCH_CHECK(
non_spec_state_indices_tensor->is_contiguous(),
"non_spec_state_indices_tensor must be contiguous");
TORCH_CHECK(
csrc/xpu/gdn_attn/gdn_attn_interface.cpp:155
- In the spec-decoding branch (num_spec_decodes > 0), spec_query_start_loc is a std::optional but is dereferenced without checking has_value(). If it is None this is undefined behavior.
if (num_spec_decodes > 0) {
TORCH_CHECK(
spec_query_start_loc->is_contiguous(),
"spec_query_start_loc must be contiguous");
TORCH_CHECK(
csrc/xpu/gdn_attn/gdn_attn_interface.cpp:169
- In the spec-decoding branch (num_spec_decodes > 0), spec_token_indx is a std::optional but is dereferenced without checking has_value(). If it is None this is undefined behavior.
TORCH_CHECK(
spec_token_indx->is_contiguous(), "spec_token_indx must be contiguous");
TORCH_CHECK(
spec_token_indx->dtype() == torch::kInt32,
"spec_token_indx must be of int32 dtype");
csrc/xpu/gdn_attn/gdn_attn_interface.cpp:179
- In the spec-decoding branch (num_spec_decodes > 0), spec_state_indices_tensor is a std::optional but is dereferenced without checking has_value(). If it is None this is undefined behavior.
TORCH_CHECK(
spec_state_indices_tensor->is_contiguous(),
"spec_state_indices_tensor must be contiguous");
TORCH_CHECK(
spec_state_indices_tensor->dtype() == torch::kInt32,
csrc/xpu/gdn_attn/gdn_attn_interface.cpp:196
- In the spec-decoding branch (num_spec_decodes > 0), num_accepted_tokens is a std::optional but is dereferenced without checking has_value(). If it is None this is undefined behavior.
TORCH_CHECK(
num_accepted_tokens->is_contiguous(),
"num_accepted_tokens must be contiguous");
TORCH_CHECK(
num_accepted_tokens->dtype() == torch::kInt32,
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
|
@YangQun1 @wuxun-zhang can you take a look? it will benefit qwen3.5/3.6 I suppose. |
| // Checkpoint the rolling conv state at every step into the cache slot | ||
| // that the next decoding round will read when `num_accepted_tokens == | ||
| // t_local + 1`. Writing each column (not only the last one) keeps the | ||
| // scheduler's per-acceptance rollback consistent: column `t_local` | ||
| // holds the conv state right after consuming spec token `t_local`. | ||
| if (Width > 1) { | ||
| const int save_state_id = | ||
| cache_indices[batch_id * cache_indices_stride_0 + t_local]; | ||
| if (save_state_id != pad_slot_id) { | ||
| T* save_state_ptr = | ||
| conv_states + save_state_id * conv_states_stride_0; | ||
| #pragma unroll | ||
| for (int i = 0; i < Width - 1; ++i) { | ||
| #pragma unroll | ||
| for (int e = 0; e < elems_per_item; ++e) { | ||
| save_state_ptr[i * conv_elems + reordered_elems_id + e] = | ||
| local_input[Width * e + i + 1]; | ||
| } | ||
| } | ||
| } | ||
| } |
There was a problem hiding this comment.
It's true that writing each column 0..K-1 keeps per-acceptance rollback consistent, but FLA's _causal_conv1d_update_kernel provides the same rollback property without using K distinct slots: it only ever touches cache_indices[seq, 0], treats that single slot as a length-state_len rolled history, and uses num_accepted_tokens - 1 as a within-slot column offset.
Why prefer the per-K-slot snapshot layout here over FLA's rolled-history-at-slot-0 layout?
There was a problem hiding this comment.
I think both can be right.
I just followed the way ssm states use cache_indices.
There was a problem hiding this comment.
And, I tested the overhead of causal_conv1d_spec_kernel, launching one time only spend about 5 us on B60 if num spec tokens is 2, so I think we do not need to spend much effort to this kernel if functionality is right.
* [XE2] Support MTP of QWEN model (vllm-project#368) * change API for MTP Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * support MTP Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * add xe2 path Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * refine XE2 path Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * format Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * fix UT Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * refine ref_gdn_attention and add mtp UT Signed-off-by: mayuyuace <qiming1.zhang@intel.com> --------- Signed-off-by: mayuyuace <qiming1.zhang@intel.com> * [XE2] gdn_attention: drop rectangular spec_token assertion for ragged MTP batches The MTP spec-decode path crashed at output length ~96/128 with "Expected spec_token == num_spec_decodes * (num_speculative_tokens + 1)". Root cause: when a request accepts 0 draft tokens the scheduler produces a ragged/short spec batch (num_decode_draft_tokens == -1), and gdn_attn.py deliberately truncates spec_token_indx via min(num_spec_decodes*(n+1), query_start_loc[-1]). So spec_token can be smaller than the rectangular value. This TORCH_CHECK (added in 2b533d2) is a rectangular assumption the upstream CUDA reference kernel (gdn_linear_attn.py forward_cuda) never makes -- it index_selects by the actual spec_token_indx length. The kernel body already works off the real spec_token and walks per-request ranges via spec_query_start_loc (causal_conv1d / gated_delta_rule), so a short spec region is handled correctly. Drop the rectangular check; keep the conservation check non_spec_token + spec_token == num_actual_tokens. Verified: Qwen3.6-35B-A3B fp8 TP2 MTP + FULL_DECODE_ONLY no longer aborts on this assertion (0 hits at output 128 where it previously crashed). Co-Authored-By: Claude <noreply@anthropic.com> --------- Signed-off-by: mayuyuace <qiming1.zhang@intel.com> Co-authored-by: Qiming Zhang <qiming1.zhang@intel.com> Co-authored-by: Claude <noreply@anthropic.com>
Vllm PR: vllm-project/vllm#43565
Tested with Qwen/Qwen3-Next-80B-A3B-Instruct and triton flash attention.
lm_eval results of "num_speculative_tokens": 2:

Throughputs on B60, i/o 1024/512:

For comparation, throughputs without mtp:
