Fix mamba radix cache ssm state indexing (#37836)
Co-authored-by: adityakamat24 <adityakamat007@gmail.com>
This commit is contained in:
@@ -25,6 +25,33 @@ def is_int_pow_2(n):
|
||||
return isinstance(n, int) and n > 0 and (n & (n - 1)) == 0
|
||||
|
||||
|
||||
def _track_states_at(
|
||||
B,
|
||||
x,
|
||||
dt,
|
||||
dA_cumsum,
|
||||
states,
|
||||
cu_seqlens,
|
||||
initial_states,
|
||||
track_seq_idx,
|
||||
track_end_locs,
|
||||
):
|
||||
starts = cu_seqlens[track_seq_idx]
|
||||
ends = track_end_locs.to(device=starts.device, dtype=starts.dtype)
|
||||
bounds = torch.stack([starts, ends], dim=1).flatten().contiguous()
|
||||
if initial_states is not None:
|
||||
initial_states = initial_states[track_seq_idx].repeat_interleave(2, dim=0)[:-1]
|
||||
return chunk_state_varlen(
|
||||
B.squeeze(0),
|
||||
x.squeeze(0),
|
||||
dt.squeeze(0),
|
||||
dA_cumsum.squeeze(0),
|
||||
bounds,
|
||||
states.squeeze(0),
|
||||
initial_states=initial_states,
|
||||
)[::2].unsqueeze(0)
|
||||
|
||||
|
||||
def _mamba_chunk_scan_combined_fwd(
|
||||
x,
|
||||
dt,
|
||||
@@ -44,6 +71,8 @@ def _mamba_chunk_scan_combined_fwd(
|
||||
dt_limit=(0.0, float("inf")),
|
||||
state_dtype=None,
|
||||
out=None,
|
||||
track_seq_idx=None,
|
||||
track_end_locs=None,
|
||||
):
|
||||
assert is_int_pow_2(chunk_size), "chunk_size must be integer power of 2"
|
||||
batch, seqlen, nheads, headdim = x.shape
|
||||
@@ -175,7 +204,24 @@ def _mamba_chunk_scan_combined_fwd(
|
||||
states.squeeze(0),
|
||||
initial_states=initial_states,
|
||||
)
|
||||
return out_x, dt, dA_cumsum, states, final_states, varlen_states
|
||||
track_states = None
|
||||
if (
|
||||
track_seq_idx is not None
|
||||
and track_end_locs is not None
|
||||
and track_end_locs.numel() > 0
|
||||
):
|
||||
track_states = _track_states_at(
|
||||
B,
|
||||
x,
|
||||
dt,
|
||||
dA_cumsum,
|
||||
states,
|
||||
cu_seqlens,
|
||||
initial_states,
|
||||
track_seq_idx,
|
||||
track_end_locs,
|
||||
)
|
||||
return out_x, dt, dA_cumsum, states, final_states, varlen_states, track_states
|
||||
|
||||
|
||||
def mamba_chunk_scan_combined(
|
||||
@@ -200,6 +246,9 @@ def mamba_chunk_scan_combined(
|
||||
return_varlen_states=False,
|
||||
return_intermediate_states=False,
|
||||
state_dtype=None,
|
||||
return_track_states=False,
|
||||
track_seq_idx=None,
|
||||
track_end_locs=None,
|
||||
):
|
||||
"""
|
||||
Argument:
|
||||
@@ -246,8 +295,17 @@ def mamba_chunk_scan_combined(
|
||||
dt_limit=dt_limit,
|
||||
out=out,
|
||||
state_dtype=state_dtype,
|
||||
track_seq_idx=track_seq_idx,
|
||||
track_end_locs=track_end_locs,
|
||||
)
|
||||
)
|
||||
if return_track_states:
|
||||
assert return_varlen_states, (
|
||||
"return_track_states requires return_varlen_states (cu_seqlens mode)"
|
||||
)
|
||||
# `states` rides along: a tracked position on the grid is read from it
|
||||
# directly, which stays bit-exact against a recompute-free run.
|
||||
return states, rest[0], rest[1]
|
||||
if return_intermediate_states:
|
||||
if return_varlen_states:
|
||||
varlen_states = rest[0]
|
||||
|
||||
@@ -122,6 +122,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
track_ssm_h_dst = None
|
||||
track_ssm_final_src = None
|
||||
track_ssm_final_dst = None
|
||||
track_ssm_seq_idx = None
|
||||
track_ssm_end_locs = None
|
||||
track_ssm_recompute_dst = None
|
||||
|
||||
mamba_cache_indices = self.req_to_token_pool.get_mamba_indices(
|
||||
forward_batch.req_pool_indices
|
||||
@@ -252,6 +255,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
track_ssm_h_dst,
|
||||
track_ssm_final_src,
|
||||
track_ssm_final_dst,
|
||||
track_ssm_seq_idx,
|
||||
track_ssm_end_locs,
|
||||
track_ssm_recompute_dst,
|
||||
) = self._init_track_ssm_indices(mamba_cache_indices, forward_batch)
|
||||
else:
|
||||
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode=}")
|
||||
@@ -270,6 +276,9 @@ 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,
|
||||
track_ssm_seq_idx=track_ssm_seq_idx,
|
||||
track_ssm_end_locs=track_ssm_end_locs,
|
||||
track_ssm_recompute_dst=track_ssm_recompute_dst,
|
||||
has_mamba_track_mask=has_mamba_track_mask,
|
||||
replayssm_write_pos=replayssm_write_pos,
|
||||
replayssm_force_flush=replayssm_force_flush,
|
||||
@@ -358,17 +367,8 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
mamba_track_seqlens = forward_batch.mamba_track_seqlens.cpu()
|
||||
prefix_lens = forward_batch.extend_prefix_lens.cpu()
|
||||
|
||||
if isinstance(self, Mamba2AttnBackend):
|
||||
num_h_states = extend_seq_lens // state_chunk_size
|
||||
else:
|
||||
num_h_states = (extend_seq_lens - 1) // state_chunk_size + 1
|
||||
|
||||
track_ssm_src_offset = torch.zeros_like(num_h_states)
|
||||
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
|
||||
|
||||
lens_to_track = mamba_track_seqlens - prefix_lens
|
||||
lens_masked = lens_to_track[mamba_track_mask]
|
||||
offset_masked = track_ssm_src_offset[mamba_track_mask]
|
||||
dst_masked = mamba_track_indices[mamba_track_mask]
|
||||
|
||||
is_aligned = (lens_masked % state_chunk_size) == 0
|
||||
@@ -379,16 +379,50 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
|
||||
# Unaligned: intermediate state from h.
|
||||
not_aligned = ~is_aligned
|
||||
track_ssm_h_src = offset_masked[not_aligned] + (
|
||||
lens_masked[not_aligned] // state_chunk_size
|
||||
)
|
||||
track_ssm_h_dst = dst_masked[not_aligned]
|
||||
track_ssm_seq_idx = None
|
||||
track_ssm_end_locs = None
|
||||
track_ssm_recompute_dst = None
|
||||
|
||||
if isinstance(self, Mamba2AttnBackend):
|
||||
# `h` is one chunk grid over the whole flattened batch, so the wanted
|
||||
# state is in it only when query_start_loc[i] is chunk-aligned.
|
||||
query_start_loc = torch.zeros_like(extend_seq_lens)
|
||||
query_start_loc[1:] = torch.cumsum(extend_seq_lens[:-1], dim=0)
|
||||
aligned_len = (
|
||||
lens_masked[not_aligned] // state_chunk_size
|
||||
) * state_chunk_size
|
||||
seq_idx_unaligned = torch.nonzero(mamba_track_mask, as_tuple=True)[0][
|
||||
not_aligned
|
||||
]
|
||||
end_locs = query_start_loc[seq_idx_unaligned] + aligned_len
|
||||
on_grid = (end_locs % state_chunk_size) == 0
|
||||
|
||||
track_ssm_h_src = end_locs[on_grid] // state_chunk_size
|
||||
track_ssm_h_dst = track_ssm_h_dst[on_grid]
|
||||
|
||||
track_ssm_seq_idx = seq_idx_unaligned[~on_grid]
|
||||
track_ssm_end_locs = end_locs[~on_grid]
|
||||
track_ssm_recompute_dst = dst_masked[not_aligned][~on_grid]
|
||||
else:
|
||||
num_h_states = (extend_seq_lens - 1) // state_chunk_size + 1
|
||||
track_ssm_src_offset = torch.zeros_like(num_h_states)
|
||||
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
|
||||
track_ssm_h_src = track_ssm_src_offset[mamba_track_mask][not_aligned] + (
|
||||
lens_masked[not_aligned] // state_chunk_size
|
||||
)
|
||||
|
||||
def to_device(t):
|
||||
return None if t is None else t.to(self.device, non_blocking=True)
|
||||
|
||||
return (
|
||||
track_ssm_h_src.to(self.device, non_blocking=True),
|
||||
track_ssm_h_dst.to(self.device, non_blocking=True),
|
||||
track_ssm_final_src.to(self.device, non_blocking=True),
|
||||
track_ssm_final_dst.to(self.device, non_blocking=True),
|
||||
to_device(track_ssm_h_src),
|
||||
to_device(track_ssm_h_dst),
|
||||
to_device(track_ssm_final_src),
|
||||
to_device(track_ssm_final_dst),
|
||||
to_device(track_ssm_seq_idx),
|
||||
to_device(track_ssm_end_locs),
|
||||
to_device(track_ssm_recompute_dst),
|
||||
)
|
||||
|
||||
def init_forward_metadata_capture_cpu_graph(
|
||||
@@ -847,6 +881,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
h: Optional[torch.Tensor],
|
||||
ssm_states: torch.Tensor,
|
||||
forward_metadata: ForwardMetadata,
|
||||
track_states: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""Copy extend SSM state at the last chunk boundary to track slots (source
|
||||
depends on chunk alignment; see `_init_track_ssm_indices`)."""
|
||||
@@ -859,6 +894,14 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
ssm_states[forward_metadata.track_ssm_h_dst] = h[
|
||||
forward_metadata.track_ssm_h_src
|
||||
].to(ssm_states.dtype, copy=False)
|
||||
if (
|
||||
forward_metadata.track_ssm_recompute_dst is not None
|
||||
and forward_metadata.track_ssm_recompute_dst.numel() > 0
|
||||
):
|
||||
assert track_states is not None
|
||||
ssm_states[forward_metadata.track_ssm_recompute_dst] = (
|
||||
track_states.squeeze(0).to(ssm_states.dtype, copy=False)
|
||||
)
|
||||
if forward_metadata.track_ssm_final_src.numel() > 0:
|
||||
ssm_states[forward_metadata.track_ssm_final_dst] = ssm_states[
|
||||
forward_metadata.track_ssm_final_src
|
||||
@@ -941,7 +984,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
||||
use_triton_causal_conv or get_memory().enable_page_major_kv_layout
|
||||
)
|
||||
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||
mixer_out, intermediate_states = mixer.forward(
|
||||
mixer_out, intermediate_states, track_states = mixer.forward(
|
||||
hidden_states=hidden_states,
|
||||
output=output,
|
||||
layer_cache=layer_cache,
|
||||
@@ -957,6 +1000,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
||||
intermediate_states,
|
||||
layer_cache.temporal,
|
||||
self.forward_metadata,
|
||||
track_states=track_states,
|
||||
)
|
||||
|
||||
if self.forward_metadata.num_decodes > 0:
|
||||
|
||||
@@ -461,6 +461,7 @@ class MambaMixer2(torch.nn.Module):
|
||||
conv_state = layer_cache.conv[0]
|
||||
ssm_state = layer_cache.temporal
|
||||
intermediate_states = None
|
||||
track_states = None
|
||||
|
||||
query_start_loc = metadata.query_start_loc
|
||||
|
||||
@@ -600,7 +601,7 @@ class MambaMixer2(torch.nn.Module):
|
||||
)
|
||||
|
||||
# NOTE: final output is an in-place update of out tensor
|
||||
intermediate_states, varlen_state = mamba_chunk_scan_combined(
|
||||
intermediate_states, varlen_state, track_states = mamba_chunk_scan_combined(
|
||||
hidden_states_p.view(
|
||||
1, num_prefill_tokens, local_num_heads, self.head_dim
|
||||
),
|
||||
@@ -619,7 +620,9 @@ class MambaMixer2(torch.nn.Module):
|
||||
initial_states=initial_states,
|
||||
return_varlen_states=True,
|
||||
return_final_states=False,
|
||||
return_intermediate_states=True,
|
||||
return_track_states=True,
|
||||
track_seq_idx=metadata.track_ssm_seq_idx,
|
||||
track_end_locs=metadata.track_ssm_end_locs,
|
||||
dt_softplus=True,
|
||||
dt_limit=(0.0, float("inf")),
|
||||
out=preallocated_ssm_out_p.view(
|
||||
@@ -763,7 +766,7 @@ class MambaMixer2(torch.nn.Module):
|
||||
if output is not None:
|
||||
output[:padded_num_tokens].copy_(mixer_out)
|
||||
|
||||
return mixer_out, intermediate_states
|
||||
return mixer_out, intermediate_states, track_states
|
||||
|
||||
@property
|
||||
def mamba_type(self) -> str:
|
||||
|
||||
@@ -58,6 +58,9 @@ class ForwardMetadata:
|
||||
state_checkpoint_cu_starts: Optional[torch.Tensor] = None
|
||||
num_state_checkpoints: int = 0
|
||||
state_checkpoint_every_n_tokens: int = 0
|
||||
track_ssm_seq_idx: Optional[torch.Tensor] = None
|
||||
track_ssm_end_locs: Optional[torch.Tensor] = None
|
||||
track_ssm_recompute_dst: Optional[torch.Tensor] = None
|
||||
|
||||
is_target_verify: bool = False
|
||||
draft_token_num: int = 1
|
||||
@@ -202,6 +205,9 @@ class Mamba2Metadata(ForwardMetadata):
|
||||
track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
|
||||
track_ssm_final_src=forward_metadata.track_ssm_final_src,
|
||||
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
|
||||
track_ssm_seq_idx=forward_metadata.track_ssm_seq_idx,
|
||||
track_ssm_end_locs=forward_metadata.track_ssm_end_locs,
|
||||
track_ssm_recompute_dst=forward_metadata.track_ssm_recompute_dst,
|
||||
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
|
||||
num_decodes=len(seq_lens) if num_decodes is None else num_decodes,
|
||||
num_prefills=0,
|
||||
@@ -303,6 +309,9 @@ class Mamba2Metadata(ForwardMetadata):
|
||||
track_ssm_h_dst=forward_metadata.track_ssm_h_dst,
|
||||
track_ssm_final_src=forward_metadata.track_ssm_final_src,
|
||||
track_ssm_final_dst=forward_metadata.track_ssm_final_dst,
|
||||
track_ssm_seq_idx=forward_metadata.track_ssm_seq_idx,
|
||||
track_ssm_end_locs=forward_metadata.track_ssm_end_locs,
|
||||
track_ssm_recompute_dst=forward_metadata.track_ssm_recompute_dst,
|
||||
has_mamba_track_mask=forward_metadata.has_mamba_track_mask,
|
||||
mamba_track_mask_indices=mamba_track_mask_indices,
|
||||
conv_states_mask_indices=conv_states_mask_indices,
|
||||
|
||||
Reference in New Issue
Block a user