@@ -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
0 commit comments