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)
|
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(
|
def poll_and_all_reduce(
|
||||||
pollers,
|
pollers,
|
||||||
gloo_group: dist.ProcessGroup,
|
gloo_group: dist.ProcessGroup,
|
||||||
@@ -219,9 +226,7 @@ def poll_and_all_reduce(
|
|||||||
and server_args is not None
|
and server_args is not None
|
||||||
):
|
):
|
||||||
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
|
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
|
||||||
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
|
return _all_reduce_polls(polls, gloo_group)
|
||||||
dist.all_reduce(tensor_to_reduce, op=dist.ReduceOp.MIN, group=gloo_group)
|
|
||||||
return tensor_to_reduce.tolist()
|
|
||||||
|
|
||||||
|
|
||||||
def poll_and_all_reduce_attn_cp_tp_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
|
# Then sync across attn-cp ranks, so all TPxCP participants in one DP shard
|
||||||
# converge to the same global status.
|
# converge to the same global status.
|
||||||
tensor_to_reduce = torch.tensor(polls, dtype=torch.uint8, device="cpu")
|
return _all_reduce_polls(polls, attn_cp_cpu_group)
|
||||||
dist.all_reduce(
|
|
||||||
tensor_to_reduce,
|
|
||||||
op=dist.ReduceOp.MIN,
|
|
||||||
group=attn_cp_cpu_group,
|
|
||||||
)
|
|
||||||
return tensor_to_reduce.tolist()
|
|
||||||
|
|
||||||
|
|
||||||
def poll_and_all_reduce_with_staging(
|
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.
|
# 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:
|
if metadata_buffers is not None and server_args is not None:
|
||||||
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
|
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
|
||||||
poll_tensor = torch.tensor(raw_polls, dtype=torch.uint8, device="cpu")
|
return _all_reduce_polls(raw_polls, gloo_group)
|
||||||
dist.all_reduce(poll_tensor, op=dist.ReduceOp.MIN, group=gloo_group)
|
|
||||||
return poll_tensor.tolist()
|
|
||||||
|
|
||||||
|
|
||||||
#########################
|
#########################
|
||||||
|
|||||||
Reference in New Issue
Block a user