Skip to content
Merged
Show file tree
Hide file tree
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
8 changes: 4 additions & 4 deletions model2vec/train/classifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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):
Expand Down
21 changes: 14 additions & 7 deletions model2vec/train/trainer.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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()}
Expand All @@ -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.
Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Comment thread
stephantul marked this conversation as resolved.
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}"
Comment thread
stephantul marked this conversation as resolved.
pbar.set_postfix(postfix)

if val_check_interval is not None and global_step % val_check_interval == 0:
Expand Down
2 changes: 1 addition & 1 deletion tests/test_trainable.py
Original file line number Diff line number Diff line change
Expand Up @@ -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())

Expand Down
Loading