Fix CuTe DSL DSA paged MQA export (#30627)
This commit is contained in:
@@ -1,4 +1,12 @@
|
|||||||
from sglang.srt.utils import is_hip
|
from sglang.srt.utils import is_cuda
|
||||||
|
|
||||||
|
_is_cuda = is_cuda()
|
||||||
|
|
||||||
|
if _is_cuda:
|
||||||
|
from .cutedsl_paged_mqa_logits import CuteDSLPagedMQALogitsRunner, pick_dsl_expand
|
||||||
|
else:
|
||||||
|
CuteDSLPagedMQALogitsRunner = None
|
||||||
|
pick_dsl_expand = None
|
||||||
|
|
||||||
from .paged_mqa_logits import (
|
from .paged_mqa_logits import (
|
||||||
aiter_paged_mqa_logits,
|
aiter_paged_mqa_logits,
|
||||||
@@ -7,10 +15,6 @@ from .paged_mqa_logits import (
|
|||||||
deepgemm_paged_mqa_logits_split,
|
deepgemm_paged_mqa_logits_split,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not is_hip():
|
|
||||||
# Preserve the original eager import behavior on non-ROCm platforms.
|
|
||||||
from .cutedsl_paged_mqa_logits import CuteDSLPagedMQALogitsRunner, pick_dsl_expand
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"CuteDSLPagedMQALogitsRunner",
|
"CuteDSLPagedMQALogitsRunner",
|
||||||
"pick_dsl_expand",
|
"pick_dsl_expand",
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ global _use_multi_stream
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
|
||||||
if not _is_npu:
|
if not _is_npu:
|
||||||
from sglang.jit_kernel.dsa import (
|
from sglang.jit_kernel.dsa import (
|
||||||
aiter_paged_mqa_logits,
|
aiter_paged_mqa_logits,
|
||||||
@@ -71,11 +72,11 @@ else:
|
|||||||
deepgemm_paged_mqa_logits_native = None
|
deepgemm_paged_mqa_logits_native = None
|
||||||
deepgemm_paged_mqa_logits_split = None
|
deepgemm_paged_mqa_logits_split = None
|
||||||
|
|
||||||
if not _is_hip and not _is_npu:
|
if _is_cuda:
|
||||||
# Preserve the original eager import behavior on non-ROCm platforms.
|
|
||||||
from sglang.jit_kernel.dsa import pick_dsl_expand
|
from sglang.jit_kernel.dsa import pick_dsl_expand
|
||||||
else:
|
else:
|
||||||
pick_dsl_expand = None
|
pick_dsl_expand = None
|
||||||
|
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
_is_fp8_fnuz = is_fp8_fnuz()
|
_is_fp8_fnuz = is_fp8_fnuz()
|
||||||
_is_gfx95_supported = is_gfx95_supported()
|
_is_gfx95_supported = is_gfx95_supported()
|
||||||
@@ -928,7 +929,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
and forward_batch.forward_mode.is_target_verify()
|
and forward_batch.forward_mode.is_target_verify()
|
||||||
and next_n >= 2
|
and next_n >= 2
|
||||||
):
|
):
|
||||||
assert pick_dsl_expand is not None, "Not supported on AMD/ROCm. "
|
assert pick_dsl_expand is not None, "CuTe DSL paged MQA is CUDA-only."
|
||||||
dsl_expand_factor, dsl_atom = pick_dsl_expand(
|
dsl_expand_factor, dsl_atom = pick_dsl_expand(
|
||||||
next_n,
|
next_n,
|
||||||
batch_size=B,
|
batch_size=B,
|
||||||
|
|||||||
Reference in New Issue
Block a user