fix(disagg): allow fake transfer with decode DCP (#35409)

Signed-off-by: Alexandre Milesi <milesial@users.noreply.github.com>
This commit is contained in:
milesial
2026-08-19 13:54:37 -07:00
committed by GitHub
parent defb2a3100
commit ed12d6827d
2 changed files with 20 additions and 4 deletions
@@ -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_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( raise ValueError(
"PD decode DCP requires --disaggregation-transfer-backend " "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}." f"{server_args.disaggregation_transfer_backend!r}."
) )
if server_args.disaggregation_decode_enable_radix_cache: if server_args.disaggregation_decode_enable_radix_cache:
@@ -524,12 +524,22 @@ class TestLoadBalanceMethod(unittest.TestCase):
def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self): def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", 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_mode="decode",
disaggregation_transfer_backend="fake", disaggregation_transfer_backend="fake",
dcp_size=4, dcp_size=4,
) )
with self.assertRaisesRegex(ValueError, "mooncake or nixl"): self.assertTrue(server_args.disable_radix_cache)
server_args._handle_pd_disaggregation()
def test_pd_decode_dcp_rejects_radix_cache(self): def test_pd_decode_dcp_rejects_radix_cache(self):
server_args = ServerArgs( server_args = ServerArgs(