Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
79 commits
Select commit Hold shift + click to select a range
c5a5223
add persistent path for DCP GLM5.2
zhuyuhua-v Aug 21, 2026
b2bbbbf
add ATOM_DCP_REPLICATE_INDEX_CACHE to control dcp index cache
zhuyuhua-v Aug 24, 2026
acf9ab2
remove contiguous and support lmcache
zhuyuhua-v Aug 24, 2026
571d5be
Merge main and preserve replicated DCP index layout
zhuyuhua-v Aug 24, 2026
82a5fa6
feat(dcp): migrate the test_pd_dcp branch onto the DCP KV-transfer line
Jasen2201 Aug 25, 2026
c3d933a
Adapt the PD block transfer to a DCP decode node
Jasen2201 Aug 26, 2026
445fa83
Run MTP on the DCP decode node
Jasen2201 Aug 26, 2026
2630230
support PD replicated index cache
Phi-C Aug 26, 2026
7a3d735
Move DSA index pages as pages, not as tokens
Jasen2201 Aug 26, 2026
9f20653
Mirror the DSA decode logits plane out of the CUDA graph
Jasen2201 Aug 26, 2026
ebfd52f
Evaluate GSM8K the way the GLM-5 recipe does
Jasen2201 Aug 26, 2026
852bd8d
Give the DCP PD GSM8K number a same-box control
Jasen2201 Aug 26, 2026
52fe310
Record the MTP=3 accuracy and acceptance rate on DCP decode
Jasen2201 Aug 26, 2026
1517dff
Merge branch 'main' into yuhua/dsa-dcp-persistent
zhuyuhua-v Aug 26, 2026
bbc6af7
Turn MTP on by default in the DCP launch script
Jasen2201 Aug 26, 2026
0330f82
Drop the CI and doc changes from the DCP KV-transfer branch
Jasen2201 Aug 26, 2026
1ddd09b
[CI] add test case step for p(cpp4)d(dcp4)
MengqingCao Aug 26, 2026
0725b57
remove tp4+dcp4
MengqingCao Aug 27, 2026
387ba07
[CI] fix common args
MengqingCao Aug 27, 2026
16420d5
[CI] fix HANDSHAKE_PORT
MengqingCao Aug 27, 2026
3fa18cc
[CI] make 47 firt order
MengqingCao Aug 27, 2026
778fce1
[ci] fix pp
MengqingCao Aug 27, 2026
e8827d4
[CI] keep prefill-only env out of the decode server
Jasen2201 Aug 27, 2026
b57cd66
[CI] replicate the DCP index cache on the cpp4-dcp4 decode server
Jasen2201 Aug 27, 2026
3250c16
Merge PR #1995 (yuhua/dsa-dcp-persistent) into the DCP KV-transfer br…
Jasen2201 Aug 27, 2026
f07b5d2
style: fold the dsa_logits_dump import into the existing atom.utils b…
Jasen2201 Aug 27, 2026
23c05f2
Remove the DSA indexer logits mirror
Jasen2201 Aug 27, 2026
0e8ea4f
[Dashboard] fix wrong count of gpu
MengqingCao Aug 27, 2026
a195fce
Keep the PAGE layout version at 3
Jasen2201 Aug 27, 2026
23c0ce6
Keep register_received_prefix returning the registered block count
Jasen2201 Aug 27, 2026
7dd0bc7
Restore the one-line docstring on _replicated_index_cache_transfer_su…
Jasen2201 Aug 27, 2026
d440996
Drop the index page plane comment
Jasen2201 Aug 27, 2026
8ecb035
Drop the index_slot_mapping MTP comment
Jasen2201 Aug 27, 2026
36ba408
Condense the DCP relayout comments
Jasen2201 Aug 27, 2026
f2cb8b5
[CI] make 43 the first order
MengqingCao Aug 28, 2026
dcd0420
set ATOM_SPARSE_INDEXER_LOGITS_BUDGET_MB to 2047
MengqingCao Aug 28, 2026
d3a412a
make node customized
MengqingCao Aug 28, 2026
178453e
Let ATOMesh workflow infer Slurm partition from custom nodes.
MengqingCao Aug 28, 2026
8919b45
Merge origin/main into Jasen/dcp-kv-transfer
Jasen2201 Aug 28, 2026
0c21124
Merge remote-tracking branch 'origin/Jasen/dcp-kv-transfer' into Jase…
Jasen2201 Aug 28, 2026
43c6384
Merge origin/main into Jasen/dcp-kv-transfer
Jasen2201 Aug 28, 2026
eed8ce1
revert 178453e7c563f9f2f6bfb23aa88a8c851fdb77d5
MengqingCao Aug 28, 2026
9cd8f03
revert d3a412a7251804dd8d126a74b48fcc72434a0291
MengqingCao Aug 28, 2026
39c3eb0
Shrink PR #2008 to the DCP KV-transfer code change
Jasen2201 Aug 28, 2026
4997a76
Trim DCP relayout commentary
Jasen2201 Aug 28, 2026
cd3e425
Drop the DCP sparse indptr comment
Jasen2201 Aug 28, 2026
a787950
Restore upstream NaN-guard formatting in dcp_ops
Jasen2201 Aug 28, 2026
48a5a45
Drop the RS-layout DCP merge optimization
Jasen2201 Aug 28, 2026
76095c6
Keep pd_matrix's node-count branches separate
Jasen2201 Aug 28, 2026
6212a11
Drop the LMCache chunk-floor fix
Jasen2201 Aug 28, 2026
3d3738d
Keep PD incremental KV alive when only decode runs DCP
Jasen2201 Aug 30, 2026
087d037
Merge origin/main into Jasen/dcp-kv-transfer
Jasen2201 Aug 31, 2026
43fb811
set TOPK_FORCE_PATH=one
MengqingCao Aug 31, 2026
2e22291
Give the replicated index cache a skip sentinel and a draft-step mapping
Jasen2201 Aug 31, 2026
ca2bcb2
Give both ubatch builders the indexer fields they were missing
Jasen2201 Aug 31, 2026
6d2e1fb
Filter the padded decode width, not just the scheduled one
Jasen2201 Aug 31, 2026
ee9cc0e
Size sparse MTP work buffers for the DCP-gathered query width
Jasen2201 Aug 31, 2026
2782986
Fault instead of misplacing KV when the DCP relayout cannot apply
Jasen2201 Aug 31, 2026
a41c374
Cover the MTP shape the DCP filter's token->request map exists for
Jasen2201 Aug 31, 2026
a864a81
Merge remote-tracking branch 'origin/main' into Jasen/dcp-kv-transfer
Jasen2201 Sep 1, 2026
0a97a53
Keep the two DCP relayouts off out-of-range blocks
Jasen2201 Sep 1, 2026
870dbd6
Merge branch 'main' into Jasen/dcp-kv-transfer
Jasen2201 Sep 1, 2026
0bf776b
test(mla): give the transfer-region runner double aligned_index_dim
Jasen2201 Sep 1, 2026
809de5f
Assert the row counts the DCP filter's token grid indexes
Jasen2201 Sep 1, 2026
9941008
Reject an index page that is not a whole number of token rows
Jasen2201 Sep 1, 2026
a02bfc0
Drop the replicated-index fields nothing reads
Jasen2201 Sep 1, 2026
4cb411b
Document ATOM_DCP_REPLICATE_INDEX_CACHE and the DSA+DCP+MTP it unblocks
Jasen2201 Sep 1, 2026
fed3c56
Drop the _coalesce claim that merging is the common case
Jasen2201 Sep 1, 2026
fb876da
Plan the replicated index relayout once per page geometry
Jasen2201 Sep 1, 2026
8f80ae4
Raise rather than assert on the DCP filter's row counts
Jasen2201 Sep 1, 2026
174456e
Refuse a block region the DCP relayout does not describe
Jasen2201 Sep 1, 2026
5dc621b
Check the index row against the indexer's quantization block
Jasen2201 Sep 1, 2026
858f6ca
Declare index_slot_mapping on AttentionMetaData
Jasen2201 Sep 1, 2026
b7939fd
Say that speculative decode keeps the sparse DCP persistent path
Jasen2201 Sep 1, 2026
8e3956e
Count a PP role's GPUs from the flag the server ran with
Jasen2201 Sep 1, 2026
e4cedc5
Merge branch 'main' into Jasen/dcp-kv-transfer
Jasen2201 Sep 1, 2026
d127fe1
Merge branch 'main' into Jasen/dcp-kv-transfer
Jasen2201 Sep 2, 2026
6337d9b
[MTP][DCP] Fix draft token_to_seq_idxs, drop broken index sharing
Jasen2201 Sep 2, 2026
44868cc
Merge branch 'main' into Jasen/dcp-kv-transfer
Jasen2201 Sep 2, 2026
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
72 changes: 72 additions & 0 deletions .github/benchmark/models_atomesh.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -325,6 +325,78 @@ models:
name: glm-52-mxfp4-1p1d-tp4-agentic-1m-c48
concurrency: [48]

# CPP4 prefill (PP4 x TP1) paired with TP4+DCP4 decode on one node.
- &glm52_agentic_lmcache_cpp4_dcp4
<<: *glm52_agentic_tp4_tp4
name: glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c48
topology: 1p1d_cpp4_dcp4
pd_worker_layout: single_node
concurrency: [48]
server:
common_args:
kv_cache_dtype: fp8
block_size: 16
max_num_seqs: 512
enable_prefix_caching: true
decode_max_num_seqs: 512
# Empty means "do not pass the flag", so the model config's own
# context length applies. The 16384 the model defaults carry would
# truncate this suite's 1M agentic traces.
max_model_len: ""
max_num_batched_tokens: ""
decode_max_num_batched_tokens: ""
online_quant_config: '{"global_quant_config":"ptpc_fp8","exclude_layer":["lm_head","model.embed_tokens","*.mlp.gate","*expert*"]}'
gpu_memory_utilization: 0.85
extra_args: "--level 3 --method mtp --num-speculative-tokens 3 --spec-decode-acceptance-rate 0.6633"
prefill:
workers: 1
tp: 1
cudagraph: none
extra_args: "--pipeline-parallel-size 4 --enforce-eager --max-num-batched-tokens 8192"
decode:
workers: 1
tp: 4
cudagraph: >-
[1,2,4,8,16,24,32,40,48,56,64,72,80,88,96,104,112,120,128,136,144,152,160,168,176,184,192,200,208,216,224,232,240,248,256]
extra_args: "--decode-context-parallel-size 4 --cudagraph-mode FULL"
env:
common:
ATOM_MLA_PAGE_SIZE: "1"
ATOM_ONLINE_QUANT_STREAMING: "0"
ATOM_SPARSE_INDEXER_LOGITS_BUDGET_MB: "2047"
ATOM_USE_TRITON_MLA: "0"
MAX_JOBS: "16"
PYTHONHASHSEED: "0"
TOPK_FORCE_PATH: "one"
prefill:
HIP_VISIBLE_DEVICES: "0,1,2,3"
VLLM_PP_LAYER_PARTITION: "20,20,20,18"
LMCACHE_LOCAL_CPU: "True"
LMCACHE_MAX_LOCAL_CPU_SIZE: "256"
LMCACHE_CHUNK_SIZE: "256"
OFFLOAD_PROFILE: "1"
OFFLOAD_MIN_LOAD_TOKENS: "0"
PREFILL_KV_TRANSFER_CONFIG: >-
{"kv_connector":"multi","connectors":[{"kv_connector":"mooncake","kv_role":"kv_producer","proxy_ip":"${ROLE_IP}","handshake_port":${HANDSHAKE_PORT},"protocol":"rdma"},{"kv_connector":"lmcache_offload","kv_role":"offload"}]}
decode:
HIP_VISIBLE_DEVICES: "4,5,6,7"
# Decode-only: the sharded index cache reaches its global top-k
# through a candidate all-gather that only handles qlen=1, so MTP's
# 4-token verify aborts on it. Replicating the index cache removes
# that exchange. Setting it on prefill instead raises at startup --
# prefill runs pp4/dcp1 and the layout rejects both.
ATOM_DCP_REPLICATE_INDEX_CACHE: "1"
DECODE_KV_TRANSFER_CONFIG: >-
{"kv_connector":"mooncake","kv_role":"kv_consumer","proxy_ip":"${ROLE_IP}","handshake_port":${HANDSHAKE_PORT},"protocol":"rdma"}

- <<: *glm52_agentic_lmcache_cpp4_dcp4
name: glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c32
concurrency: [32]

- <<: *glm52_agentic_lmcache_cpp4_dcp4
name: glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c40
concurrency: [40]

- &glm52_agentic_lmcache_tp4_dpa
<<: *glm52_agentic_tp4_tp4
name: glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c24
Expand Down
4 changes: 3 additions & 1 deletion .github/scripts/atomesh/pd_matrix.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,7 +206,9 @@ def role_env(
model_cfg.get("env", {}).get(role, {}),
suite_cfg.get("env", {}).get(role, {}),
)
env = resolve_env_refs_in_value(env, preserve_names={"ROLE_IP"})
# ROLE_IP and HANDSHAKE_PORT are filled in by pd_server_atom.sh at
# launch time (host IP and 6301 + ATOMESH_SERVICE_PORT_OFFSET).
env = resolve_env_refs_in_value(env, preserve_names={"ROLE_IP", "HANDSHAKE_PORT"})
return {str(key): str(value) for key, value in env.items()}


Expand Down
34 changes: 32 additions & 2 deletions .github/scripts/atomesh/pd_server_atom.sh
Original file line number Diff line number Diff line change
Expand Up @@ -242,14 +242,44 @@ dump_launch_info() {
apply_prefixed_env() {
local prefix="$1"
local role_ip="$2"
local handshake_port="${3:-${HANDSHAKE_PORT}}"
local name raw value
while IFS='=' read -r name raw; do
[[ "${name}" == "${prefix}"* ]] || continue
value="${raw//\$\{ROLE_IP\}/${role_ip}}"
value="${value//\$\{HANDSHAKE_PORT\}/${handshake_port}}"
export "${name#${prefix}}=${value}"
done < <(env)
}

# Names the last apply_role_env() exported for a role.
ROLE_ENV_NAMES=()

# Both servers are launched from this one shell, so a variable exported for one
# role stays in the environment the next role inherits. Only names the two roles
# both define get overwritten; a prefill-only name reaches decode unchanged --
# VLLM_PP_LAYER_PARTITION from a pp4 prefill aborts a pp1 decode's model build
# with "len(partitions)=4 does not match pp_size=1". Drop the previous role's
# names before applying this one's.
apply_role_env() {
local prefix="$1"
local role_ip="$2"
local handshake_port="${3:-${HANDSHAKE_PORT}}"
local name
for name in ${ROLE_ENV_NAMES[@]+"${ROLE_ENV_NAMES[@]}"}; do
unset "${name}"
done
ROLE_ENV_NAMES=()
while IFS='=' read -r name _; do
[[ "${name}" == "${prefix}"* ]] || continue
ROLE_ENV_NAMES+=("${name#${prefix}}")
done < <(env)
# A name the common block also sets was just unset with the previous role's,
# so put the common value back before the role overrides it.
apply_prefixed_env "ATOMESH_ENV_" "${role_ip}" "${handshake_port}"
apply_prefixed_env "${prefix}" "${role_ip}" "${handshake_port}"
}

host_ip="$(echo "${IPADDRS}" | tr ',' '\n' | sed -n "$((NODE_RANK + 1))p")"
if [[ -z "${host_ip}" ]]; then
host_ip="$(hostname -I 2>/dev/null | awk '{print $1}')"
Expand Down Expand Up @@ -543,7 +573,7 @@ start_prefill() {
local handshake_port="${3:-${HANDSHAKE_PORT}}"
local dp_master_port="${4:-${PREFILL_DP_MASTER_PORT}}"
local dp_base_port="${5:-${PREFILL_DP_BASE_PORT}}"
apply_prefixed_env "ATOMESH_PREFILL_ENV_" "${host_ip}"
apply_role_env "ATOMESH_PREFILL_ENV_" "${host_ip}" "${handshake_port}"
local -a prefill_cache_env=()
build_server_cache_env "prefill" "${server_port}" prefill_cache_env
local -a prefill_dp_env=()
Expand Down Expand Up @@ -580,7 +610,7 @@ start_decode() {
local handshake_port="${3:-${HANDSHAKE_PORT}}"
local dp_master_port="${4:-${DECODE_DP_MASTER_PORT}}"
local dp_base_port="${5:-${DECODE_DP_BASE_PORT}}"
apply_prefixed_env "ATOMESH_DECODE_ENV_" "${host_ip}"
apply_role_env "ATOMESH_DECODE_ENV_" "${host_ip}" "${handshake_port}"
local max_conc
max_conc="$(echo "${BENCH_MAX_CONCURRENCY}" | tr 'x,' '\n' | sort -n | tail -1)"
local decode_max_num_seqs="${MAX_NUM_SEQS}"
Expand Down
51 changes: 45 additions & 6 deletions .github/scripts/atomesh/process_result.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,9 @@
)
TOPOLOGY_RE = re.compile(r"(?P<p>\d+)p(?P<d>\d+)d", re.IGNORECASE)
TP_RE = re.compile(r"tp(?P<tp>\d+)", re.IGNORECASE)
DUAL_TP_RE = re.compile(r"tp(?P<prefill_tp>\d+)-tp(?P<decode_tp>\d+)", re.IGNORECASE)
CPP_PP_RE = re.compile(r"(?:^|[\s_-])(?:cpp|pp)(?P<pp>\d+)(?:$|[\s_-])", re.IGNORECASE)
PP_ARG_RE = re.compile(r"--pipeline-parallel-size(?:=|\s+)(\d+)", re.IGNORECASE)
EVAL_CONC_RE = re.compile(r"(?:^|[_-])c(?P<conc>\d+)(?:$|[_-])", re.IGNORECASE)
EVAL_TOPOLOGY_RE = re.compile(
r"(?:^|[_-])(?P<topology>\d+p\d+d(?:[_-]dpa)?)(?:$|[_-])",
Expand Down Expand Up @@ -101,6 +104,12 @@ def int_value(*values: Any) -> int | None:
return int(parsed) if parsed is not None else None


def pp_from_server_args(payload: dict[str, Any], key: str) -> int | None:
"""Pipeline-parallel size as the role's own launch flag set it."""
match = PP_ARG_RE.search(string_value(payload.get(key)))
return int(match.group(1)) if match else None


def round_or_none(*values: Any, digits: int = 4) -> float | None:
parsed = number(*values)
return round(parsed, digits) if parsed is not None else None
Expand Down Expand Up @@ -202,7 +211,6 @@ def topology_resources(
)
)
topology = TOPOLOGY_RE.search(text)
tp = TP_RE.search(text)
prefill_workers = int_value(
payload.get("prefill_workers"), payload.get("num_prefill_workers")
)
Expand All @@ -219,16 +227,43 @@ def topology_resources(
decode_tp = int_value(
payload.get("decode_tp"), payload.get("decode_tensor_parallel_size")
)
if tp:
prefill_tp = prefill_tp or int(tp.group("tp"))
decode_tp = decode_tp or int(tp.group("tp"))
dual_tp = DUAL_TP_RE.search(text)
if dual_tp:
prefill_tp = prefill_tp or int(dual_tp.group("prefill_tp"))
decode_tp = decode_tp or int(dual_tp.group("decode_tp"))
else:
tp = TP_RE.search(text)
if tp:
tp_size = int(tp.group("tp"))
prefill_tp = prefill_tp or tp_size
decode_tp = decode_tp or tp_size

# The launch flag outranks the topology label: it is what the server ran with.
prefill_pp = int_value(
payload.get("prefill_pp"), payload.get("prefill_pipeline_parallel_size")
)
if prefill_pp is None:
prefill_pp = pp_from_server_args(payload, "prefill_extra_server_args")
if prefill_pp is None:
cpp_pp = CPP_PP_RE.search(text)
if cpp_pp:
prefill_pp = int(cpp_pp.group("pp"))
prefill_pp = prefill_pp or 1

# No topology label encodes a decode-side PP, so the flag is the only source.
decode_pp = int_value(
payload.get("decode_pp"), payload.get("decode_pipeline_parallel_size")
)
if decode_pp is None:
decode_pp = pp_from_server_args(payload, "decode_extra_server_args")
decode_pp = decode_pp or 1

num_prefill_gpu = int_value(payload.get("num_prefill_gpu"))
num_decode_gpu = int_value(payload.get("num_decode_gpu"))
if num_prefill_gpu is None and prefill_workers and prefill_tp:
num_prefill_gpu = prefill_workers * prefill_tp
num_prefill_gpu = prefill_workers * prefill_tp * prefill_pp
if num_decode_gpu is None and decode_workers and decode_tp:
num_decode_gpu = decode_workers * decode_tp
num_decode_gpu = decode_workers * decode_tp * decode_pp
total_gpu = int_value(payload.get("total_gpu"))
if total_gpu is None and num_prefill_gpu is not None and num_decode_gpu is not None:
total_gpu = num_prefill_gpu + num_decode_gpu
Expand Down Expand Up @@ -339,6 +374,10 @@ def enrich_payload(
enriched.setdefault("decode_workers", env.get("DECODE_WORKERS"))
enriched.setdefault("prefill_tp", env.get("PREFILL_TP"))
enriched.setdefault("decode_tp", env.get("DECODE_TP"))
enriched.setdefault(
"prefill_extra_server_args", env.get("PREFILL_EXTRA_SERVER_ARGS")
)
enriched.setdefault("decode_extra_server_args", env.get("DECODE_EXTRA_SERVER_ARGS"))
runner = env.get("SLURM_SUBMIT_RUNNER", "")
if hardware:
enriched["hardware"] = hardware
Expand Down
7 changes: 7 additions & 0 deletions .github/workflows/atomesh-benchmark.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,9 @@ on:
glm-52-mxfp4-1p1d-tp4-agentic-1m-c1,glm-52-mxfp4-1p1d-tp4-agentic-1m-c2,
glm-52-mxfp4-1p1d-tp4-agentic-1m-c4,glm-52-mxfp4-1p1d-tp4-agentic-1m-c8,
glm-52-mxfp4-1p1d-tp4-agentic-1m-c48,
glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c32,
glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c40,
glm-52-mxfp4-1p1d-cpp4-dcp4-agentic-lmcache-1m-c48,
glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c24,
glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c32,
glm-52-mxfp4-1p1d-tp4-dpa-agentic-lmcache-1m-c48
Expand Down Expand Up @@ -124,6 +127,10 @@ on:
- mia1-p02-g42,mia1-p02-g44
- mia1-p02-g42,mia1-p02-g47
- mia1-p02-g44,mia1-p02-g47
# single_node cases only use the first entry, so keep orderings that
# let any one of the three nodes be selected.
- mia1-p02-g47,mia1-p02-g44
- mia1-p01-g43,mia1-p01-g36
- mia1-p01-g36,mia1-p01-g43
atomesh_2p1d_nodes:
description: 'ATOMesh 2P1D nodes'
Expand Down
13 changes: 13 additions & 0 deletions atom/distributed/dcp_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
"""

from atom.config import get_current_atom_config
from atom.utils import envs


def get_dcp_world_size() -> int:
Expand All @@ -36,6 +37,18 @@ def dcp_is_enabled() -> bool:
return get_dcp_world_size() > 1


def dcp_replicated_index_cache_enabled(atom_config=None) -> bool:
"""Whether native GLM-5.2 uses a full replicated index cache under DCP."""
if not envs.ATOM_DCP_REPLICATE_INDEX_CACHE:
return False
config = atom_config or get_current_atom_config()
hf_config = config.hf_config
return (
config.decode_context_parallel_size > 1
and getattr(hf_config, "model_type", None) == "glm_moe_dsa"
)


def get_dcp_group():
"""The DCP process group (aiter parallel state). Only valid when DCP is enabled."""
from aiter.dist.parallel_state import get_dcp_group as _get_dcp_group
Expand Down
Loading
Loading