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:
|
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:
|
if seqlens_32.dim() == 2:
|
||||||
|
if seqlens_32.size(1) == 1:
|
||||||
return seqlens_32
|
return seqlens_32
|
||||||
n = seqlens_32.numel()
|
# Fall through and re-flatten if the caller already gave us a (bs, next_n)
|
||||||
assert (
|
# view — we want (N_total, 1) regardless.
|
||||||
n % batch_size == 0
|
seqlens_32 = seqlens_32.reshape(-1)
|
||||||
), f"seqlens_32 size {n} is not a multiple of batch_size {batch_size}"
|
return seqlens_32.contiguous().view(-1, 1)
|
||||||
next_n = n // batch_size
|
|
||||||
return seqlens_32.view(batch_size, next_n)
|
|
||||||
|
|
||||||
|
|
||||||
# Reuse this workspace buffer across all NSA backend instances
|
# Reuse this workspace buffer across all NSA backend instances
|
||||||
@@ -644,7 +644,7 @@ class NativeSparseAttnBackend(
|
|||||||
if is_cuda() and (
|
if is_cuda() and (
|
||||||
forward_batch.forward_mode.is_decode_or_idle()
|
forward_batch.forward_mode.is_decode_or_idle()
|
||||||
or forward_batch.forward_mode.is_target_verify()
|
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:
|
try:
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
@@ -654,7 +654,7 @@ class NativeSparseAttnBackend(
|
|||||||
seqlens_expanded
|
seqlens_expanded
|
||||||
if (
|
if (
|
||||||
forward_batch.forward_mode.is_target_verify()
|
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
|
else cache_seqlens_int32
|
||||||
)
|
)
|
||||||
@@ -933,7 +933,7 @@ class NativeSparseAttnBackend(
|
|||||||
if is_cuda() and (
|
if is_cuda() and (
|
||||||
forward_mode.is_decode_or_idle()
|
forward_mode.is_decode_or_idle()
|
||||||
or forward_mode.is_target_verify()
|
or forward_mode.is_target_verify()
|
||||||
or forward_mode.is_draft_extend()
|
or forward_mode.is_draft_extend(include_v2=True)
|
||||||
):
|
):
|
||||||
try:
|
try:
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
@@ -942,7 +942,7 @@ class NativeSparseAttnBackend(
|
|||||||
seqlens_expanded
|
seqlens_expanded
|
||||||
if (
|
if (
|
||||||
forward_mode.is_target_verify()
|
forward_mode.is_target_verify()
|
||||||
or forward_mode.is_draft_extend()
|
or forward_mode.is_draft_extend(include_v2=True)
|
||||||
)
|
)
|
||||||
else cache_seqlens_int32
|
else cache_seqlens_int32
|
||||||
)
|
)
|
||||||
@@ -1084,7 +1084,7 @@ class NativeSparseAttnBackend(
|
|||||||
if is_cuda() and (
|
if is_cuda() and (
|
||||||
forward_mode.is_decode_or_idle()
|
forward_mode.is_decode_or_idle()
|
||||||
or forward_mode.is_target_verify()
|
or forward_mode.is_target_verify()
|
||||||
or forward_mode.is_draft_extend()
|
or forward_mode.is_draft_extend(include_v2=True)
|
||||||
):
|
):
|
||||||
try:
|
try:
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
@@ -1093,7 +1093,7 @@ class NativeSparseAttnBackend(
|
|||||||
seqlens_expanded
|
seqlens_expanded
|
||||||
if (
|
if (
|
||||||
forward_mode.is_target_verify()
|
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
|
else metadata.cache_seqlens_int32
|
||||||
)
|
)
|
||||||
@@ -1102,11 +1102,13 @@ class NativeSparseAttnBackend(
|
|||||||
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
seqlens_32_2d, 64, deep_gemm.get_num_sms()
|
||||||
)
|
)
|
||||||
if metadata.paged_mqa_schedule_metadata is None:
|
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:
|
else:
|
||||||
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
|
metadata.paged_mqa_schedule_metadata.copy_(new_schedule)
|
||||||
except (ImportError, ModuleNotFoundError):
|
except (ImportError, ModuleNotFoundError):
|
||||||
metadata.paged_mqa_schedule_metadata = None
|
object.__setattr__(metadata, "paged_mqa_schedule_metadata", None)
|
||||||
seqlens_expanded_size = seqlens_expanded.shape[0]
|
seqlens_expanded_size = seqlens_expanded.shape[0]
|
||||||
assert (
|
assert (
|
||||||
metadata.nsa_cache_seqlens_int32 is not None
|
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 = metadata.flashmla_metadata.slice(slice(0, size + 1))
|
||||||
flashmla_metadata.copy_(precomputed.flashmla_metadata)
|
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
|
self.forward_metadata = metadata
|
||||||
|
|
||||||
def forward_extend(
|
def forward_extend(
|
||||||
|
|||||||
@@ -19,7 +19,6 @@ from sglang.test.test_utils import (
|
|||||||
register_cuda_ci(
|
register_cuda_ci(
|
||||||
est_time=1048,
|
est_time=1048,
|
||||||
suite="stage-c-test-8-gpu-h200",
|
suite="stage-c-test-8-gpu-h200",
|
||||||
disabled="Disabled due to #24268. Should be fixed soon.",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
FULL_DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2"
|
FULL_DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2"
|
||||||
|
|||||||
@@ -16,7 +16,6 @@ from sglang.test.test_utils import (
|
|||||||
register_cuda_ci(
|
register_cuda_ci(
|
||||||
est_time=616,
|
est_time=616,
|
||||||
suite="stage-c-test-deepep-8-gpu-h200",
|
suite="stage-c-test-deepep-8-gpu-h200",
|
||||||
disabled="Disabled due to #24268. Should be fixed soon.",
|
|
||||||
)
|
)
|
||||||
DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2"
|
DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2"
|
||||||
|
|
||||||
|
|||||||
@@ -18,7 +18,6 @@ from sglang.test.test_utils import (
|
|||||||
register_cuda_ci(
|
register_cuda_ci(
|
||||||
est_time=1060,
|
est_time=1060,
|
||||||
suite="stage-c-test-4-gpu-b200",
|
suite="stage-c-test-4-gpu-b200",
|
||||||
disabled="Disabled due to #24268. Should be fixed soon.",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3.2-NVFP4"
|
FULL_DEEPSEEK_V3_FP4_MODEL_PATH = "nvidia/DeepSeek-V3.2-NVFP4"
|
||||||
|
|||||||
Reference in New Issue
Block a user