diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 8f33c4887..0fe4076b2 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -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) #########################