fix(loads): switch get_loads_communicator to watching mode (#22919)
This commit is contained in:
@@ -142,8 +142,12 @@ class _Communicator(Generic[T]):
|
|||||||
if obj:
|
if obj:
|
||||||
self._sender.send_pyobj(obj)
|
self._sender.send_pyobj(obj)
|
||||||
|
|
||||||
await self._result_event.wait()
|
event = self._result_event
|
||||||
result_values = copy.deepcopy(self._result_values)
|
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
|
self._result_event = self._result_values = None
|
||||||
return result_values
|
return result_values
|
||||||
|
|
||||||
@@ -247,7 +251,7 @@ class TokenizerCommunicatorMixin:
|
|||||||
self.send_to_scheduler, server_args.dp_size, mode="watching"
|
self.send_to_scheduler, server_args.dp_size, mode="watching"
|
||||||
)
|
)
|
||||||
self.get_loads_communicator = _Communicator(
|
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.dumper_control_communicator = _Communicator(
|
||||||
self.send_to_scheduler, server_args.dp_size
|
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)
|
List of GetLoadsReqOutput, one per scheduler (filtered by dp_rank if specified)
|
||||||
"""
|
"""
|
||||||
self.auto_create_handle_loop()
|
self.auto_create_handle_loop()
|
||||||
req = GetLoadsReqInput(
|
# Always request all sections from scheduler — watching mode shares
|
||||||
include=include if include else ["all"],
|
# results across concurrent callers, so we fetch full data and filter here.
|
||||||
dp_rank=dp_rank,
|
req = GetLoadsReqInput(include=["all"], dp_rank=None)
|
||||||
)
|
|
||||||
results = await self.get_loads_communicator(req)
|
results = await self.get_loads_communicator(req)
|
||||||
|
|
||||||
# Filter by dp_rank if specified
|
# Filter by dp_rank if specified
|
||||||
|
|||||||
Reference in New Issue
Block a user