diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index 93c3cd5..1f98000 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -22,9 +22,9 @@ def _classifier_metrics(head_out: torch.Tensor, y: torch.Tensor, loss: torch.Tensor) -> dict[str, float]: - """Validation metrics for single-label classification: loss and accuracy.""" + """Metrics for single-label classification: loss and accuracy.""" accuracy = (head_out.argmax(dim=1) == y).float().mean() - return {"val_loss": loss.item(), "val_accuracy": accuracy.item()} + return {"loss": loss.item(), "accuracy": accuracy.item()} def _compute_accuracy(y_true: torch.Tensor, y_pred: torch.Tensor) -> float: @@ -36,10 +36,10 @@ def _compute_accuracy(y_true: torch.Tensor, y_pred: torch.Tensor) -> float: def _multilabel_classifier_metrics(head_out: torch.Tensor, y: torch.Tensor, loss: torch.Tensor) -> dict[str, float]: - """Validation metrics for multi-label classification: loss and Jaccard accuracy.""" + """Metrics for multi-label classification: loss and Jaccard accuracy.""" preds = (torch.sigmoid(head_out) > 0.5).float() accuracy = _compute_accuracy(y, preds) - return {"val_loss": loss.item(), "val_accuracy": accuracy} + return {"loss": loss.item(), "accuracy": accuracy} class StaticModelForClassification(BaseFinetuneable): diff --git a/model2vec/train/trainer.py b/model2vec/train/trainer.py index 2c45261..6fe0493 100644 --- a/model2vec/train/trainer.py +++ b/model2vec/train/trainer.py @@ -1,7 +1,7 @@ from __future__ import annotations import copy -from collections import defaultdict +from collections import defaultdict, deque from collections.abc import Callable import torch @@ -12,11 +12,12 @@ MetricsFn = Callable[[torch.Tensor, torch.Tensor, torch.Tensor], dict[str, float]] _UNBOUNDED_MAX_EPOCHS = 9999 +_TRAIN_METRICS_WINDOW = 50 def default_metrics(head_out: torch.Tensor, y: torch.Tensor, loss: torch.Tensor) -> dict[str, float]: - """Validation metrics for tasks that only track loss (used for early stopping on val_loss).""" - return {"val_loss": loss.item()} + """Metrics for tasks that only track loss (used for early stopping on val_loss).""" + return {"loss": loss.item()} class EarlyStopper: @@ -82,7 +83,7 @@ def _run_validation( head_out = model(x) loss = loss_function(head_out, y) for key, value in compute_metrics(head_out, y, loss).items(): - weighted_sums[key] += value * batch_size + weighted_sums[f"val_{key}"] += value * batch_size total_samples += batch_size model.train() return {key: total / total_samples for key, total in weighted_sums.items()} @@ -109,7 +110,7 @@ def run_training_loop( # noqa: C901 :param model: The model to train, called as `head_out = model(x)`. :param loss_function: Computes the training and validation loss from `(head_out, y)`. :param learning_rate: The Adam learning rate. - :param val_metric: The metric key (returned by `compute_metrics`) used for early stopping. + :param val_metric: The metric key used for early stopping: a key returned by `compute_metrics`, prefixed with `val_`. :param early_stopping_direction: Either "min" or "max", the direction of improvement for `val_metric`. :param train_loader: The training data loader. :param val_loader: The validation data loader. @@ -120,7 +121,7 @@ def run_training_loop( # noqa: C901 :param device: The device to train on. :param val_check_interval: If set, validate every this many training steps. :param check_val_every_epoch: If set, validate every this many epochs. - :param compute_metrics: Computes validation metrics from `(head_out, y, loss)`. Defaults to just `val_loss`. + :param compute_metrics: Computes unprefixed metrics from `(head_out, y, loss)`. Defaults to just `loss`. :return: The model's state dict from the validation check with the best `val_metric`. """ model.to(device) @@ -149,6 +150,7 @@ def run_training_loop( # noqa: C901 current_epoch = 0 global_step = 0 postfix: dict[str, str] = {} + train_metric_windows: dict[str, deque[float]] = defaultdict(lambda: deque(maxlen=_TRAIN_METRICS_WINDOW)) latest_val_loss: float | None = None def validate_and_checkpoint() -> bool: @@ -176,7 +178,12 @@ def validate_and_checkpoint() -> bool: optimizer.step() global_step += 1 - postfix["train_loss"] = f"{loss.item():.4f}" + with torch.no_grad(): + train_metrics = compute_metrics(head_out, y, loss) + for key, value in train_metrics.items(): + train_metric_windows[f"train_{key}"].append(value) + for key, window in train_metric_windows.items(): + postfix[key] = f"{sum(window) / len(window):.4f}" pbar.set_postfix(postfix) if val_check_interval is not None and global_step % val_check_interval == 0: diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 69f6be2..0bdafa4 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -936,7 +936,7 @@ def test_run_training_loop_mid_epoch_early_stop() -> None: device=resolve_device("cpu"), val_check_interval=1, check_val_every_epoch=None, - compute_metrics=lambda head_out, y, loss: {"val_loss": 1.0}, + compute_metrics=lambda head_out, y, loss: {"loss": 1.0}, ) assert set(state_dict) == set(model.state_dict())