diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index 912c99c80..43bd301b9 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -31,16 +31,10 @@ def handle_pd_disaggregation(server_args: "ServerArgs") -> None: "--disaggregation-decode-enable-radix-cache is incompatible " "with --enable-hisparse" ) - if server_args.disaggregation_transfer_backend not in ( - "nixl", - "mooncake", - "mori", - ): + if server_args.disaggregation_transfer_backend == "fake": raise ValueError( - "--disaggregation-decode-enable-radix-cache currently " - "requires --disaggregation-transfer-backend in " - "('nixl', 'mooncake', 'mori'), but got " - f"{server_args.disaggregation_transfer_backend!r}" + "--disaggregation-decode-enable-radix-cache is incompatible " + "with --disaggregation-transfer-backend fake" ) if server_args.speculative_algorithm is not None: raise ValueError( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 399449c5d..3235fa052 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -7387,7 +7387,7 @@ class ServerArgs: parser.add_argument( "--disaggregation-decode-enable-radix-cache", action="store_true", - help="Enable radix cache on decode server (PD mode). Caches KV prefixes to avoid redundant transfers. Requires --disaggregation-transfer-backend nixl, mooncake or mori and is incompatible with --enable-hisparse.", + help="Enable radix cache on decode server (PD mode). Caches KV prefixes to avoid redundant transfers. Incompatible with --enable-hisparse, speculative decoding, and --disaggregation-transfer-backend fake.", ) parser.add_argument( "--disaggregation-decode-enable-offload-kvcache", diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index c55d85ef4..b4db117b3 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -109,7 +109,7 @@ class TestLoadBalanceMethod(unittest.TestCase): self.assertFalse(server_args.disable_radix_cache) - def test_pd_decode_radix_cache_rejects_unknown_backend(self): + def test_pd_decode_radix_cache_rejects_fake_backend(self): with self.assertRaises(ValueError) as context: ServerArgs( model_path="dummy", @@ -118,8 +118,32 @@ class TestLoadBalanceMethod(unittest.TestCase): disaggregation_transfer_backend="fake", ) - self.assertIn("('nixl', 'mooncake', 'mori')", str(context.exception)) - self.assertIn("'fake'", str(context.exception)) + self.assertIn( + "--disaggregation-decode-enable-radix-cache is incompatible " + "with --disaggregation-transfer-backend fake", + str(context.exception), + ) + + def test_pd_decode_radix_cache_allows_ascend(self): + server_args = ServerArgs( + model_path="dummy", + disaggregation_mode="decode", + disaggregation_decode_enable_radix_cache=True, + disaggregation_transfer_backend="ascend", + ) + + self.assertFalse(server_args.disable_radix_cache) + + def test_pd_decode_radix_cache_allows_mooncake_tcp(self): + server_args = ServerArgs( + model_path="dummy", + disaggregation_mode="decode", + disaggregation_decode_enable_radix_cache=True, + disaggregation_transfer_backend="mooncake_tcp", + ) + + self.assertFalse(server_args.disable_radix_cache) + self.assertEqual(server_args.disaggregation_transfer_backend, "mooncake") class TestContextParallelServerArgs(CustomTestCase):