Skip to content
Repo

TrainingGroup

from modal_training_gym.common.training_group import TrainingGroup

A parameter sweep over a base TrainConfig.

TrainingGroup(base: TrainConfig, grid: dict[str, list[Any]] | None = None, *, name: str | None = None) -> None
get_train_configs() -> list[TrainConfig]

Build validated configs for each sweep point.

Returns

Validated training configs.

iter_variants() -> list[tuple[dict[str, Any], TrainConfig]]

Expand the sweep grid.

Returns

Pairs of overrides and validated training configs.

launch(*, continue_on_error: bool = True, prepare_inputs: bool = False) -> list[TrainingRun]

Launch every variant as a detached Modal call.

Parameters

continue_on_error bool

Continue after a variant fails to launch. Default: True

prepare_inputs bool

Materialize model and dataset inputs before launching. Default: False

Returns

Launched training runs.

train(*, max_parallel: int = 1, continue_on_error: bool = True) -> list[TrainResult]

Train every variant.

Parameters

max_parallel int

Maximum number of variants to run at once. Default: 1

continue_on_error bool

Continue after a variant fails. Default: True

Returns

Successful training results.