Skip to content
Closed
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
2 changes: 2 additions & 0 deletions model2vec/train/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -363,6 +363,7 @@ def _train(
validation_steps: int | None,
compute_metrics: MetricsFn = default_metrics,
token_dropout: float = 0.0,
max_steps: int | None = None,
) -> None:
if not 0.0 <= token_dropout < 1.0:
raise ValueError("token_dropout must be in the range [0, 1).")
Expand All @@ -387,6 +388,7 @@ def _train(
val_check_interval=val_check_interval,
check_val_every_epoch=check_val_every_epoch,
compute_metrics=compute_metrics,
max_steps=max_steps,
)

self.load_state_dict(state_dict)
Expand Down
4 changes: 4 additions & 0 deletions model2vec/train/classifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,7 @@ def fit(
validation_steps: int | None = None,
random_seed: int = DEFAULT_RANDOM_SEED,
token_dropout: float = 0.0,
max_steps: int | None = None,
) -> StaticModelForClassification:
"""Fit a model.

Expand Down Expand Up @@ -181,6 +182,8 @@ def fit(
:param random_seed: The random seed to use. Defaults to 42.
:param token_dropout: The fraction of tokens to randomly drop from each training sample.
Has no effect during validation. Must be in the range [0, 1).
:param max_steps: The maximum number of training steps. If None, the number of steps is not limited.
When it is reached, training stops, even before `min_epochs`.
:return: The fitted model.
:raises ValueError: If either X_val or y_val are provided, but not both.
"""
Expand Down Expand Up @@ -224,6 +227,7 @@ def fit(
validation_steps=validation_steps,
compute_metrics=compute_metrics,
token_dropout=token_dropout,
max_steps=max_steps,
)

return self
Expand Down
4 changes: 4 additions & 0 deletions model2vec/train/pairs.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,7 @@ def fit(
labels_val: list[int] | None = None,
validation_steps: int | None = None,
random_seed: int = DEFAULT_RANDOM_SEED,
max_steps: int | None = None,
) -> T:
"""Fit a model that maximizes the cosine similarity between paired texts.

Expand Down Expand Up @@ -188,6 +189,8 @@ def fit(
:param labels_val: The label for each validation pair. If None, every validation pair is labeled 1.
:param validation_steps: The number of steps to run validation for. If None, validation steps are estimated from the data.
:param random_seed: The random seed to use. Defaults to 42.
:param max_steps: The maximum number of training steps. If None, the number of steps is not limited.
When it is reached, training stops, even before `min_epochs`.
:return: The fitted model.
"""
seed_everything(random_seed)
Expand Down Expand Up @@ -218,6 +221,7 @@ def fit(
max_epochs=max_epochs,
device=device,
validation_steps=validation_steps,
max_steps=max_steps,
)

return self
Expand Down
4 changes: 4 additions & 0 deletions model2vec/train/similarity.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,7 @@ def fit(
validation_steps: int | None = None,
random_seed: int = DEFAULT_RANDOM_SEED,
token_dropout: float = 0.0,
max_steps: int | None = None,
) -> T:
"""Fit a model.

Expand Down Expand Up @@ -108,6 +109,8 @@ def fit(
:param random_seed: The random seed to use. Defaults to 42.
:param token_dropout: The fraction of tokens to randomly drop from each training sample.
Has no effect during validation. Must be in the range [0, 1).
:param max_steps: The maximum number of training steps. If None, the number of steps is not limited.
When it is reached, training stops, even before `min_epochs`.
:return: The fitted model.
"""
seed_everything(random_seed)
Expand All @@ -131,6 +134,7 @@ def fit(
device=device,
validation_steps=validation_steps,
token_dropout=token_dropout,
max_steps=max_steps,
)

return self
Expand Down
15 changes: 14 additions & 1 deletion model2vec/train/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,7 @@ def run_training_loop( # noqa: C901
val_check_interval: int | None,
check_val_every_epoch: int | None,
compute_metrics: MetricsFn = default_metrics,
max_steps: int | None = None,
) -> dict[str, torch.Tensor]:
"""Train `model` with a plain torch loop, validating and checkpointing on the configured cadence.

Expand All @@ -121,8 +122,13 @@ def run_training_loop( # noqa: C901
: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 max_steps: The maximum number of training steps. When it is reached, the model is validated one last
time and training stops, even before `min_epochs`. If None, the number of steps is not limited.
:return: The model's state dict from the validation check with the best `val_metric`.
:raises ValueError: If `max_steps` is smaller than 1.
"""
if max_steps is not None and max_steps < 1:
raise ValueError("max_steps must be at least 1.")
model.to(device)
loss_function.to(device)

Expand Down Expand Up @@ -150,9 +156,11 @@ def run_training_loop( # noqa: C901
global_step = 0
postfix: dict[str, str] = {}
latest_val_loss: float | None = None
last_validated_step = -1

def validate_and_checkpoint() -> bool:
nonlocal best_checkpoint, best_val_metric, latest_val_loss
nonlocal best_checkpoint, best_val_metric, latest_val_loss, last_validated_step
last_validated_step = global_step
metrics = _run_validation(model, loss_function, compute_metrics, val_loader, device)
latest_val_loss = metrics["val_loss"]
current = metrics[val_metric]
Expand Down Expand Up @@ -185,6 +193,11 @@ def validate_and_checkpoint() -> bool:
if should_stop and (min_epochs is None or current_epoch >= min_epochs):
return best_checkpoint

if max_steps is not None and global_step >= max_steps:
if last_validated_step != global_step:
validate_and_checkpoint()
return best_checkpoint

current_epoch += 1

if check_val_every_epoch is not None and current_epoch % check_val_every_epoch == 0:
Expand Down
91 changes: 91 additions & 0 deletions tests/test_trainable.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import logging
from tempfile import TemporaryDirectory
from typing import Any

import numpy as np
import pytest
Expand Down Expand Up @@ -849,3 +850,93 @@ def counting_step(self: torch.optim.lr_scheduler.ReduceLROnPlateau, metrics: flo
check_val_every_epoch=None,
)
assert len(step_calls) == 3


class _CountingLoss(nn.Module):
def __init__(self, model: nn.Module) -> None:
super().__init__()
self.model = model
self.train_calls = 0
self.val_calls = 0
self.mse = nn.MSELoss()

def forward(self, head_out: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
if self.model.training:
self.train_calls += 1
else:
self.val_calls += 1
return self.mse(head_out, y)


def _run_counting_loop(
max_steps: int | None, val_check_interval: int | None = None, min_epochs: int | None = None
) -> _CountingLoss:
model = nn.Linear(3, 2)
loss = _CountingLoss(model)
run_training_loop(
model=model,
loss_function=loss,
learning_rate=1e-3,
val_metric="val_loss",
early_stopping_direction="min",
train_loader=_make_loader(10),
val_loader=_make_loader(2),
early_stopping_patience=None,
min_epochs=min_epochs,
max_epochs=None,
device=resolve_device("cpu"),
val_check_interval=val_check_interval,
check_val_every_epoch=None if val_check_interval else 1,
max_steps=max_steps,
)
return loss


def test_run_training_loop_stops_at_max_steps() -> None:
"""Training stops after max_steps, across epochs, and validates once at the end."""
loss = _run_counting_loop(max_steps=13)
assert loss.train_calls == 13
assert loss.val_calls == 2 * 2
Comment on lines +895 to +899

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Best checkpoint remains untested The new test checks training and validation call counts, but it never checks the returned checkpoint. It would pass even if training validated at max_steps and then returned the final weights instead of an earlier, better checkpoint. A test with controlled validation results would protect that behavior.

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!



def test_run_training_loop_max_steps_does_not_validate_twice() -> None:
"""If the last step is also a validation step, the model is validated once."""
loss = _run_counting_loop(max_steps=6, val_check_interval=3)
assert loss.train_calls == 6
assert loss.val_calls == 2 * 2


def test_run_training_loop_max_steps_overrides_min_epochs() -> None:
"""max_steps stops training even before min_epochs is reached."""
loss = _run_counting_loop(max_steps=4, min_epochs=5)
assert loss.train_calls == 4


def test_run_training_loop_rejects_invalid_max_steps() -> None:
"""max_steps must be at least 1."""
with pytest.raises(ValueError):
_run_counting_loop(max_steps=0)


@pytest.mark.parametrize(
"model_class, y",
[
(StaticModelForClassification, ["a", "b"] * 4),
(StaticModelForRegression, torch.ones(8, 2)),
(StaticModelForPairSimilarity, ["word2", "word3"] * 4),
],
)
def test_fit_passes_max_steps(
monkeypatch: pytest.MonkeyPatch, mock_vectors: np.ndarray, mock_tokenizer: Tokenizer, model_class: Any, y: Any
) -> None:
"""max_steps is passed from fit to the training loop."""
captured: list[object] = []

def fake_run_training_loop(**kwargs: Any) -> dict[str, torch.Tensor]:
captured.append(kwargs["max_steps"])
return kwargs["model"].state_dict()

monkeypatch.setattr("model2vec.train.base.run_training_loop", fake_run_training_loop)
model = model_class(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer)
model.fit(["word1", "word2", "word3", "word1 word2"] * 2, y, max_steps=7)
assert captured == [7]
Loading