MCPcopy Create free account
hub / github.com/apple/axlearn / ManagerClient

Class ManagerClient

axlearn/ft/manager_client.py:23–360  ·  view source on GitHub ↗

FT trainer manager gRPC client with simplified, consistent interface.

Source from the content-addressed store, hash-verified

21
22
23class ManagerClient:
24 """FT trainer manager gRPC client with simplified, consistent interface."""
25
26 def __init__(self, timeout: float = 10.0, port: int = DEFAULT_MANAGER_PORT):
27 """Initialize manager client.
28
29 Args:
30 timeout: gRPC request timeout in seconds
31 port: Default port for manager connections
32 """
33 self.timeout = timeout
34 self.port = port
35 self._channels: dict[tuple[str, int], grpc.Channel] = {}
36 self._channels_lock = threading.Lock()
37 logging.debug("ManagerClient initialized: timeout=%.1fs, port=%d", timeout, port)
38
39 def _get_stub(self, hostname: str, port: int) -> manager_pb2_grpc.ManagerServiceStub:
40 """Return a cached gRPC stub for the target host, creating the channel if needed."""
41 key = (hostname, port)
42 with self._channels_lock:
43 if key not in self._channels:
44 self._channels[key] = grpc.insecure_channel(f"{hostname}:{port}")
45 return manager_pb2_grpc.ManagerServiceStub(self._channels[key])
46
47 def _grpc_call(
48 self,
49 hostname: str,
50 port: int,
51 method_name: str,
52 request,
53 log_prefix: Optional[str] = None,
54 ) -> Any:
55 """Execute a gRPC call with uniform debug logging.
56
57 Args:
58 hostname: Target hostname.
59 port: Target port.
60 method_name: Name of the gRPC stub method to invoke (e.g. "ReportStatus",
61 "RestartReplicaTraining"). Must match the RPC name in the proto service definition.
62 request: Protobuf request message.
63 log_prefix: Short string used in log messages; defaults to method_name.
64
65 Returns:
66 The response message. Raises exceptions on failure.
67 """
68 prefix = log_prefix or method_name
69 logging.debug("%s: sending to %s:%d", prefix, hostname, port)
70 stub = self._get_stub(hostname, port)
71 response = getattr(stub, method_name)(request, timeout=self.timeout)
72 logging.debug(
73 "%s: response from %s:%d acknowledged=%s",
74 prefix,
75 hostname,
76 port,
77 response.acknowledged,
78 )
79 return response
80

Callers 4

_restart_all_replicasMethod · 0.90
_handle_pod_shutdownMethod · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected