fix(loads): switch get_loads_communicator to watching mode (#22919)

This commit is contained in:
ybyang
2026-04-16 02:12:22 -07:00
committed by GitHub
parent fbd6dc3565
commit 03fef357a6
@@ -142,9 +142,13 @@ class _Communicator(Generic[T]):
if obj:
self._sender.send_pyobj(obj)
await self._result_event.wait()
result_values = copy.deepcopy(self._result_values)
self._result_event = self._result_values = None
event = self._result_event
values = self._result_values
await event.wait()
# Capture list ref before await so later awaiters survive clearing.
result_values = copy.deepcopy(values)
if self._result_event is event:
self._result_event = self._result_values = None
return result_values
async def __call__(self, obj):
@@ -247,7 +251,7 @@ class TokenizerCommunicatorMixin:
self.send_to_scheduler, server_args.dp_size, mode="watching"
)
self.get_loads_communicator = _Communicator(
self.send_to_scheduler, server_args.dp_size
self.send_to_scheduler, server_args.dp_size, mode="watching"
)
self.dumper_control_communicator = _Communicator(
self.send_to_scheduler, server_args.dp_size
@@ -1051,10 +1055,9 @@ class TokenizerCommunicatorMixin:
List of GetLoadsReqOutput, one per scheduler (filtered by dp_rank if specified)
"""
self.auto_create_handle_loop()
req = GetLoadsReqInput(
include=include if include else ["all"],
dp_rank=dp_rank,
)
# Always request all sections from scheduler — watching mode shares
# results across concurrent callers, so we fetch full data and filter here.
req = GetLoadsReqInput(include=["all"], dp_rank=None)
results = await self.get_loads_communicator(req)
# Filter by dp_rank if specified