[DCP] Resolve --dcp-comm-backend to fi_a2a/a2a by default for every model (#39165)

Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
Khoa Pham
2026-09-11 23:48:42 -07:00
committed by GitHub
co-authored by Claude Fable 5.1
parent 981b947568
commit 6ba96d329f
10 changed files with 104 additions and 52 deletions
@@ -3407,5 +3407,45 @@ class TestTpLmHeadAllToAllNcclGraphRegister(unittest.TestCase):
self.assertNotIn("NCCL_GRAPH_REGISTER", os.environ)
class TestDcpCommBackendDefault(CustomTestCase):
def _resolved(self, **fields):
args = ServerArgs(model_path="dummy", tp_size=8, **fields)
parallel_hook.handle_decode_context_parallelism(args)
return resolution_result(args, "dcp_comm_backend")
def test_no_dcp_is_ag_rs(self):
self.assertEqual(self._resolved(dcp_size=1), "ag_rs")
@override_platform(is_cuda=True, is_hip=False)
def test_fi_a2a_where_supported(self):
with patch(
"sglang.srt.arg_groups.overrides.is_fi_a2a_supported", return_value=True
):
self.assertEqual(self._resolved(dcp_size=4), "fi_a2a")
@override_platform(is_cuda=True, is_hip=False)
def test_a2a_on_cuda_without_mnnvl(self):
with patch(
"sglang.srt.arg_groups.overrides.is_fi_a2a_supported", return_value=False
):
self.assertEqual(self._resolved(dcp_size=4), "a2a")
@override_platform(is_cuda=False, is_hip=False)
def test_ag_rs_off_cuda(self):
with patch(
"sglang.srt.arg_groups.overrides.is_fi_a2a_supported", return_value=False
):
self.assertEqual(self._resolved(dcp_size=4), "ag_rs")
@override_platform(is_cuda=True, is_hip=False)
def test_explicit_value_wins(self):
with patch(
"sglang.srt.arg_groups.overrides.is_fi_a2a_supported", return_value=True
):
self.assertEqual(
self._resolved(dcp_size=4, dcp_comm_backend="ag_rs"), "ag_rs"
)
if __name__ == "__main__":
unittest.main()