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
9 changes: 8 additions & 1 deletion model2vec/train/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.

Expand All @@ -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.
Expand All @@ -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)

Expand Down
2 changes: 1 addition & 1 deletion model2vec/train/classifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
10 changes: 8 additions & 2 deletions model2vec/train/pairs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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.
Expand All @@ -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(
Expand Down Expand Up @@ -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:
Expand Down
2 changes: 1 addition & 1 deletion model2vec/train/similarity.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
4 changes: 3 additions & 1 deletion model2vec/train/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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:
Expand Down
37 changes: 37 additions & 0 deletions tests/test_trainable.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"]
Expand Down
Loading