[Mamba] eliminate D2H if tracking mamba states (#20522)
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user