From 6ba96d329fb838d7028b5f3db4f5dbc3fb8e366d Mon Sep 17 00:00:00 2001 From: Khoa Pham Date: Fri, 11 Sep 2026 23:48:42 -0700 Subject: [PATCH] [DCP] Resolve --dcp-comm-backend to fi_a2a/a2a by default for every model (#39165) Co-authored-by: Claude Fable 5.1 --- docs/docs/advanced_features/dcp.mdx | 10 ++--- .../advanced_features/server_arguments.mdx | 4 +- .../sglang/srt/arg_groups/fields/parallel.py | 7 ++-- .../srt/arg_groups/model_overrides/kimi_k3.py | 26 ++---------- python/sglang/srt/arg_groups/overrides.py | 27 +++++++++++++ python/sglang/srt/arg_groups/parallel_hook.py | 11 ++--- python/sglang/srt/arg_groups/pipeline.py | 4 +- python/sglang/srt/layers/dcp/comm.py | 15 +------ python/sglang/srt/utils/common.py | 12 ++++++ .../unit/server_args/test_server_args.py | 40 +++++++++++++++++++ 10 files changed, 104 insertions(+), 52 deletions(-) diff --git a/docs/docs/advanced_features/dcp.mdx b/docs/docs/advanced_features/dcp.mdx index e7bd9ed62..262b69c4c 100644 --- a/docs/docs/advanced_features/dcp.mdx +++ b/docs/docs/advanced_features/dcp.mdx @@ -15,7 +15,7 @@ On Kimi K3, DCP applies only to MLA layers. Request-indexed KDA state is unchang | Setting | Behavior | | --- | --- | | `--dcp-size N` | Enables DCP and widens virtual capacity and page size by `N`. Alias: `--decode-context-parallel-size`. | -| `--dcp-comm-backend` | Selects the MLA partial-output merge: `ag_rs`, `a2a`, or `fi_a2a`. | +| `--dcp-comm-backend` | Selects the MLA partial-output merge: `ag_rs`, `a2a`, or `fi_a2a`. Unset resolves to `fi_a2a` where the DCP group shares an MNNVL domain on Blackwell, otherwise `a2a` on CUDA/ROCm. | | `--dcp-replicate-q-proj` | Removes the query all-gather for supported MLA weights. Use `--no-dcp-replicate-q-proj` to disable a model-specific default. | | `--enable-dp-attention` | Composes only when DCP groups nest inside attention TP. | @@ -44,7 +44,7 @@ sglang serve \ --port 30000 ``` -Kimi K3 on 32 GPUs with DPA and DCP 8. The model override enables replicated Q by default and picks `fi_a2a` on MNNVL systems, otherwise `a2a`: +Kimi K3 on 32 GPUs with DPA and DCP 8. The model override enables replicated Q by default; the communication backend resolves generically to `fi_a2a` on MNNVL systems, otherwise `a2a`: ```bash Command sglang serve \ @@ -130,12 +130,12 @@ Normally each layer all-gathers its absorbed and rotary query components. `--dcp | Backend | DCP collectives per MLA layer | Notes | | --- | --- | --- | -| `ag_rs` | Query AG + LSE AG + FP32 output RS | Generic fallback | +| `ag_rs` | Query AG + LSE AG + FP32 output RS | Fallback for platforms without the a2a path; opt in with `--dcp-comm-backend ag_rs` | | `a2a`, gathered Q | Query AG + packed NCCL A2A | Two collectives | | `a2a`, replicated Q | One packed NCCL A2A | Canonical low-latency path | -| `fi_a2a`, replicated Q | One FlashInfer MNNVL A2A | Requires CUDA, SM90+, and MNNVL fabric | +| `fi_a2a`, replicated Q | One FlashInfer MNNVL A2A | Requires Blackwell and a DCP group inside one MNNVL domain (GB200/GB300 fabric, or a single node) | -Kimi K3 enables replicated Q by default and picks `fi_a2a` on MNNVL systems, otherwise `a2a`. Its decode backend is `cutedsl_mla`. The generic defaults remain `ag_rs` and model-resolved query replication. +When `--dcp-comm-backend` is unset, every DCP launch resolves it the same way: `fi_a2a` where the DCP group shares one MNNVL domain on Blackwell, otherwise `a2a` on CUDA/ROCm, and `ag_rs` only on platforms without the a2a path. Query replication stays model-resolved: Kimi K3 enables it by default on the a2a family, and its decode backend is `cutedsl_mla`. DCP adds a context-independent decode collective while cutting context-dependent KV storage and reads by about `c`. Extend is outside that decode cost model: it may gather cached prefix shards, restore token order, and append new tokens, and that work grows with prefix length. diff --git a/docs/docs/advanced_features/server_arguments.mdx b/docs/docs/advanced_features/server_arguments.mdx index 224b8fa52..b573d1a2f 100644 --- a/docs/docs/advanced_features/server_arguments.mdx +++ b/docs/docs/advanced_features/server_arguments.mdx @@ -668,8 +668,8 @@ Please consult the documentation below and [server_args.py](https://github.com/s `--dcp-comm-backend` - Communication backend for the DCP attention reduction: AllGather + ReduceScatter, fused NCCL All-to-All, or FlashInfer MNNVL All-to-All. - `ag_rs` + Communication backend for the DCP attention reduction: AllGather + ReduceScatter, fused NCCL All-to-All, or FlashInfer MNNVL All-to-All. Unset resolves to fi_a2a where the DCP group shares an MNNVL domain on Blackwell, else a2a on CUDA/ROCm, else ag_rs. + None ag_rs, a2a, fi_a2a diff --git a/python/sglang/srt/arg_groups/fields/parallel.py b/python/sglang/srt/arg_groups/fields/parallel.py index 4e1c6eb52..dec06cd44 100644 --- a/python/sglang/srt/arg_groups/fields/parallel.py +++ b/python/sglang/srt/arg_groups/fields/parallel.py @@ -121,17 +121,18 @@ class Parallel(msgspec.Struct): ), ] = 1 dcp_comm_backend: A[ - str, + Optional[str], Arg( help="Communication backend for the decode context-parallel (DCP) " "attention reduction: 'ag_rs' (AllGather + ReduceScatter), 'a2a' " "(fused NCCL All-to-All exchange of output+LSE + local Triton LSE " "combine), or 'fi_a2a' (FlashInfer MNNVL All-to-All kernel; requires " - "SM90+ and MNNVL fabric memory, e.g. GB200 NVL72).", + "Blackwell and a DCP group within one MNNVL domain). Unset resolves " + "to 'fi_a2a' where supported, else 'a2a' on CUDA/ROCm, else 'ag_rs'.", choices=["ag_rs", "a2a", "fi_a2a"], resolvable=True, ), - ] = "ag_rs" + ] = None dcp_replicate_q_proj: A[ Optional[bool], Arg( diff --git a/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py b/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py index fe2b361d6..033450943 100644 --- a/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py +++ b/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py @@ -16,7 +16,6 @@ from sglang.srt.arg_groups.model_override_base import ( resolving_view, ) from sglang.srt.runtime_context import get_platform -from sglang.srt.utils.common import get_device_name logger = logging.getLogger(__name__) @@ -130,29 +129,12 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla', 'tokenspeed_mla' or 'aiter', got {decode_backend!r}." ) - if cfg.dcp_replicate_q_proj is None: + if cfg.dcp_replicate_q_proj is None and cfg.dcp_comm_backend in ( + "a2a", + "fi_a2a", + ): logger.info("Kimi-K3 DCP enables replicated Q projection by default.") overrides["dcp_replicate_q_proj"] = True - - from sglang.srt.layers.dcp.comm import is_fi_a2a_supported - - device_name = get_device_name() - dcp_comm_backend = ( - "fi_a2a" - if is_fi_a2a_supported( - dcp_size=cfg.dcp_size, - tp_size=cfg.tp_size, - pp_size=cfg.pp_size, - nnodes=cfg.nnodes, - ) - else "a2a" - ) - logger.info( - "Kimi-K3 DCP selects communication backend on " - f"{device_name!r}: {cfg.dcp_comm_backend!r} -> " - f"{dcp_comm_backend!r}." - ) - overrides["dcp_comm_backend"] = dcp_comm_backend return overrides if not (get_platform().is_sm100 and get_platform().device_sm in (100, 103)): diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index cd6b64db7..c6730c8bf 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -78,6 +78,7 @@ from sglang.srt.runtime_context import ( ) from sglang.srt.utils.common import ( get_quantization_config, + is_fi_a2a_supported, is_gfx95_supported, xpu_has_xmx_support, ) @@ -1424,6 +1425,32 @@ def _data_parallelism_defaults(view: Any) -> dict: return {} +@register_post_process +def _dcp_comm_backend_default(view: Any) -> dict: + if view.dcp_comm_backend is not None: + return {} + if view.dcp_size <= 1: + return {"dcp_comm_backend": "ag_rs"} + platform = get_platform() + if is_fi_a2a_supported( + dcp_size=view.dcp_size, + tp_size=view.tp_size, + pp_size=view.pp_size, + nnodes=view.nnodes, + ): + backend = "fi_a2a" + elif platform.is_cuda or platform.is_hip: + backend = "a2a" + else: + backend = "ag_rs" + logger.info( + "DCP (dcp_size=%d) selects communication backend %r.", + view.dcp_size, + backend, + ) + return {"dcp_comm_backend": backend} + + @register_post_process def _tp_lm_head_all_to_all_default(view: Any) -> dict: """Enable the TP LM-head all-to-all path only for pure-DP decode nodes. diff --git a/python/sglang/srt/arg_groups/parallel_hook.py b/python/sglang/srt/arg_groups/parallel_hook.py index 286a0790d..8e412fa20 100644 --- a/python/sglang/srt/arg_groups/parallel_hook.py +++ b/python/sglang/srt/arg_groups/parallel_hook.py @@ -9,6 +9,7 @@ from typing import Any from sglang.srt.arg_groups.overrides import ( _data_parallelism_defaults, + _dcp_comm_backend_default, _dp_lm_head_validation, _tp_lm_head_all_to_all_default, declare_resolution, @@ -114,7 +115,8 @@ def handle_context_parallelism(server_args: Any): ) -def handle_dcp_validation(server_args: Any): +def handle_decode_context_parallelism(server_args: Any): + run_post_process_pass(server_args, _dcp_comm_backend_default) cfg = resolving_view(server_args) if cfg.dcp_size < 1: raise ValueError( @@ -132,10 +134,9 @@ def handle_dcp_validation(server_args: Any): if cfg.dcp_comm_backend == "fi_a2a" and not get_platform().is_cuda: raise ValueError( "--dcp-comm-backend fi_a2a delegates the exchange to FlashInfer's " - "MNNVL All-to-All kernel, which requires an NVIDIA CUDA platform " - "with SM90+ and MNNVL fabric memory (e.g. GB200 NVL72). The " - "authoritative fabric probe runs at model-runner init; use 'a2a' " - "or 'ag_rs' on clusters without MNNVL." + "MNNVL All-to-All kernel, which requires Blackwell and a DCP group " + "within one MNNVL domain. Use 'a2a' or 'ag_rs' elsewhere, or leave " + "the flag unset to resolve it." ) if cfg.dcp_replicate_q_proj: if cfg.dcp_size <= 1: diff --git a/python/sglang/srt/arg_groups/pipeline.py b/python/sglang/srt/arg_groups/pipeline.py index b51a627c2..298675869 100644 --- a/python/sglang/srt/arg_groups/pipeline.py +++ b/python/sglang/srt/arg_groups/pipeline.py @@ -163,7 +163,7 @@ def run_resolution_pipeline(server_args: Any) -> None: from sglang.srt.arg_groups.parallel_hook import ( handle_context_parallelism, handle_data_parallelism, - handle_dcp_validation, + handle_decode_context_parallelism, handle_dwdp, handle_elastic_ep, handle_eplb_and_dispatch, @@ -171,7 +171,7 @@ def run_resolution_pipeline(server_args: Any) -> None: ) validate_prefill_only_disable_kv_cache_args(server_args) - handle_dcp_validation(server_args) + handle_decode_context_parallelism(server_args) # Model-arch prefill CUDA-graph default must land before cuda-graph # resolution (the declarative registry materializes too late to affect diff --git a/python/sglang/srt/layers/dcp/comm.py b/python/sglang/srt/layers/dcp/comm.py index 41e5ce284..ae2d1422d 100644 --- a/python/sglang/srt/layers/dcp/comm.py +++ b/python/sglang/srt/layers/dcp/comm.py @@ -36,9 +36,9 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.distributed.parallel_state import GroupCoordinator -from sglang.srt.runtime_context import get_parallel, get_platform +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import is_hip -from sglang.srt.utils.common import is_mnnvl_fabric_device +from sglang.srt.utils.common import is_fi_a2a_supported _is_hip = is_hip() @@ -390,17 +390,6 @@ def all_gather_kv_cache_for_dcp( _FI_A2A_STATE: Optional[dict] = None -def is_fi_a2a_supported( - *, dcp_size: int, tp_size: int, pp_size: int, nnodes: int -) -> bool: - if not get_platform().is_sm100: - return False - if is_mnnvl_fabric_device(): - return True - tp_size_per_node = tp_size // max(nnodes // pp_size, 1) - return tp_size_per_node % dcp_size == 0 - - def init_fi_a2a_workspace(cp_group: "GroupCoordinator") -> None: # Call once per process BEFORE CUDA-graph capture: the FlashInfer init syncs # the stream and barriers cross-rank, neither of which is capturable. diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index bbffe4795..62ad5159f 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -106,6 +106,7 @@ from sglang.srt.runtime_context import ( get_flags, get_model, get_parallel, + get_platform, get_spec, ) from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper @@ -875,6 +876,17 @@ def is_mnnvl_fabric_device() -> bool: return any(tag in name for tag in ("GB200", "GB300")) +def is_fi_a2a_supported( + *, dcp_size: int, tp_size: int, pp_size: int, nnodes: int +) -> bool: + if not get_platform().is_sm100: + return False + if is_mnnvl_fabric_device(): + return True + tp_size_per_node = tp_size // max(nnodes // pp_size, 1) + return tp_size_per_node % dcp_size == 0 + + @lru_cache(maxsize=1) def is_habana_available() -> bool: return find_spec("habana_frameworks") is not None diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 7cef34025..cbaa1fc9a 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -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()