Fix stuck when enabling MTP on DSA models (#24635)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user