[Spec] Mamba scatter cleanup; fix multi-layer positional bug; dflash naming (#25029)

This commit is contained in:
Liangsheng Yin
2026-05-11 20:36:50 -07:00
committed by GitHub
parent 10375a1037
commit f526e3fa27
10 changed files with 49 additions and 51 deletions
@@ -219,7 +219,7 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
def update_mamba_state_after_mtp_verify( def update_mamba_state_after_mtp_verify(
self, self,
accept_steps: torch.Tensor, last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor], mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor], mamba_steps_to_track: Optional[torch.Tensor],
model, model,
@@ -233,7 +233,7 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
- index_select kernel launches - index_select kernel launches
- nonzero kernel launches - nonzero kernel launches
""" """
request_number = accept_steps.shape[0] request_number = last_correct_step_indices.shape[0]
state_indices_tensor = ( state_indices_tensor = (
self.linear_attn_backend.forward_metadata.mamba_cache_indices[ self.linear_attn_backend.forward_metadata.mamba_cache_indices[
@@ -254,7 +254,7 @@ class AscendHybridLinearAttnBackend(HybridLinearAttnBackend):
device=dst_indices_tensor.device, device=dst_indices_tensor.device,
dtype=torch.int64, dtype=torch.int64,
) )
last_steps = accept_steps.to(torch.int64) # [N] last_steps = last_correct_step_indices.to(torch.int64) # [N]
move_intermediate_cache( move_intermediate_cache(
ssm_states, ssm_states,
@@ -936,7 +936,7 @@ class HybridLinearAttnBackend(AttentionBackend):
def update_mamba_state_after_mtp_verify( def update_mamba_state_after_mtp_verify(
self, self,
accept_steps: torch.Tensor, last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor], mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor], mamba_steps_to_track: Optional[torch.Tensor],
model, model,
@@ -950,7 +950,7 @@ class HybridLinearAttnBackend(AttentionBackend):
- index_select kernel launches - index_select kernel launches
- nonzero kernel launches - nonzero kernel launches
""" """
request_number = accept_steps.shape[0] request_number = last_correct_step_indices.shape[0]
state_indices_tensor = ( state_indices_tensor = (
self.linear_attn_backend.forward_metadata.mamba_cache_indices[ self.linear_attn_backend.forward_metadata.mamba_cache_indices[
@@ -973,13 +973,13 @@ class HybridLinearAttnBackend(AttentionBackend):
ssm_states, ssm_states,
intermediate_state_cache, intermediate_state_cache,
state_indices_tensor, state_indices_tensor,
accept_steps, last_correct_step_indices,
) )
fused_mamba_state_scatter_with_mask( fused_mamba_state_scatter_with_mask(
conv_states, conv_states,
intermediate_conv_window_cache, intermediate_conv_window_cache,
state_indices_tensor, state_indices_tensor,
accept_steps, last_correct_step_indices,
) )
# Track indices used for tracking mamba states for prefix cache # Track indices used for tracking mamba states for prefix cache
@@ -17,7 +17,7 @@ def _fused_mamba_state_scatter_with_mask_kernel(
dst_ptr, dst_ptr,
# Raw index arrays (before index_select) # Raw index arrays (before index_select)
dst_indices_raw_ptr, # [total_requests] - state_indices_tensor dst_indices_raw_ptr, # [total_requests] - state_indices_tensor
step_indices_raw_ptr, # [total_requests] - accept_steps or mamba_steps_to_track step_indices_raw_ptr, # [total_requests] - last_correct_step_indices or mamba_steps_to_track
elem_per_entry: tl.constexpr, elem_per_entry: tl.constexpr,
src_layer_stride, src_layer_stride,
src_req_stride, src_req_stride,
@@ -106,7 +106,7 @@ class SchedulerMetricsMixin:
# Cumulative spec-decoding counters (reset every decode_log_interval). # Cumulative spec-decoding counters (reset every decode_log_interval).
# Each update adds (num_correct_drafts + bs, bs). # Each update adds (num_correct_drafts + bs, bs).
# `*_accepted_tokens` = drafts + bonus; `*_accepted_drafts` = drafts-only. # `*_accept_tokens` = drafts + bonus; `*_correct_drafts` = drafts-only.
self.spec_num_accept_tokens = 0 # per-log-interval self.spec_num_accept_tokens = 0 # per-log-interval
self.spec_num_forward_ct = 0 self.spec_num_forward_ct = 0
self.spec_total_num_accept_tokens = 0 # lifetime self.spec_total_num_accept_tokens = 0 # lifetime
+4 -4
View File
@@ -368,7 +368,7 @@ class DFlashVerifyInput(SpecInput):
and not sampling_info.is_all_greedy and not sampling_info.is_all_greedy
and is_dflash_sampling_verify_available() and is_dflash_sampling_verify_available()
): ):
accept_len, bonus = compute_dflash_sampling_correct_drafts_and_bonus( correct_len, bonus = compute_dflash_sampling_correct_drafts_and_bonus(
candidates=candidates, candidates=candidates,
next_token_logits=logits_output.next_token_logits, next_token_logits=logits_output.next_token_logits,
sampling_info=sampling_info, sampling_info=sampling_info,
@@ -377,14 +377,14 @@ class DFlashVerifyInput(SpecInput):
target_predict = torch.argmax(logits_output.next_token_logits, dim=-1).view( target_predict = torch.argmax(logits_output.next_token_logits, dim=-1).view(
bs, self.draft_token_num bs, self.draft_token_num
) )
accept_len, bonus = compute_dflash_correct_drafts_and_bonus( correct_len, bonus = compute_dflash_correct_drafts_and_bonus(
candidates=candidates, candidates=candidates,
target_predict=target_predict, target_predict=target_predict,
) )
# Single D2H transfer: candidates[1:] + accept_len + bonus # Single D2H transfer: candidates[1:] + correct_len + bonus
packed = torch.cat( packed = torch.cat(
[candidates[:, 1:], accept_len.unsqueeze(1), bonus.unsqueeze(1)], dim=1 [candidates[:, 1:], correct_len.unsqueeze(1), bonus.unsqueeze(1)], dim=1
).cpu() ).cpu()
max_acc = self.draft_token_num - 1 max_acc = self.draft_token_num - 1
@@ -432,8 +432,8 @@ def compute_dflash_correct_drafts_and_bonus(
Shape: [bs, block_size]. target_predict[:, t] corresponds to argmax at position t. Shape: [bs, block_size]. target_predict[:, t] corresponds to argmax at position t.
Returns: Returns:
accept_len: int32 tensor [bs], number of accepted *draft* tokens (excluding current token and bonus token). correct_len: int32 tensor [bs], number of accepted *draft* tokens (excluding current token and bonus token).
bonus: int64 tensor [bs], the target-predicted token at index accept_len (the "bonus" token to append). bonus: int64 tensor [bs], the target-predicted token at index correct_len (the "bonus" token to append).
Notes: Notes:
Matches the reference implementation rule: Matches the reference implementation rule:
@@ -454,9 +454,9 @@ def compute_dflash_correct_drafts_and_bonus(
raise ValueError(f"block_size must be positive, got {block_size}.") raise ValueError(f"block_size must be positive, got {block_size}.")
matches = candidates[:, 1:] == target_predict[:, :-1] matches = candidates[:, 1:] == target_predict[:, :-1]
accept_len = matches.to(torch.int32).cumprod(dim=1).sum(dim=1) correct_len = matches.to(torch.int32).cumprod(dim=1).sum(dim=1)
bonus = target_predict[torch.arange(bs, device=target_predict.device), accept_len] bonus = target_predict[torch.arange(bs, device=target_predict.device), correct_len]
return accept_len, bonus.to(torch.int64) return correct_len, bonus.to(torch.int64)
def compute_dflash_sampling_correct_drafts_and_bonus( def compute_dflash_sampling_correct_drafts_and_bonus(
@@ -631,8 +631,8 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
deterministic=True, deterministic=True,
) )
accept_len = accept_token_num correct_len = accept_token_num
row_ids = torch.arange(bs, dtype=torch.long, device=device) row_ids = torch.arange(bs, dtype=torch.long, device=device)
accept_pos = accept_index[row_ids, accept_len.to(torch.long)].to(torch.long) accept_pos = accept_index[row_ids, correct_len.to(torch.long)].to(torch.long)
bonus = predicts[accept_pos].to(torch.int64) bonus = predicts[accept_pos].to(torch.int64)
return accept_len, bonus return correct_len, bonus
@@ -1080,7 +1080,7 @@ class DFlashWorker:
if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"): if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"):
return return
accept_steps = commit_lens.to(torch.int64) - 1 last_correct_step_indices = commit_lens.to(torch.int64) - 1
mamba_steps_to_track = None mamba_steps_to_track = None
if batch.mamba_track_indices is not None: if batch.mamba_track_indices is not None:
@@ -1103,7 +1103,7 @@ class DFlashWorker:
) )
attn_backend.update_mamba_state_after_mtp_verify( attn_backend.update_mamba_state_after_mtp_verify(
accept_steps=accept_steps, last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices, mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track, mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model, model=self.target_worker.model_runner.model,
+8 -10
View File
@@ -1003,15 +1003,12 @@ class EAGLEWorker(TpModelWorker):
if batch.forward_mode.is_idle(): if batch.forward_mode.is_idle():
return return
num_accept_tokens = ( num_correct_drafts = torch.tensor(
torch.tensor(
res.num_correct_drafts_per_req_cpu, res.num_correct_drafts_per_req_cpu,
device=logits_output.hidden_states.device, device=logits_output.hidden_states.device,
dtype=torch.int64, dtype=torch.int64,
) )
+ 1 cumulative_num_accept_tokens = torch.cumsum(num_correct_drafts + 1, dim=0)
)
cumulative_num_accept_tokens = torch.cumsum(num_accept_tokens, dim=0)
# prepend 0 to the cumulative_num_accept_tokens # prepend 0 to the cumulative_num_accept_tokens
accepted_indices_start = torch.cat( accepted_indices_start = torch.cat(
[ [
@@ -1037,14 +1034,15 @@ class EAGLEWorker(TpModelWorker):
# accepted_indices=[0,2,3,4,5,7,9,10,11], num_accept_tokens=[4, 3, 2], cumulative_num_accept_tokens=[4, 7, 9] # accepted_indices=[0,2,3,4,5,7,9,10,11], num_accept_tokens=[4, 3, 2], cumulative_num_accept_tokens=[4, 7, 9]
# first_token_indices_per_req=prepend(0, accepted_indices[cumulative_num_accept_tokens[:-1]]) = [0, 5, 10] # first_token_indices_per_req=prepend(0, accepted_indices[cumulative_num_accept_tokens[:-1]]) = [0, 5, 10]
# last_token_indices_per_req=accepted_indices[cumulative_num_accept_tokens - 1] = [4, 9, 11] (last token ID of each req) # last_token_indices_per_req=accepted_indices[cumulative_num_accept_tokens - 1] = [4, 9, 11] (last token ID of each req)
# accept_steps = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches # last_correct_step_indices = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches
# first_token_indices_per_req = res.accepted_indices[accepted_indices_start] # equivalent: last_correct_step_indices = last_token_indices_per_req - first_token_indices_per_req;
accept_steps = ( # `accepted_indices_offset` equals `first_token_indices_per_req` because the first accepted slot of each req is its "current token" at logical position i * draft_token_num.
last_correct_step_indices = (
res.accepted_indices[cumulative_num_accept_tokens - 1] res.accepted_indices[cumulative_num_accept_tokens - 1]
- accepted_indices_offset - accepted_indices_offset
) )
else: else:
accept_steps = num_accept_tokens - 1 last_correct_step_indices = num_correct_drafts
if batch.mamba_track_indices is not None: if batch.mamba_track_indices is not None:
# If after verify, the request's seq_lens has crossed a mamba track interval, # If after verify, the request's seq_lens has crossed a mamba track interval,
@@ -1068,7 +1066,7 @@ class EAGLEWorker(TpModelWorker):
mamba_steps_to_track = None mamba_steps_to_track = None
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify( self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
accept_steps=accept_steps, last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices, mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track, mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model, model=self.target_worker.model_runner.model,
@@ -1097,7 +1097,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
): ):
"""Update mamba state for hybrid GDN models after verification.""" """Update mamba state for hybrid GDN models after verification."""
# `accept_lens` already includes the bonus token (drafts + 1 per req). # `accept_lens` already includes the bonus token (drafts + 1 per req).
num_accept_tokens = accept_lens
if not batch.forward_mode.is_idle() and accept_index.numel() > 0: if not batch.forward_mode.is_idle() and accept_index.numel() > 0:
if verify_input.topk != 1: if verify_input.topk != 1:
raise ValueError("Spec v2 currently only supports topk = 1.") raise ValueError("Spec v2 currently only supports topk = 1.")
@@ -1106,16 +1105,16 @@ class EAGLEWorkerV2(BaseSpecWorker):
0, 0,
bs * self.speculative_num_draft_tokens, bs * self.speculative_num_draft_tokens,
step=self.speculative_num_draft_tokens, step=self.speculative_num_draft_tokens,
dtype=num_accept_tokens.dtype, dtype=accept_lens.dtype,
device=num_accept_tokens.device, device=accept_lens.device,
) )
accept_steps = num_accept_tokens - 1 last_correct_step_indices = accept_lens - 1
if batch.mamba_track_indices is not None: if batch.mamba_track_indices is not None:
# If after verify, the request's seq_lens has crossed a mamba track interval, # If after verify, the request's seq_lens has crossed a mamba track interval,
# we need to update the mamba state for the request at the crossing point. # we need to update the mamba state for the request at the crossing point.
seq_lens_pre_verify = batch.seq_lens seq_lens_pre_verify = batch.seq_lens
seq_lens_post_verify = batch.seq_lens + num_accept_tokens seq_lens_post_verify = batch.seq_lens + accept_lens
mamba_track_interval = self.server_args.mamba_track_interval mamba_track_interval = self.server_args.mamba_track_interval
to_track_mask = ( to_track_mask = (
seq_lens_pre_verify // mamba_track_interval seq_lens_pre_verify // mamba_track_interval
@@ -1130,7 +1129,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
req_idx = torch.arange( req_idx = torch.arange(
bs, bs,
dtype=torch.int64, dtype=torch.int64,
device=num_accept_tokens.device, device=accept_lens.device,
) )
candidate_track_steps = ( candidate_track_steps = (
accept_index[req_idx, to_track_ith] - accepted_indices_offset accept_index[req_idx, to_track_ith] - accepted_indices_offset
@@ -1144,7 +1143,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
mamba_steps_to_track = None mamba_steps_to_track = None
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify( self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
accept_steps=accept_steps, last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices, mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track, mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model, model=self.target_worker.model_runner.model,
@@ -1157,12 +1156,13 @@ class EAGLEWorkerV2(BaseSpecWorker):
num_correct_drafts: torch.Tensor, num_correct_drafts: torch.Tensor,
): ):
""" """
Move accepted tokens to the target KV cache. Move accepted tokens (drafts + bonus) to the target KV cache.
Args: Args:
batch: The batch to run. batch: The batch to run.
accept_index: The index of the accepted tokens. accept_index: The index of the accepted tokens (incl. bonus).
num_correct_drafts: The length of the accepted tokens. num_correct_drafts: Per-req count of correct drafts (excludes bonus);
seq_lens is advanced by ``num_correct_drafts + 1`` to cover the bonus slot.
""" """
bs = len(batch.seq_lens) bs = len(batch.seq_lens)
size = bs * self.speculative_num_draft_tokens size = bs * self.speculative_num_draft_tokens
@@ -1179,7 +1179,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
batch.req_pool_indices, batch.req_pool_indices,
self.req_to_token_pool.req_to_token, self.req_to_token_pool.req_to_token,
batch.seq_lens, batch.seq_lens,
batch.seq_lens + num_correct_drafts, batch.seq_lens + num_correct_drafts + 1,
tgt_cache_loc, tgt_cache_loc,
self.req_to_token_pool.req_to_token.shape[1], self.req_to_token_pool.req_to_token.shape[1],
next_power_of_2(bs), next_power_of_2(bs),
@@ -576,7 +576,7 @@ class MultiLayerEagleWorker(TpModelWorker):
# accepted_indices=[0,2,3,4,5,7,9,10,11], num_accept_tokens=[4, 3, 2], cumulative_num_accept_tokens=[4, 7, 9] # accepted_indices=[0,2,3,4,5,7,9,10,11], num_accept_tokens=[4, 3, 2], cumulative_num_accept_tokens=[4, 7, 9]
# first_token_indices_per_req=prepend(0, accepted_indices[cumulative_num_accept_tokens[:-1]]) = [0, 5, 10] # first_token_indices_per_req=prepend(0, accepted_indices[cumulative_num_accept_tokens[:-1]]) = [0, 5, 10]
# last_token_indices_per_req=accepted_indices[cumulative_num_accept_tokens - 1] = [4, 9, 11] (last token ID of each req) # last_token_indices_per_req=accepted_indices[cumulative_num_accept_tokens - 1] = [4, 9, 11] (last token ID of each req)
# max_relative_indices_per_req = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches # last_correct_step_indices = [4,4,1]; those are the per-req spec-decoding step offsets that contain the correct mamba caches
cumulative_num_accept_tokens = torch.cumsum(num_accept_tokens, dim=0) cumulative_num_accept_tokens = torch.cumsum(num_accept_tokens, dim=0)
req_start_positions = torch.cat( req_start_positions = torch.cat(
[ [
@@ -592,13 +592,13 @@ class MultiLayerEagleWorker(TpModelWorker):
last_token_indices_per_req = res.accepted_indices[ last_token_indices_per_req = res.accepted_indices[
cumulative_num_accept_tokens - 1 cumulative_num_accept_tokens - 1
] ]
max_relative_indices_per_req = ( last_correct_step_indices = (
last_token_indices_per_req - first_token_indices_per_req last_token_indices_per_req - first_token_indices_per_req
) )
else: else:
max_relative_indices_per_req = num_accept_tokens - 1 last_correct_step_indices = num_accept_tokens - 1
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify( self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
max_relative_indices_per_req, self.target_worker.model_runner.model last_correct_step_indices, self.target_worker.model_runner.model
) )
if batch.return_logprob: if batch.return_logprob: