Diffusers
Safetensors
HY / trainer /worker /executor.py
Cccccz's picture
Upload batch 65: 500 files (0.01 GiB)
74da989 verified
Raw History Blame Contribute Delete
3.45 kB
# 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