Skip to content

Commit 088746e

Browse files
authored
fix: norm underflow for float16 (#369)
1 parent d9ac0d9 commit 088746e

2 files changed

Lines changed: 15 additions & 2 deletions

File tree

model2vec/model.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -476,8 +476,9 @@ def _encode_batch(self, sentences: Sequence[str], normalize: bool) -> np.ndarray
476476
out[i] = emb.mean(axis=0)
477477

478478
if normalize:
479-
norm = np.linalg.norm(out, axis=1, keepdims=True) + 1e-32
480-
np.divide(out, norm, out=out)
479+
out32 = out.astype(np.float32)
480+
norm = np.linalg.norm(out32, axis=1, keepdims=True) + 1e-32
481+
return (out32 / norm).astype(out.dtype)
481482

482483
return out
483484

tests/test_model.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,18 @@ def test_encode_single_sentence_empty(
5353
assert np.all(encoded == 0)
5454

5555

56+
def test_encode_single_sentence_empty_float16(
57+
mock_vectors: np.ndarray, mock_tokenizer: Tokenizer, mock_config: dict[str, str]
58+
) -> None:
59+
"""Test encoding of a single empty sentence with float16 embeddings."""
60+
model = StaticModel(vectors=mock_vectors.astype(np.float16), tokenizer=mock_tokenizer, config=mock_config)
61+
model.normalize = True
62+
encoded = model.encode("")
63+
assert not np.isnan(encoded).any()
64+
assert np.all(encoded == 0)
65+
assert encoded.dtype == model.embedding.dtype
66+
67+
5668
def test_encode_multiple_sentences(
5769
mock_vectors: np.ndarray, mock_tokenizer: Tokenizer, mock_config: dict[str, str]
5870
) -> None:

0 commit comments

Comments
 (0)