Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions python/sglang/srt/arg_groups/pd_disaggregation_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
14 changes: 12 additions & 2 deletions test/registered/unit/server_args/test_server_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -516,12 +516,22 @@ def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self):
server_args = ServerArgs(
model_path="dummy",
disaggregation_mode="decode",
disaggregation_transfer_backend="fake",
disaggregation_transfer_backend="mori",
dcp_size=4,
)
with self.assertRaisesRegex(ValueError, "mooncake or nixl"):
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,
)
self.assertTrue(server_args.disable_radix_cache)

def test_pd_decode_dcp_rejects_radix_cache(self):
server_args = ServerArgs(
model_path="dummy",
Expand Down
Loading