[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:
Clint
2026-06-19 11:59:45 -07:00
committed by GitHub
co-authored by clintg6 HAI
parent ca88b7f1d2
commit c436a8161a
12 changed files with 971 additions and 78 deletions
@@ -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()