# SPDX-License-Identifier: Apache-2.0 from abc import ABC, abstractmethod from collections.abc import Callable from typing import Any, TypeVar, cast from trainer.trainer_args import TrainerArgs from trainer.pipelines import ForwardBatch from trainer.utils import init_logger logger = init_logger(__name__) _R = TypeVar("_R") class Executor(ABC): def __init__(self, trainer_args: TrainerArgs): self.trainer_args = trainer_args self._init_executor() @abstractmethod def _init_executor(self) -> None: raise NotImplementedError @classmethod def get_class(cls, trainer_args: TrainerArgs) -> type["Executor"]: if trainer_args.distributed_executor_backend == "mp": from trainer.worker.multiproc_executor import MultiprocExecutor return cast(type["Executor"], MultiprocExecutor) else: raise ValueError( f"Unsupported distributed executor backend: {trainer_args.distributed_executor_backend}" ) def execute_forward( self, forward_batch: ForwardBatch, trainer_args: TrainerArgs, ) -> ForwardBatch: outputs: list[dict[str, Any]] = self.collective_rpc("execute_forward", kwargs={ "forward_batch": forward_batch, "trainer_args": trainer_args }) return cast(ForwardBatch, outputs[0]["output_batch"]) @abstractmethod def set_lora_adapter(self, lora_nickname: str, lora_path: str | None = None) -> None: """ Set the LoRA adapter for the workers. """ raise NotImplementedError @abstractmethod def collective_rpc(self, method: str | Callable[..., _R], timeout: float | None = None, args: tuple = (), kwargs: dict[str, Any] | None = None) -> list[_R]: """ Execute an RPC call on all workers. Args: method: Name of the worker method to execute, or a callable that is serialized and sent to all workers to execute. If the method is a callable, it should accept an additional `self` argument, in addition to the arguments passed in `args` and `kwargs`. The `self` argument will be the worker object. timeout: Maximum time in seconds to wait for execution. Raises a :exc:`TimeoutError` on timeout. `None` means wait indefinitely. args: Positional arguments to pass to the worker method. kwargs: Keyword arguments to pass to the worker method. Returns: A list containing the results from each worker. Note: It is recommended to use this API to only pass control messages, and set up data-plane communication to pass data. """ raise NotImplementedError @abstractmethod def shutdown(self) -> None: """ Shutdown the executor. """ raise NotImplementedError