refactor(disagg): extract _all_reduce_polls helper (#35886)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
#########################
|
||||
|
||||
Reference in New Issue
Block a user