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()