From 2f3f613c4e2ad5bf906fd481965ff4c21c98ec15 Mon Sep 17 00:00:00 2001 From: stephantul Date: Fri, 9 Oct 2026 09:01:04 +0200 Subject: [PATCH] fix: random seed controls splitting --- model2vec/train/base.py | 9 ++++++++- model2vec/train/classifier.py | 2 +- model2vec/train/pairs.py | 10 ++++++++-- model2vec/train/similarity.py | 2 +- model2vec/train/utils.py | 4 +++- tests/test_trainable.py | 37 +++++++++++++++++++++++++++++++++++ 6 files changed, 58 insertions(+), 6 deletions(-) diff --git a/model2vec/train/base.py b/model2vec/train/base.py index 06f75e4..4615545 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -25,6 +25,7 @@ ) from model2vec.train.trainer import MetricsFn, default_metrics, resolve_device, run_training_loop from model2vec.train.utils import ( + DEFAULT_RANDOM_SEED, MAX_VALIDATION_SIZE, split_indices, to_pipeline, @@ -524,6 +525,7 @@ def _create_datasets( y_val: Any | None, test_size: float | int, stratify_by: Sequence[Any] | None = None, + random_seed: int = DEFAULT_RANDOM_SEED, ) -> tuple[TextDataset, TextDataset]: """Create the training and validation datasets. @@ -535,6 +537,7 @@ def _create_datasets( `MAX_VALIDATION_SIZE` rows, or a number of rows if it is an int. :param stratify_by: Validated single labels to stratify the validation split by. If None, the split is not stratified. + :param random_seed: The random seed of the validation split. :return: The train and validation datasets. :raises ValueError: If only one of `X_val` and `y_val` is given, or if the texts and labels have different lengths. @@ -549,7 +552,11 @@ def _create_datasets( return self._text_dataset(rows), self._text_dataset(ColumnRows(text=X_val, label=y_val)) train_indices, val_indices = split_indices( - len(rows), test_size, max_test_size=MAX_VALIDATION_SIZE, stratify_by=stratify_by + len(rows), + test_size, + max_test_size=MAX_VALIDATION_SIZE, + stratify_by=stratify_by, + random_seed=random_seed, ) return self._text_dataset(rows, train_indices), self._text_dataset(rows, val_indices) diff --git a/model2vec/train/classifier.py b/model2vec/train/classifier.py index 322fa09..ccc33a2 100644 --- a/model2vec/train/classifier.py +++ b/model2vec/train/classifier.py @@ -392,7 +392,7 @@ def fit( self._initialize() resolved_class_weight = self._resolve_class_weight(class_weight, label_counts) train_dataset, val_dataset = self._create_datasets( - X, y, X_val, y_val, test_size, stratify_by=None if self.multilabel else y + X, y, X_val, y_val, test_size, stratify_by=None if self.multilabel else y, random_seed=random_seed ) if self.multilabel: diff --git a/model2vec/train/pairs.py b/model2vec/train/pairs.py index 0f910bd..43d06c9 100644 --- a/model2vec/train/pairs.py +++ b/model2vec/train/pairs.py @@ -237,6 +237,7 @@ def _create_pair_datasets( text_a_val: Sequence[str] | None, text_b_val: Sequence[str] | None, test_size: float | int, + random_seed: int = DEFAULT_RANDOM_SEED, ) -> tuple[PairDataset, PairDataset]: """Create the training and validation datasets of pairs. @@ -247,6 +248,7 @@ def _create_pair_datasets( :param text_b_val: The second half of each validation pair. :param test_size: The size of the validation split if `text_a_val` is None: a fraction of the pairs, capped at `MAX_VALIDATION_SIZE` rows, or a number of pairs if it is an int. + :param random_seed: The random seed of the validation split. :return: The train and validation datasets. :raises ValueError: If only one of `text_a_val` and `text_b_val` is given, or if the halves of the pairs have different lengths. @@ -260,7 +262,9 @@ def _create_pair_datasets( self._check_aligned(text_a_val=text_a_val, text_b_val=text_b_val) return self._pair_dataset(rows), self._pair_dataset(ColumnRows(text_a=text_a_val, text_b=text_b_val)) - train_indices, val_indices = split_indices(len(rows), test_size, max_test_size=MAX_VALIDATION_SIZE) + train_indices, val_indices = split_indices( + len(rows), test_size, max_test_size=MAX_VALIDATION_SIZE, random_seed=random_seed + ) return self._pair_dataset(rows, train_indices), self._pair_dataset(rows, val_indices) def fit( @@ -324,7 +328,9 @@ def fit( self._check_inputs(text_a=text_a, text_b=text_b, text_a_val=text_a_val, text_b_val=text_b_val) loss_function = PairInfoNCELoss(temperature=temperature) - train_dataset, val_dataset = self._create_pair_datasets(text_a, text_b, text_a_val, text_b_val, test_size) + train_dataset, val_dataset = self._create_pair_datasets( + text_a, text_b, text_a_val, text_b_val, test_size, random_seed=random_seed + ) self._check_pair_splits(len(train_dataset), len(val_dataset)) batch_size = self._determine_batch_size(batch_size, len(train_dataset)) if batch_size < 2: diff --git a/model2vec/train/similarity.py b/model2vec/train/similarity.py index 86568a0..3921294 100644 --- a/model2vec/train/similarity.py +++ b/model2vec/train/similarity.py @@ -242,7 +242,7 @@ def fit( if y_val is not None and (val_dim := _vector_dim(y_val, "y_val")) != out_dim: raise ValueError(f"The vectors in y_val have dimension {val_dim}, but those in y have dimension {out_dim}.") - train_dataset, val_dataset = self._create_datasets(X, y, X_val, y_val, test_size) + train_dataset, val_dataset = self._create_datasets(X, y, X_val, y_val, test_size, random_seed=random_seed) self.out_dim = out_dim self._initialize() self._train( diff --git a/model2vec/train/utils.py b/model2vec/train/utils.py index 562112e..3226ede 100644 --- a/model2vec/train/utils.py +++ b/model2vec/train/utils.py @@ -76,6 +76,7 @@ def split_indices( test_size: float | int, max_test_size: int | None = None, stratify_by: Sequence[Any] | None = None, + random_seed: int = DEFAULT_RANDOM_SEED, ) -> tuple[np.ndarray, np.ndarray]: """Randomly split the indices `0..n-1` into sorted train and test indices. @@ -88,12 +89,13 @@ def split_indices( :param stratify_by: The single label of each item, as strings or integers that have been validated. If every label occurs at least twice, each label is split separately, in the same proportion. If None, the split is not stratified. + :param random_seed: The random seed of the split. :return: The train indices and the test indices. :raises ValueError: If `test_size` is a bool. """ if isinstance(test_size, bool): raise ValueError("test_size must be a float or an int, not a bool.") - rng = np.random.default_rng(DEFAULT_RANDOM_SEED) + rng = np.random.default_rng(random_seed) if isinstance(test_size, numbers.Integral): n_test = int(test_size) else: diff --git a/tests/test_trainable.py b/tests/test_trainable.py index d8afdc2..060de7c 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -1208,6 +1208,17 @@ def test_split_indices() -> None: assert list(test) == sorted(test) +def test_split_indices_random_seed() -> None: + """The same seed gives the same split, and a different seed a different one.""" + labels = ["a"] * 50 + ["b"] * 50 + for stratify_by in (None, labels): + first = split_indices(100, 0.2, stratify_by=stratify_by, random_seed=1) + second = split_indices(100, 0.2, stratify_by=stratify_by, random_seed=1) + other = split_indices(100, 0.2, stratify_by=stratify_by, random_seed=2) + assert np.array_equal(first[1], second[1]) + assert not np.array_equal(first[1], other[1]) + + def test_split_indices_absolute_and_capped_sizes() -> None: """An int test size is a number of items, and a fractional test size can be capped.""" assert len(split_indices(100, 7)[1]) == 7 @@ -1314,6 +1325,32 @@ def test_fit_only_stratifies_single_labels( assert (stratify_by is y) if stratified else (stratify_by is None) +@pytest.mark.parametrize( + ("model_class", "labels", "module"), + [ + (StaticModelForClassification, ["a", "b"] * 4, "base"), + (StaticModelForSimilarity, [[0.5, 1.0]] * 8, "base"), + (StaticModelForRegression, [[0.5, 1.0]] * 8, "base"), + (StaticModelForPairSimilarity, None, "pairs"), + ], + ids=["classification", "similarity", "regression", "pairs"], +) +def test_fit_seeds_validation_split( + model_class: type[BaseFinetuneable], + labels: list[Any] | None, + module: str, + mock_vectors: np.ndarray, + mock_tokenizer: Tokenizer, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The validation split uses the random seed passed to fit.""" + monkeypatch.setattr("model2vec.train.base.run_training_loop", lambda **kwargs: kwargs["model"].state_dict()) + model = model_class(vectors=torch.from_numpy(mock_vectors).float(), tokenizer=mock_tokenizer) + with patch(f"model2vec.train.{module}.split_indices", wraps=split_indices) as split_mock: + model.fit(_TRAIN_TEXTS, labels or _TRAIN_TEXTS, test_size=0.5, random_seed=7) # type: ignore[attr-defined] + assert split_mock.call_args.kwargs["random_seed"] == 7 + + def test_column_strata_match_list_strata_across_batches() -> None: """Columns are grouped by label in batches, matching lists, in order of first occurrence.""" labels = ["b", "a", "c", "a", "b", "c", "b", "a", "d", "d"]