[NPU][Bugfix] fix MTP accuracy regression on Qwen3 hybrid models (#27202)
This commit is contained in:
+33
-7
@@ -157,16 +157,11 @@ class AscendMambaAttnBackendBase(MambaAttnBackendBase):
|
|||||||
)
|
)
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
ssm_state_indices = torch.arange(
|
ssm_state_indices = torch.arange(
|
||||||
len(mamba_indices[: bs - num_padding]) * spec_info.draft_token_num,
|
bs * spec_info.draft_token_num,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=mamba_indices.device,
|
device=mamba_indices.device,
|
||||||
)
|
)
|
||||||
self.state_indices_list_gdn[bs - 1][
|
self.state_indices_list_gdn[bs - 1].copy_(ssm_state_indices)
|
||||||
: len(mamba_indices[: bs - num_padding]) * spec_info.draft_token_num
|
|
||||||
].copy_(ssm_state_indices)
|
|
||||||
self.state_indices_list_gdn[bs - 1][
|
|
||||||
len(mamba_indices[: bs - num_padding]) * spec_info.draft_token_num :
|
|
||||||
] = 0
|
|
||||||
if num_padding == 0:
|
if num_padding == 0:
|
||||||
self.query_start_loc_list[bs - 1].copy_(
|
self.query_start_loc_list[bs - 1].copy_(
|
||||||
self.cached_cuda_graph_verify_query_start_loc[: bs + 1]
|
self.cached_cuda_graph_verify_query_start_loc[: bs + 1]
|
||||||
@@ -269,6 +264,28 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
draft_token_num = intermediate_state_cache.shape[2]
|
draft_token_num = intermediate_state_cache.shape[2]
|
||||||
|
if mamba_track_indices is not None:
|
||||||
|
assert mamba_steps_to_track is not None
|
||||||
|
mamba_track_indices = mamba_track_indices.to(torch.int64)
|
||||||
|
mamba_steps_to_track = mamba_steps_to_track.to(torch.int64)
|
||||||
|
|
||||||
|
move_intermediate_cache(
|
||||||
|
ssm_states,
|
||||||
|
intermediate_state_cache,
|
||||||
|
mamba_track_indices,
|
||||||
|
src_indices_tensor,
|
||||||
|
mamba_steps_to_track,
|
||||||
|
)
|
||||||
|
|
||||||
|
track_mask = mamba_steps_to_track >= 0
|
||||||
|
# Track conv state from the verify-time window before rolling back
|
||||||
|
# the working slot; NPU does not keep per-step conv intermediates.
|
||||||
|
track_indices = mamba_track_indices[track_mask]
|
||||||
|
if track_indices.numel() > 0:
|
||||||
|
conv_states[:, track_indices] = conv_states[
|
||||||
|
:, dst_indices_tensor[track_mask]
|
||||||
|
]
|
||||||
|
|
||||||
if dst_indices_tensor.numel() > 0:
|
if dst_indices_tensor.numel() > 0:
|
||||||
conv_state_rollback(
|
conv_state_rollback(
|
||||||
conv_states,
|
conv_states,
|
||||||
@@ -276,6 +293,15 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
|
|||||||
last_steps,
|
last_steps,
|
||||||
draft_token_num,
|
draft_token_num,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if mamba_track_indices is not None and mamba_track_indices.numel() > 0:
|
||||||
|
conv_state_rollback(
|
||||||
|
conv_states,
|
||||||
|
mamba_track_indices,
|
||||||
|
mamba_steps_to_track,
|
||||||
|
draft_token_num,
|
||||||
|
)
|
||||||
|
|
||||||
return
|
return
|
||||||
|
|
||||||
def update_verify_buffers_to_fill_after_draft(
|
def update_verify_buffers_to_fill_after_draft(
|
||||||
|
|||||||
@@ -109,6 +109,11 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
|
|||||||
self.k_buffer.append(k_buffer_layer)
|
self.k_buffer.append(k_buffer_layer)
|
||||||
self.v_buffer.append(v_buffer_layer)
|
self.v_buffer.append(v_buffer_layer)
|
||||||
|
|
||||||
|
def _init_kv_copy_and_warmup(self):
|
||||||
|
# implementation relies on self.data_strides / self.data_ptrs, which the
|
||||||
|
# NPU paged buffer layout never builds.
|
||||||
|
self._kv_copy_config = None
|
||||||
|
|
||||||
# for disagg
|
# for disagg
|
||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
# layer_num x [seq_len, head_num, head_dim]
|
# layer_num x [seq_len, head_num, head_dim]
|
||||||
|
|||||||
Reference in New Issue
Block a user