From 6dffecf61bf7f09bf33f782662bdaab377c219f1 Mon Sep 17 00:00:00 2001 From: stephantul Date: Sun, 27 Sep 2026 19:15:42 +0200 Subject: [PATCH] feat(train): add max_steps to all trainers fit now accepts max_steps, which stops training after that many steps, also in the middle of an epoch. When it is reached, the model is validated one last time, unless that step was already a validation step, and the best checkpoint is kept. max_steps takes precedence over min_epochs. --- model2vec/train/base.py | 2 + model2vec/train/classifier.py | 4 ++ model2vec/train/pairs.py | 4 ++ model2vec/train/similarity.py | 4 ++ model2vec/train/trainer.py | 15 +++++- tests/test_trainable.py | 91 +++++++++++++++++++++++++++++++++++ 6 files changed, 119 insertions(+), 1 deletion(-) diff --git a/model2vec/train/base.py b/model2vec/train/base.py index 982176e..5c4705e 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -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).") @@ -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) diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index cb23568..e578050 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -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. @@ -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. """ @@ -224,6 +227,7 @@ def fit( validation_steps=validation_steps, compute_metrics=compute_metrics, token_dropout=token_dropout, + max_steps=max_steps, ) return self diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index 07f9254..a92c143 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -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. @@ -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) @@ -218,6 +221,7 @@ def fit( max_epochs=max_epochs, device=device, validation_steps=validation_steps, + max_steps=max_steps, ) return self diff --git a/model2vec/train/similarity.py b/model2vec/train/similarity.py index 367f075..076f5d0 100644 --- a/model2vec/train/similarity.py +++ b/model2vec/train/similarity.py @@ -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. @@ -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) @@ -131,6 +134,7 @@ def fit( device=device, validation_steps=validation_steps, token_dropout=token_dropout, + max_steps=max_steps, ) return self diff --git a/model2vec/train/trainer.py b/model2vec/train/trainer.py index 2c45261..e2d1118 100644 --- a/model2vec/train/trainer.py +++ b/model2vec/train/trainer.py @@ -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. @@ -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) @@ -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] @@ -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: diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 7d90fb0..551576b 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -1,5 +1,6 @@ import logging from tempfile import TemporaryDirectory +from typing import Any import numpy as np import pytest @@ -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 + + +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]