[Mamba] eliminate D2H if tracking mamba states (#20522)

Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
Henson-Zh-Ali
2026-04-08 00:17:26 +08:00
committed by GitHub
co-authored by hzh0425
parent 5ae00ecd48
commit 727a182067
3 changed files with 27 additions and 11 deletions
@@ -217,6 +217,11 @@ class MambaAttnBackendBase(AttentionBackend):
else:
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}")
has_mamba_track_mask = bool(
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
)
return ForwardMetadata(
query_start_loc=query_start_loc,
mamba_cache_indices=mamba_cache_indices,
@@ -228,6 +233,7 @@ class MambaAttnBackendBase(AttentionBackend):
track_ssm_h_dst=track_ssm_h_dst,
track_ssm_final_src=track_ssm_final_src,
track_ssm_final_dst=track_ssm_final_dst,
has_mamba_track_mask=has_mamba_track_mask,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
@@ -613,10 +619,7 @@ class MambaAttnBackendBase(AttentionBackend):
Note: Conv state tracking for extend is handled separately via gather operations
using indices computed by `_init_track_conv_indices`.
"""
if (
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
if forward_metadata.has_mamba_track_mask:
h = h.squeeze(0)
if forward_metadata.track_ssm_h_src.numel() > 0:
@@ -256,6 +256,18 @@ class GDNAttnBackend(MambaAttnBackendBase):
prefill_backend = get_linear_attn_prefill_backend()
self.kernel_dispatcher = GDNKernelDispatcher(decode_backend, prefill_backend)
def init_forward_metadata(self, forward_batch: ForwardBatch):
super().init_forward_metadata(forward_batch)
if self.forward_metadata.has_mamba_track_mask:
self.forward_metadata.mamba_track_mask_indices = (
forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0]
)
self.forward_metadata.conv_states_mask_indices = (
forward_batch.mamba_track_indices[
self.forward_metadata.mamba_track_mask_indices
]
)
def forward_decode(
self,
layer: RadixLinearAttention,
@@ -394,16 +406,13 @@ class GDNAttnBackend(MambaAttnBackendBase):
mixed_qkv = mixed_qkv_processed.transpose(1, 2).view(seq_len, -1)
else:
mixed_qkv = mixed_qkv.transpose(0, 1)
if (
forward_batch.mamba_track_mask is not None
and forward_batch.mamba_track_mask.any()
):
conv_dst = forward_batch.mamba_track_indices
if forward_metadata.has_mamba_track_mask:
mixed_qkv_to_track = mixed_qkv[
:, forward_metadata.track_conv_indices
].transpose(0, 1)
mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0]
conv_states[conv_dst[mask_indices]] = mixed_qkv_to_track
conv_states[forward_metadata.conv_states_mask_indices] = (
mixed_qkv_to_track
)
mixed_qkv = causal_conv1d_fn(
mixed_qkv,
@@ -41,6 +41,10 @@ class ForwardMetadata:
is_target_verify: bool = False
draft_token_num: int = 1
has_mamba_track_mask: bool = False
mamba_track_mask_indices: Optional[torch.Tensor] = None
conv_states_mask_indices: Optional[torch.Tensor] = None
@dataclass(kw_only=True)
class Mamba2Metadata(ForwardMetadata):