Compare commits

...
Author SHA1 Message Date
Dmitrii Ogn a32500fa6d Merge branch 'main' into mean_pooling 2025-01-30 15:59:50 +03:00
d.rudenko e41197dc0b Fix 2025-01-15 12:43:37 +01:00
d.rudenko c6facad6e8 Fix 2025-01-15 12:24:16 +01:00
d.rudenko 66839cf36a Fix 2025-01-15 12:23:24 +01:00
d.rudenko 0c0581d04a Fix 2025-01-15 12:22:36 +01:00
d.rudenko 276de5a49e Added pooling + tests 2025-01-15 12:18:40 +01:00
d.rudenko 16c41dff18 HF sources for all models 2024-12-27 12:39:51 +01:00
2 changed files with 72 additions and 0 deletions
+38
View File
@@ -0,0 +1,38 @@
import numpy as np
from numpy import ufunc
class LateInteractionPooler(object):
def __init__(self, agg_type="row", agg="mean"):
self.agg = agg
self.type = agg_type
def _pick_operation(self) -> ufunc:
if self.agg == "mean":
return np.mean
elif self.agg == "max":
return np.max
elif self.agg == "min":
return np.min
else:
raise NotImplementedError(
f"LateInteractionPooler only supports agg=mean,min,max, provided {self.agg}"
)
def pool(self, embeddings_batch) -> np.array:
if isinstance(embeddings_batch, np.ndarray) and len(embeddings_batch.shape) == 2:
embeddings_batch = [embeddings_batch]
if self.type == "row":
pooled_embedding = self.pool_row(embeddings_batch)
elif self.type == "col":
pooled_embedding = self.pool_col(embeddings_batch)
else:
raise ValueError("type must be 'row' or 'col'")
return pooled_embedding
def pool_row(self, embeddings_batch) -> np.array:
return self._pick_operation()(embeddings_batch, axis=-1)
def pool_col(self, embeddings_batch) -> np.array:
return self._pick_operation()(embeddings_batch, axis=-2)
+34
View File
@@ -0,0 +1,34 @@
import os
import numpy as np
from fastembed.late_interaction.late_interaction_text_embedding import (
LateInteractionTextEmbedding,
)
from fastembed.common.pooling import LateInteractionPooler
from tests.utils import delete_model_cache
CANONICAL_COLUMN_VALUES = {
"colbert-ir/colbertv2.0": np.array(
[4.0727495e-03, -2.4026826e-03, -6.8204990e-04, -7.1383954e-05, 4.4963313e-03]
),
}
docs = ["Hello World"]
def test_batch_embedding():
is_ci = os.getenv("CI")
docs_to_embed = docs * 10
for model_name, expected_result in CANONICAL_COLUMN_VALUES.items():
print("evaluating", model_name)
model = LateInteractionTextEmbedding(model_name=model_name)
pooler = LateInteractionPooler()
result = list(model.embed(docs_to_embed, batch_size=6))
pooled_result = pooler.pool(result)
assert np.allclose(pooled_result[0], expected_result, atol=2e-3)
if is_ci:
delete_model_cache(model.model._model_dir)