[AMD] Enable HiSparse on ROCm (#26639)
Co-authored-by: clintg6 <7388379+clintg6@users.noreply.github.com> Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -147,6 +147,115 @@ class TestLoadBalanceMethod(unittest.TestCase):
|
||||
self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake")
|
||||
|
||||
|
||||
class TestHiSparseDsaBackendPolicy(unittest.TestCase):
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=False)
|
||||
def test_hisparse_defaults_to_flashmla_sparse_on_cuda_bfloat16(self, _mock_is_hip):
|
||||
server_args = ServerArgs(model_path="dummy", enable_hisparse=True)
|
||||
|
||||
server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9)
|
||||
|
||||
self.assertEqual(server_args.dsa_prefill_backend, "flashmla_sparse")
|
||||
self.assertEqual(server_args.dsa_decode_backend, "flashmla_sparse")
|
||||
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=False)
|
||||
def test_hisparse_defaults_to_flashmla_kv_on_cuda_fp8(self, _mock_is_hip):
|
||||
server_args = ServerArgs(model_path="dummy", enable_hisparse=True)
|
||||
|
||||
server_args._set_default_dsa_backends(kv_cache_dtype="fp8_e4m3", major=9)
|
||||
|
||||
self.assertEqual(server_args.dsa_prefill_backend, "flashmla_kv")
|
||||
self.assertEqual(server_args.dsa_decode_backend, "flashmla_kv")
|
||||
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=True)
|
||||
def test_hisparse_defaults_to_tilelang_on_rocm(self, _mock_is_hip):
|
||||
server_args = ServerArgs(model_path="dummy", enable_hisparse=True)
|
||||
|
||||
server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9)
|
||||
|
||||
self.assertEqual(server_args.dsa_prefill_backend, "tilelang")
|
||||
self.assertEqual(server_args.dsa_decode_backend, "tilelang")
|
||||
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=True)
|
||||
def test_hisparse_preserves_rocm_user_backend_and_defaults_missing_side(
|
||||
self, _mock_is_hip
|
||||
):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
dsa_prefill_backend="tilelang",
|
||||
)
|
||||
|
||||
server_args._set_default_dsa_backends(kv_cache_dtype="bfloat16", major=9)
|
||||
|
||||
self.assertEqual(server_args.dsa_prefill_backend, "tilelang")
|
||||
self.assertEqual(server_args.dsa_decode_backend, "tilelang")
|
||||
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=True)
|
||||
def test_hisparse_accepts_aiter_backend_on_rocm(self, _mock_is_hip):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
kv_cache_dtype="bfloat16",
|
||||
dsa_prefill_backend="aiter",
|
||||
dsa_decode_backend="aiter",
|
||||
)
|
||||
|
||||
server_args._validate_hisparse_dsa_backend("dsa_prefill_backend", "prefill")
|
||||
server_args._validate_hisparse_dsa_backend("dsa_decode_backend", "decode")
|
||||
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=True)
|
||||
def test_hisparse_rejects_cuda_backend_on_rocm(self, _mock_is_hip):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
kv_cache_dtype="bfloat16",
|
||||
dsa_prefill_backend="flashmla_sparse",
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "tilelang"):
|
||||
server_args._validate_hisparse_dsa_backend("dsa_prefill_backend", "prefill")
|
||||
|
||||
@patch("sglang.srt.server_args.is_hip", return_value=False)
|
||||
def test_hisparse_rejects_rocm_backend_on_cuda(self, _mock_is_hip):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
kv_cache_dtype="bfloat16",
|
||||
dsa_decode_backend="tilelang",
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "flashmla_sparse"):
|
||||
server_args._validate_hisparse_dsa_backend("dsa_decode_backend", "decode")
|
||||
|
||||
def test_hisparse_accepts_bfloat16_kv_cache_dtype(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
kv_cache_dtype="bfloat16",
|
||||
)
|
||||
|
||||
server_args._validate_hisparse_kv_cache_dtype()
|
||||
|
||||
def test_hisparse_accepts_fp8_e4m3_kv_cache_dtype(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
kv_cache_dtype="fp8_e4m3",
|
||||
)
|
||||
|
||||
server_args._validate_hisparse_kv_cache_dtype()
|
||||
|
||||
def test_hisparse_rejects_unsupported_kv_cache_dtype(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
enable_hisparse=True,
|
||||
kv_cache_dtype="float16",
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, r"fp8_e4m3"):
|
||||
server_args._validate_hisparse_kv_cache_dtype()
|
||||
|
||||
|
||||
class TestContextParallelServerArgs(CustomTestCase):
|
||||
def setUp(self):
|
||||
self.parser = server_args_module.argparse.ArgumentParser()
|
||||
|
||||
Reference in New Issue
Block a user