Skip to content

Commit d638401

Browse files
committed
perf: bypass redundant qwen3.5 decode projection packing.
- Use separate qkv/z/b/a projections directly in decode and speculative verification. - Preserve the packed qkvz/ba fused-split path for Qwen3Next. - Clarify the packed-token padding helper and its shape contract.
1 parent 653031e commit d638401

4 files changed

Lines changed: 41 additions & 34 deletions

File tree

xllm/core/layers/npu_torch/qwen3_5_gated_delta_net.cpp

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -151,17 +151,17 @@ Qwen3_5GatedDeltaNetImpl::project_flat_inputs(
151151

152152
std::optional<
153153
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>>
154-
Qwen3_5GatedDeltaNetImpl::project_prefill_split_inputs(
154+
Qwen3_5GatedDeltaNetImpl::project_split_inputs(
155155
const torch::Tensor& hidden_states,
156156
const AttentionMetadata& attn_metadata) {
157-
auto qkv = reshape_qkvz_with_pad(attn_metadata,
158-
in_proj_qkv_->forward(hidden_states));
159-
auto z_proj =
160-
reshape_qkvz_with_pad(attn_metadata, in_proj_z_->forward(hidden_states));
161-
auto b_proj =
162-
reshape_qkvz_with_pad(attn_metadata, in_proj_b_->forward(hidden_states));
163-
auto a_proj =
164-
reshape_qkvz_with_pad(attn_metadata, in_proj_a_->forward(hidden_states));
157+
auto qkv = reshape_projected_tokens_with_pad(
158+
attn_metadata, in_proj_qkv_->forward(hidden_states));
159+
auto z_proj = reshape_projected_tokens_with_pad(
160+
attn_metadata, in_proj_z_->forward(hidden_states));
161+
auto b_proj = reshape_projected_tokens_with_pad(
162+
attn_metadata, in_proj_b_->forward(hidden_states));
163+
auto a_proj = reshape_projected_tokens_with_pad(
164+
attn_metadata, in_proj_a_->forward(hidden_states));
165165

166166
const int64_t batch_size = qkv.size(0);
167167
const int64_t seq_len = qkv.size(1);

xllm/core/layers/npu_torch/qwen3_5_gated_delta_net.h

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -42,8 +42,8 @@ class Qwen3_5GatedDeltaNetImpl : public Qwen3NextGatedDeltaNetImpl {
4242
const torch::Tensor& hidden_states) override;
4343
std::optional<
4444
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>>
45-
project_prefill_split_inputs(const torch::Tensor& hidden_states,
46-
const AttentionMetadata& attn_metadata) override;
45+
project_split_inputs(const torch::Tensor& hidden_states,
46+
const AttentionMetadata& attn_metadata) override;
4747
bool use_fla_ssm_state_layout() const override { return true; }
4848

4949
void load_projection_state_dict(const StateDict& state_dict) override;

xllm/core/layers/npu_torch/qwen3_gated_delta_net_base.cpp

Lines changed: 19 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -523,8 +523,8 @@ Qwen3GatedDeltaNetBaseImpl::project_padded_inputs(
523523
const AttentionMetadata& attn_metadata) {
524524
if (attn_metadata.is_prefill || attn_metadata.is_chunked_prefill) {
525525
auto [qkvz_flat, ba_flat] = project_flat_inputs(hidden_states);
526-
return {reshape_qkvz_with_pad(attn_metadata, qkvz_flat),
527-
reshape_qkvz_with_pad(attn_metadata, ba_flat)};
526+
return {reshape_projected_tokens_with_pad(attn_metadata, qkvz_flat),
527+
reshape_projected_tokens_with_pad(attn_metadata, ba_flat)};
528528
}
529529
return project_decode_inputs(hidden_states);
530530
}
@@ -543,12 +543,12 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::forward(
543543
int64_t batch_size = 0;
544544
int64_t seq_len = 0;
545545

546-
auto prefill_split_inputs =
547-
(!use_spec_verify && is_any_prefill)
548-
? project_prefill_split_inputs(hidden_states, attn_metadata)
549-
: std::nullopt;
550-
if (prefill_split_inputs.has_value()) {
551-
std::tie(mixed_qkv, z, b, a) = prefill_split_inputs.value();
546+
// Qwen3.5 stores qkv, z, b, and a as separate projection weights, so it can
547+
// use their outputs directly in every forward mode. Qwen3Next stores qkvz
548+
// and ba as packed weights and uses the fused-split fallback below.
549+
auto split_inputs = project_split_inputs(hidden_states, attn_metadata);
550+
if (split_inputs.has_value()) {
551+
std::tie(mixed_qkv, z, b, a) = split_inputs.value();
552552
batch_size = mixed_qkv.size(0);
553553
seq_len = mixed_qkv.size(1);
554554
} else {
@@ -611,7 +611,7 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::forward(
611611
xllm::npu::kCausalConv1dGraphPadSlotId,
612612
xllm::npu::kCausalConv1dRunModeForward);
613613

614-
mixed_qkv = reshape_qkvz_with_pad(attn_metadata, mixed_qkv);
614+
mixed_qkv = reshape_projected_tokens_with_pad(attn_metadata, mixed_qkv);
615615
mixed_qkv = mixed_qkv.transpose(1, 2);
616616
} else {
617617
if (use_spec_verify) {
@@ -691,7 +691,7 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::forward(
691691
}
692692
}
693693
}
694-
mixed_qkv = reshape_qkvz_with_pad(attn_metadata, mixed_qkv);
694+
mixed_qkv = reshape_projected_tokens_with_pad(attn_metadata, mixed_qkv);
695695
mixed_qkv = mixed_qkv.transpose(1, 2);
696696
}
697697
const bool fla_ssm_state_layout = use_fla_ssm_state_layout();
@@ -938,9 +938,9 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::forward(
938938
auto rearranged_norm =
939939
norm_out.reshape({norm_out.size(0), norm_out.size(1) * norm_out.size(2)});
940940
rearranged_norm = reshape_qkvz_unpad(attn_metadata, rearranged_norm);
941-
// For chunked prefill or spec verify, reshape_qkvz_with_pad may pad each
942-
// batch to max_len, causing output tokens > original_num_tokens. We need to
943-
// slice back to original_num_tokens to match residual shape for add_rms_norm.
941+
// For chunked prefill or spec verify, reshape_projected_tokens_with_pad may
942+
// pad each batch to max_len, causing output tokens > original_num_tokens. We
943+
// need to slice back to original_num_tokens to match the residual shape.
944944
if (rearranged_norm.size(0) > original_num_tokens) {
945945
// Slice excess padding tokens
946946
rearranged_norm =
@@ -999,9 +999,9 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::get_linear_state_indices(
999999
torch::TensorOptions().dtype(torch::kInt).device(device));
10001000
}
10011001

1002-
torch::Tensor Qwen3GatedDeltaNetBaseImpl::reshape_qkvz_with_pad(
1002+
torch::Tensor Qwen3GatedDeltaNetBaseImpl::reshape_projected_tokens_with_pad(
10031003
const AttentionMetadata& attn_metadata,
1004-
const torch::Tensor& qkvz) const {
1004+
const torch::Tensor& projected_tokens) const {
10051005
const bool has_host_lens = !attn_metadata.q_seq_lens_vec.empty();
10061006
int64_t bs = has_host_lens
10071007
? static_cast<int64_t>(attn_metadata.q_seq_lens_vec.size())
@@ -1011,11 +1011,11 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::reshape_qkvz_with_pad(
10111011
const bool need_padding =
10121012
attn_metadata.is_prefill || attn_metadata.is_chunked_prefill;
10131013
if (!need_padding) {
1014-
return qkvz.view({bs, -1, qkvz.size(-1)});
1014+
return projected_tokens.view({bs, -1, projected_tokens.size(-1)});
10151015
}
10161016
if (has_host_lens && bs == 1 && attn_metadata.q_seq_lens_vec[0] == max_len &&
1017-
qkvz.dim() == 2 && qkvz.size(0) == max_len) {
1018-
return qkvz.view({1, max_len, qkvz.size(-1)});
1017+
projected_tokens.dim() == 2 && projected_tokens.size(0) == max_len) {
1018+
return projected_tokens.view({1, max_len, projected_tokens.size(-1)});
10191019
}
10201020
std::vector<torch::Tensor> batches;
10211021
batches.reserve(bs);
@@ -1024,7 +1024,7 @@ torch::Tensor Qwen3GatedDeltaNetBaseImpl::reshape_qkvz_with_pad(
10241024
int64_t cur_len = has_host_lens ? attn_metadata.q_seq_lens_vec[b]
10251025
: start_loc[b].template item<int64_t>();
10261026
torch::Tensor batch =
1027-
qkvz.slice(/*dim=*/0, idx, idx + cur_len).contiguous();
1027+
projected_tokens.slice(/*dim=*/0, idx, idx + cur_len).contiguous();
10281028
idx = idx + cur_len;
10291029
if (batch.size(0) != max_len) {
10301030
batch = batch.size(0) > max_len

xllm/core/layers/npu_torch/qwen3_gated_delta_net_base.h

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -56,10 +56,13 @@ class Qwen3GatedDeltaNetBaseImpl : public torch::nn::Module {
5656
const torch::Tensor& hidden_states) = 0;
5757
virtual std::pair<torch::Tensor, torch::Tensor> project_flat_inputs(
5858
const torch::Tensor& hidden_states) = 0;
59+
// Qwen3.5 overrides this to project and reshape its separate qkv/z/b/a
60+
// weights in every forward mode. Qwen3Next keeps qkvz/ba packed and returns
61+
// nullopt to select the fused-split fallback.
5962
virtual std::optional<
6063
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor>>
61-
project_prefill_split_inputs(const torch::Tensor& hidden_states,
62-
const AttentionMetadata& attn_metadata) {
64+
project_split_inputs(const torch::Tensor& hidden_states,
65+
const AttentionMetadata& attn_metadata) {
6366
return std::nullopt;
6467
}
6568
virtual bool use_fla_ssm_state_layout() const { return false; }
@@ -77,8 +80,12 @@ class Qwen3GatedDeltaNetBaseImpl : public torch::nn::Module {
7780
torch::Tensor reshape_qkvz_unpad(const AttentionMetadata& attn_metadata,
7881
const torch::Tensor& padded_qkvz) const;
7982

80-
torch::Tensor reshape_qkvz_with_pad(const AttentionMetadata& attn_metadata,
81-
const torch::Tensor& qkvz) const;
83+
// Projection outputs are packed as [total_tokens, dim], while GDN kernels
84+
// consume dense [batch, max_query_len, dim] tensors. Split the packed tokens
85+
// by query length and pad each sequence before entering the kernels.
86+
torch::Tensor reshape_projected_tokens_with_pad(
87+
const AttentionMetadata& attn_metadata,
88+
const torch::Tensor& projected_tokens) const;
8289

8390
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> process_mixed_qkv(
8491
torch::Tensor& mixed_qkv) const;

0 commit comments

Comments
 (0)