From facb553fd178c930f67eca25a44de2f835e3684c Mon Sep 17 00:00:00 2001 From: stephantul Date: Tue, 29 Sep 2026 16:22:53 +0200 Subject: [PATCH 1/3] add other metrics to the classifier --- model2vec/train/trainer.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/model2vec/train/trainer.py b/model2vec/train/trainer.py index 2c45261..e89cca9 100644 --- a/model2vec/train/trainer.py +++ b/model2vec/train/trainer.py @@ -176,7 +176,10 @@ 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(): + postfix[key.replace("val_", "train_", 1)] = f"{value:.4f}" pbar.set_postfix(postfix) if val_check_interval is not None and global_step % val_check_interval == 0: From 388aaffb1fef9fe3d9ee248f7cb7ebed836953f6 Mon Sep 17 00:00:00 2001 From: stephantul Date: Tue, 29 Sep 2026 16:26:58 +0200 Subject: [PATCH 2/3] feat: add metrics to classifiers --- model2vec/train/trainer.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/model2vec/train/trainer.py b/model2vec/train/trainer.py index e89cca9..156f8e0 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,6 +12,7 @@ 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]: @@ -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: @@ -179,7 +181,10 @@ def validate_and_checkpoint() -> bool: with torch.no_grad(): train_metrics = compute_metrics(head_out, y, loss) for key, value in train_metrics.items(): - postfix[key.replace("val_", "train_", 1)] = f"{value:.4f}" + window = train_metric_windows[key.replace("val_", "train_", 1)] + window.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: From 92fbf4a1adebd5c83fb33b8fdcd5f3b49517c820 Mon Sep 17 00:00:00 2001 From: stephantul Date: Tue, 29 Sep 2026 16:41:49 +0200 Subject: [PATCH 3/3] refine metrics --- model2vec/train/classifier.py | 8 ++++---- model2vec/train/trainer.py | 13 ++++++------- tests/test_trainable.py | 2 +- 3 files changed, 11 insertions(+), 12 deletions(-) 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 156f8e0..6fe0493 100644 --- a/model2vec/train/trainer.py +++ b/model2vec/train/trainer.py @@ -16,8 +16,8 @@ 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: @@ -83,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()} @@ -110,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. @@ -121,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) @@ -181,8 +181,7 @@ def validate_and_checkpoint() -> bool: with torch.no_grad(): train_metrics = compute_metrics(head_out, y, loss) for key, value in train_metrics.items(): - window = train_metric_windows[key.replace("val_", "train_", 1)] - window.append(value) + 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) 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())