Fix stuck when enabling MTP on DSA models (#24635)

This commit is contained in:
Baizhou Zhang
2026-05-07 17:06:28 -07:00
committed by GitHub
parent 95fb722dd2
commit c4bb3ce273
4 changed files with 44 additions and 18 deletions
@@ -68,14 +68,14 @@ else:
def _to_2d_context_lens(seqlens_32: torch.Tensor, batch_size: int) -> torch.Tensor:
# Always normalize to (N_total, 1) layout, to avoid deadlock at deep_gemm.fp8_paged_mqa_logits
if seqlens_32.dim() == 2:
return seqlens_32
n = seqlens_32.numel()
assert (
n % batch_size == 0
), f"seqlens_32 size {n} is not a multiple of batch_size {batch_size}"
next_n = n // batch_size
return seqlens_32.view(batch_size, next_n)
if seqlens_32.size(1) == 1:
return seqlens_32
# Fall through and re-flatten if the caller already gave us a (bs, next_n)
# view — we want (N_total, 1) regardless.
seqlens_32 = seqlens_32.reshape(-1)
return seqlens_32.contiguous().view(-1, 1)
# Reuse this workspace buffer across all NSA backend instances
@@ -644,7 +644,7 @@ class NativeSparseAttnBackend(
if is_cuda() and (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
):
try:
import deep_gemm
@@ -654,7 +654,7 @@ class NativeSparseAttnBackend(
seqlens_expanded
if (
forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend(include_v2=True)
)
else cache_seqlens_int32
)
@@ -933,7 +933,7 @@ class NativeSparseAttnBackend(
if is_cuda() and (
forward_mode.is_decode_or_idle()
or forward_mode.is_target_verify()
or forward_mode.is_draft_extend()
or forward_mode.is_draft_extend(include_v2=True)
):
try:
import deep_gemm
@@ -942,7 +942,7 @@ class NativeSparseAttnBackend(
seqlens_expanded
if (
forward_mode.is_target_verify()
or forward_mode.is_draft_extend()
or forward_mode.is_draft_extend(include_v2=True)
)
else cache_seqlens_int32
)
@@ -1084,7 +1084,7 @@ class NativeSparseAttnBackend(
if is_cuda() and (
forward_mode.is_decode_or_idle()
or forward_mode.is_target_verify()
or forward_mode.is_draft_extend()
or forward_mode.is_draft_extend(include_v2=True)
):
try:
import deep_gemm
@@ -1093,7 +1093,7 @@ class NativeSparseAttnBackend(
seqlens_expanded
if (
forward_mode.is_target_verify()
or forward_mode.is_draft_extend()
or forward_mode.is_draft_extend(include_v2=True)
)
else metadata.cache_seqlens_int32
)
@@ -1102,11 +1102,13 @@ class NativeSparseAttnBackend(
seqlens_32_2d, 64, deep_gemm.get_num_sms()
)
if metadata.paged_mqa_schedule_metadata is None:
metadata.paged_mqa_schedule_metadata = new_schedule
object.__setattr__(
metadata, "paged_mqa_schedule_metadata", new_schedule
)
else:
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
except (ImportError, ModuleNotFoundError):
metadata.paged_mqa_schedule_metadata = None
object.__setattr__(metadata, "paged_mqa_schedule_metadata", None)
seqlens_expanded_size = seqlens_expanded.shape[0]
assert (
metadata.nsa_cache_seqlens_int32 is not None
@@ -1293,6 +1295,33 @@ class NativeSparseAttnBackend(
flashmla_metadata = metadata.flashmla_metadata.slice(slice(0, size + 1))
flashmla_metadata.copy_(precomputed.flashmla_metadata)
# Refresh DeepGEMM paged MQA schedule metadata for the actual seqlens of
# this replay (the captured graph holds stale data otherwise, which can
# deadlock the kernel when the runtime work decomposition diverges from
# the captured one).
if is_cuda():
try:
import deep_gemm
if forward_mode.is_decode_or_idle():
seqlens_32 = metadata.cache_seqlens_int32
else:
seqlens_32 = metadata.nsa_seqlens_expanded[
: precomputed.seqlens_expanded_size
]
seqlens_32_2d = _to_2d_context_lens(seqlens_32, bs)
new_schedule = deep_gemm.get_paged_mqa_logits_metadata(
seqlens_32_2d, 64, deep_gemm.get_num_sms()
)
if metadata.paged_mqa_schedule_metadata is None:
object.__setattr__(
metadata, "paged_mqa_schedule_metadata", new_schedule
)
else:
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
except (ImportError, ModuleNotFoundError):
pass
self.forward_metadata = metadata
def forward_extend(