From 0338367d1de0d4d198633e87edee98def7a5f686 Mon Sep 17 00:00:00 2001 From: stephantul Date: Sun, 27 Sep 2026 19:10:44 +0200 Subject: [PATCH] fix(train): avoid NaN gradients for texts without tokens Mean pooling divided by the number of tokens plus 1e-16. For a text without any tokens, the backward pass divided the gradient by 1e-16 after normalize had already scaled it up, which overflowed to inf and turned into NaN in the token weight gradient. Clamp the length to at least 1 instead: non-empty texts are unaffected, and empty texts get a zero embedding and a zero gradient. --- model2vec/train/base.py | 3 +-- tests/test_trainable.py | 14 ++++++++++++++ 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/model2vec/train/base.py b/model2vec/train/base.py index 982176e..04fff14 100644 --- a/model2vec/train/base.py +++ b/model2vec/train/base.py @@ -230,8 +230,7 @@ def _encode(self, input_ids: torch.Tensor) -> torch.Tensor: """ zeros = (input_ids != self.pad_id).float() zeros = self._apply_token_dropout(zeros) - # Add a small epsilon to avoid division by zero - length = zeros.sum(1) + 1e-16 + length = zeros.sum(1).clamp(min=1) input_ids_embeddings = self.token_mapping[input_ids] embedded = self.embeddings(input_ids_embeddings) diff --git a/tests/test_trainable.py b/tests/test_trainable.py index 7d90fb0..e039e67 100644 --- a/tests/test_trainable.py +++ b/tests/test_trainable.py @@ -53,6 +53,20 @@ def test_init_base_class(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> assert head[0].in_features == mock_vectors.shape[1] +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) + model = StaticModelForClassification( + vectors=torch.from_numpy(mock_vectors).float() * 1e20, tokenizer=mock_tokenizer, n_layers=0 + ) + dataset = model._prepare_dataset(["word1 word2", ""], torch.tensor([0, 1]), max_length=None) + batch, y = next(iter(dataset.to_dataloader(shuffle=False, batch_size=2))) + + nn.functional.cross_entropy(model(batch), y).backward() + + assert all(torch.isfinite(p.grad).all() for p in model.parameters() if p.grad is not None) + + def test_init_base_from_model(mock_vectors: np.ndarray, mock_tokenizer: Tokenizer) -> None: """Test initializion from a static model.""" model = StaticModel(vectors=mock_vectors, tokenizer=mock_tokenizer)