refactor(disagg): extract _all_reduce_polls helper (#35886)

This commit is contained in:
Shangming Cai
2026-08-22 02:30:33 +08:00
committed by GitHub
parent c3735625de
commit 729a050ea3
+10 -13
View File
@@ -202,6 +202,13 @@ def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> N
polls[i] = int(KVPoll.Transferring)
def _all_reduce_polls(polls: List[int], group: dist.ProcessGroup) -> List[int]:
"""MIN-reduce poll states so no rank commits ahead of its peers."""
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=group)
return tensor_to_reduce.tolist()
def poll_and_all_reduce(
pollers,
gloo_group: dist.ProcessGroup,
@@ -219,9 +226,7 @@ def poll_and_all_reduce(
and server_args is not None
):
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=gloo_group)
return tensor_to_reduce.tolist()
return _all_reduce_polls(polls, gloo_group)
def poll_and_all_reduce_attn_cp_tp_group(
@@ -235,13 +240,7 @@ def poll_and_all_reduce_attn_cp_tp_group(
# Then sync across attn-cp ranks, so all TPxCP participants in one DP shard
# converge to the same global status.
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(
tensor_to_reduce,
op=dist.ReduceOp.MIN,
group=attn_cp_cpu_group,
)
return tensor_to_reduce.tolist()
return _all_reduce_polls(polls, attn_cp_cpu_group)
def poll_and_all_reduce_with_staging(
@@ -277,9 +276,7 @@ def poll_and_all_reduce_with_staging(
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
if metadata_buffers is not None and server_args is not None:
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
poll_tensor = torch.tensor(raw_polls, dtype=torch.uint8, device="cpu")
dist.all_reduce(poll_tensor, op=dist.ReduceOp.MIN, group=gloo_group)
return poll_tensor.tolist()
return _all_reduce_polls(raw_polls, gloo_group)
#########################