Skip to content

TrainerCollection

Group trainers that are already built.

Pass trainers directly; ids come from trainer.run_id. When flattening a :class:TrainerSuite, ids are the suite child ids from suite.items() (which match trainer.run_id when the suite has no parent run_id).

Source code in gradiend/trainer/suite/collection.py
def __init__(
    self,
    *trainers: Trainer,
    retain_models_in_memory: bool = True,
) -> None:
    self.retain_models_in_memory = bool(retain_models_in_memory)
    self.trainers: Dict[str, Trainer] = {}
    for trainer in trainers:
        self._add_trainer(trainer)

retain_models_in_memory instance-attribute

retain_models_in_memory = bool(retain_models_in_memory)

trainers instance-attribute

trainers = {}

__iter__

__iter__()
Source code in gradiend/trainer/suite/collection.py
def __iter__(self) -> Iterator[str]:
    return iter(self.trainers)

__len__

__len__()
Source code in gradiend/trainer/suite/collection.py
def __len__(self) -> int:
    return len(self.trainers)

_add_part

_add_part(part)
Source code in gradiend/trainer/suite/collection.py
def _add_part(self, part: Union[Trainer, TrainerSuite, TrainerCollection]) -> None:
    if isinstance(part, TrainerCollection):
        for trainer_id, trainer in part.trainers.items():
            self._add_trainer_with_id(trainer_id, trainer)
    elif isinstance(part, TrainerSuite):
        for child_id, trainer in part.items():
            self._add_trainer_with_id(str(child_id), trainer)
    elif isinstance(part, Trainer):
        self._add_trainer(part)
    else:
        raise TypeError(
            "TrainerCollection.merge expected Trainer, TrainerSuite, or TrainerCollection; "
            f"got {type(part).__name__}"
        )

_add_trainer

_add_trainer(trainer)
Source code in gradiend/trainer/suite/collection.py
def _add_trainer(self, trainer: Trainer) -> None:
    self._add_trainer_with_id(_require_trainer_run_id(trainer), trainer)

_add_trainer_with_id

_add_trainer_with_id(trainer_id, trainer)
Source code in gradiend/trainer/suite/collection.py
def _add_trainer_with_id(self, trainer_id: str, trainer: Trainer) -> None:
    if trainer_id in self.trainers:
        raise ValueError(f"Duplicate trainer id {trainer_id!r}")
    self.trainers[trainer_id] = trainer

get_trainer

get_trainer(trainer_id)
Source code in gradiend/trainer/suite/collection.py
def get_trainer(self, trainer_id: str) -> Trainer:
    return self.trainers[trainer_id]

items

items()
Source code in gradiend/trainer/suite/collection.py
def items(self) -> Iterable[Tuple[str, Trainer]]:
    return self.trainers.items()

keys

keys()
Source code in gradiend/trainer/suite/collection.py
def keys(self) -> Iterable[str]:
    return self.trainers.keys()

merge classmethod

merge(*parts, retain_models_in_memory=True)

Combine trainers, suites, and collections into one group.

Source code in gradiend/trainer/suite/collection.py
@classmethod
def merge(
    cls,
    *parts: Union[Trainer, TrainerSuite, TrainerCollection],
    retain_models_in_memory: bool = True,
) -> TrainerCollection:
    """Combine trainers, suites, and collections into one group."""
    collection = cls(retain_models_in_memory=retain_models_in_memory)
    for part in parts:
        collection._add_part(part)
    return collection

train

train(*, use_cache=True)
Source code in gradiend/trainer/suite/collection.py
def train(self, *, use_cache: bool = True) -> None:
    for trainer in self.trainers.values():
        trainer.train(use_cache=use_cache)
        used_cache = bool(getattr(trainer, "_last_train_used_cache", False))
        if not self.retain_models_in_memory and not used_cache and hasattr(trainer, "unload_model"):
            trainer.unload_model()

values

values()
Source code in gradiend/trainer/suite/collection.py
def values(self) -> Iterable[Trainer]:
    return self.trainers.values()