diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index ddaf80744..65be40b5f 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -35,10 +35,16 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: ) if server_args.disaggregation_mode == "decode" and server_args.dcp_size > 1: - if server_args.disaggregation_transfer_backend not in ("mooncake", "nixl"): + # Fake transfer moves no KV and is only used for synthetic decode + # benchmarks, so it does not need the DCP relayout from Mooncake/NIXL. + if server_args.disaggregation_transfer_backend not in ( + "mooncake", + "nixl", + "fake", + ): raise ValueError( "PD decode DCP requires --disaggregation-transfer-backend " - "mooncake or nixl, got " + "mooncake, nixl, or fake for synthetic benchmarking, got " f"{server_args.disaggregation_transfer_backend!r}." ) if server_args.disaggregation_decode_enable_radix_cache: diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index c94adc5e7..faf2cbe9f 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -524,12 +524,22 @@ class TestLoadBalanceMethod(unittest.TestCase): def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self): server_args = ServerArgs( model_path="dummy", + disaggregation_mode="decode", + disaggregation_transfer_backend="mori", + dcp_size=4, + ) + with self.assertRaisesRegex( + ValueError, "mooncake, nixl, or fake for synthetic benchmarking" + ): + server_args._handle_pd_disaggregation() + + def test_pd_decode_dcp_allows_fake_transfer_backend(self): + server_args = self._load_balance_args( disaggregation_mode="decode", disaggregation_transfer_backend="fake", dcp_size=4, ) - with self.assertRaisesRegex(ValueError, "mooncake or nixl"): - server_args._handle_pd_disaggregation() + self.assertTrue(server_args.disable_radix_cache) def test_pd_decode_dcp_rejects_radix_cache(self): server_args = ServerArgs(