diff --git a/model2vec/model.py b/model2vec/model.py index a55dfdb..4e56589 100644 --- a/model2vec/model.py +++ b/model2vec/model.py @@ -80,11 +80,7 @@ def __init__( self.token_mapping: np.ndarray | None = token_mapping self.tokenizer = copy.deepcopy(tokenizer) - padding = self.tokenizer.padding - if padding is not None: - self.tokenizer.enable_padding( - pad_id=padding["pad_id"], pad_token=padding["pad_token"], pad_type_id=padding["pad_type_id"], length=0 - ) + _disable_padding(self.tokenizer) self.unk_token_id = _get_unk_token_id(self.tokenizer) self.median_token_length = int(np.median([len(token) for token in self.tokens])) @@ -615,6 +611,15 @@ def _loading_helper( ) +def _disable_padding(tokenizer: Tokenizer) -> None: + """Stop the tokenizer from padding, while keeping its pad token.""" + padding = tokenizer.padding + if padding is not None: + tokenizer.enable_padding( + pad_id=padding["pad_id"], pad_token=padding["pad_token"], pad_type_id=padding["pad_type_id"], length=0 + ) + + def _get_unk_token_id(tokenizer: Tokenizer) -> int | None: """Get the unk token id.""" model = tokenizer.model diff --git a/model2vec/train/base.py b/model2vec/train/base.py index 04fff14..2d7aa8b 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -13,7 +13,7 @@ from tqdm import trange from model2vec.inference import StaticModelPipeline -from model2vec.model import DEFAULT_MAX_LENGTH, PathLike, StaticModel, _get_unk_token_id +from model2vec.model import DEFAULT_MAX_LENGTH, PathLike, StaticModel, _disable_padding, _get_unk_token_id from model2vec.train.dataset import PairDataset, TextDataset from model2vec.train.trainer import MetricsFn, default_metrics, resolve_device, run_training_loop from model2vec.train.utils import ( @@ -93,6 +93,7 @@ def __init__( # Truncation happens here through `max_length`; a StaticModel's tokenizer carries its own setting. self.tokenizer = copy.deepcopy(tokenizer) self.tokenizer.no_truncation() + _disable_padding(self.tokenizer) self.unk_token_id = _get_unk_token_id(self.tokenizer) def _remove_unk(self, token_ids: list[int]) -> list[int]: diff --git a/tests/test_trainable.py b/tests/test_trainable.py index e039e67..3ada3d4 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -53,6 +53,16 @@ def test_init_base_class(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> assert head[0].in_features == mock_vectors.shape[1] +def test_trainable_tokenizer_does_not_pad(mock_trained_pair_similarity_pipeline: StaticModelForPairSimilarity) -> None: + """The tokenizer of a trainable model keeps its pad token, but doesn't pad.""" + model = mock_trained_pair_similarity_pipeline + assert model.tokenizer.padding is not None + assert ( + model._tokenize_texts(["word1 word2", "word2"], max_length=None)[1] + == model._tokenize_texts(["word2"], max_length=None)[0] + ) + + def test_empty_texts_have_finite_gradients(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """Texts without any tokens encode to zero vectors and don't produce NaN gradients.""" torch.manual_seed(0)