Skip to content

[XE2] Support MTP of QWEN model - #368

Merged
mayuyuace merged 7 commits into
mainfrom
qiming/mtp
May 26, 2026
Merged

mayuyuace merged 7 commits into
mainfrom
qiming/mtp

Conversation

@mayuyuace

@mayuyuace mayuyuace commented May 25, 2026 •

Copy link
Copy Markdown
Collaborator

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:
image

Throughputs on B60, i/o 1024/512:
image

For comparation, throughputs without mtp:
image

mayuyuace added 5 commits May 25, 2026 01:26
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>
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
Copilot AI review requested due to automatic review settings May 25, 2026 04:26

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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_attention Torch op schema + C++ API with num_spec_decodes and 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.

Comment thread csrc/xpu/gdn_attn/gdn_attn_interface.cpp
Comment thread csrc/xpu/gdn_attn/gdn_attn_interface.cpp
Comment thread csrc/xpu/gdn_attn/gdn_attn_interface.cpp
Comment thread csrc/xpu/gdn_attn/gated_delta_rule.hpp
Comment thread csrc/xpu/gdn_attn/causal_conv1d.hpp
Comment thread csrc/xpu/torch_bindings.cpp
mayuyuace added 2 commits May 25, 2026 04:33
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
Signed-off-by: mayuyuace <qiming1.zhang@intel.com>
@jikunshang

Copy link
Copy Markdown
Member

@YangQun1 @wuxun-zhang can you take a look? it will benefit qwen3.5/3.6 I suppose.

@mayuyuace
mayuyuace merged commit 8037256 into main May 26, 2026
9 checks passed
Comment on lines +803 to +823
// 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];
}
}
}
}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think both can be right.
I just followed the way ssm states use cache_indices.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@mayuyuace
mayuyuace deleted the qiming/mtp branch May 27, 2026 02:22
hzjane added a commit to analytics-zoo/vllm-xpu-kernels that referenced this pull request Jul 9, 2026
* [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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants