Skip to content
Draft
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 71 additions & 48 deletions src/itwinai/torch/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,8 @@ class TorchTrainer(Trainer, LogMixin):
Defaults to 'ddp'.
test_every (int | None, optional): run a test epoch every ``test_every`` epochs.
Disabled if None. Defaults to None.
validation_every (int, optional): run a validation epoch every ``validation_every``
epochs. Disabled if set to zero. Defaults to 1.
random_seed (int | None, optional): set random seed for reproducibility. If None, the
seed is not set. Defaults to None.
logger (Logger | None, optional): logger for ML tracking. Defaults to None.
Expand Down Expand Up @@ -158,6 +160,10 @@ class TorchTrainer(Trainer, LogMixin):
validation_dataloader: DataLoader | None = None
#: PyTorch ``DataLoader`` for test dataset.
test_dataloader: DataLoader | None = None
#: How often to perform the validation step
validation_every: int = 1
#: How often to perform the test step
test_every: int | None = None
#: PyTorch model to train.
model: nn.Module | None = None
#: Loss criterion.
Expand Down Expand Up @@ -204,6 +210,7 @@ def __init__(
model: nn.Module | None = None,
strategy: Literal["ddp", "deepspeed", "horovod"] = "ddp",
test_every: int | None = None,
validation_every: int = 1,
random_seed: int | None = None,
logger: Logger | None = None,
metrics: Dict[str, Metric] | None = None,
Expand Down Expand Up @@ -249,6 +256,7 @@ def __init__(
self.model = model
self.strategy = strategy
self.test_every = test_every
self.validation_every = validation_every
self.random_seed = random_seed
self.logger = logger
self.metrics = metrics if metrics is not None else {}
Expand Down Expand Up @@ -352,6 +360,19 @@ def _detect_distributed_strategy(self, strategy: str) -> TorchDistributedStrateg
py_logger.info(
f"Ray cluster was detected, thus the Ray equivalent for {strategy} is used"
)
if self.validation_every == 0 or self.validation_every > self.epochs:
raise ValueError(
"Ray is activated but with 'validation_every' set to"
f" {self.validation_every} and 'epochs' set to {self.epochs}, you will"
" never report any metrics. Ray requires you to report at least one"
" validation metric!"
)
if self.validation_every != 1:
py_logger.warning(
"Ray is activated, but 'validation_every' is not set to one. Keep in mind"
" that the validation metrics are used for tuning, so this could lead"
" to suboptimal results."
)
Comment on lines +372 to +375

@matbun matbun Jul 14, 2025

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

here I would emphasize more the fact that this flag influences the reporting frequency of ray, which has implications when tuning (and not only)


match strategy, ray_cluster_is_running():
case "ddp", True:
Expand Down Expand Up @@ -861,9 +882,9 @@ def _execute_with_ray(
if contains_mlflow_logger(self.logger):
# Create mlflow runs per trial (will be started by the trial's main worker)
client = mlflow.tracking.MlflowClient()
experiment_id = (
client.get_experiment_by_name(self.experiment_name).experiment_id
)
experiment_id = client.get_experiment_by_name(
self.experiment_name
).experiment_id

for trial_idx in range(self.ray_tune_config.num_samples):
# create a mlflow run for each trial (without starting it)
Expand Down Expand Up @@ -1193,59 +1214,61 @@ def train(self) -> None:

self.set_epoch()
self.train_epoch()
val_metric = self.validation_epoch()

# Periodic checkpointing
periodic_ckpt_path = self.save_checkpoint(name=f"epoch_{self.current_epoch}")

# Checkpointing current best model
best_ckpt_path = None
worker_val_metrics = self.strategy.gather(val_metric, dst_rank=0)

if self.strategy.is_main_worker:
avg_metric = torch.mean(torch.stack(worker_val_metrics)).detach().cpu()
if avg_metric < self.best_validation_metric:
best_ckpt_path = self.save_checkpoint(
name="best_model",
best_validation_metric=avg_metric,
force=True,
)
self.best_validation_metric = avg_metric

# Report validation metrics to Ray (useful for tuning!)
metric_name = _get_tuning_metric_name(self.ray_tune_config)
if metric_name is None:
raise ValueError("Could not find a metric in the TuneConfig")
if self.strategy.is_main_worker and self.strategy.is_distributed:
assert epoch_time_logger is not None
epoch_time = default_timer() - epoch_start_time
epoch_time_logger.add_epoch_time(self.current_epoch + 1, epoch_time)

if (
self.time_ray
and self.logger is not None
and isinstance(self.strategy, RayTorchDistributedStrategy)
):
time_and_log(
func=partial(
self.ray_report,
if self.validation_every and (self.current_epoch + 1) % self.validation_every == 0:
val_metric = self.validation_epoch()

# Checkpointing current best model
best_ckpt_path = None
worker_val_metrics = self.strategy.gather(val_metric, dst_rank=0)

if self.strategy.is_main_worker:
avg_metric = torch.mean(torch.stack(worker_val_metrics)).detach().cpu()
if avg_metric < self.best_validation_metric:
best_ckpt_path = self.save_checkpoint(
name="best_model",
best_validation_metric=avg_metric,
force=True,
)
self.best_validation_metric = avg_metric

# Periodic checkpointing
periodic_ckpt_path = self.save_checkpoint(name=f"epoch_{self.current_epoch}")

# Report validation metrics to Ray (useful for tuning!)
metric_name = _get_tuning_metric_name(self.ray_tune_config)
if metric_name is None:
raise ValueError("Could not find a metric in the TuneConfig")

if (
self.time_ray
and self.logger is not None
and isinstance(self.strategy, RayTorchDistributedStrategy)
):
time_and_log(
func=partial(
self.ray_report,
metrics={metric_name: val_metric.item()},
checkpoint_dir=best_ckpt_path or periodic_ckpt_path,
),
logger=self.logger,
identifier="ray_report_time_s_per_epoch",
step=self.current_epoch,
)
else:
self.ray_report(
Comment thread
jarlsondre marked this conversation as resolved.
metrics={metric_name: val_metric.item()},
checkpoint_dir=best_ckpt_path or periodic_ckpt_path,
),
logger=self.logger,
identifier="ray_report_time_s_per_epoch",
step=self.current_epoch,
)
else:
self.ray_report(
metrics={metric_name: val_metric.item()},
checkpoint_dir=best_ckpt_path or periodic_ckpt_path,
)
)

if self.test_every and (self.current_epoch + 1) % self.test_every == 0:
self.test_epoch()

if self.strategy.is_main_worker and self.strategy.is_distributed:
assert epoch_time_logger is not None
epoch_time = default_timer() - epoch_start_time
epoch_time_logger.add_epoch_time(self.current_epoch + 1, epoch_time)

def train_epoch(self) -> torch.Tensor:
"""Perform a complete sweep over the training dataset, completing an epoch of training.

Expand Down