mirror of
https://github.com/qdrant/fastembed.git
synced 2026-09-22 14:07:51 -05:00
Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a32500fa6d | ||
|
|
e41197dc0b | ||
|
|
c6facad6e8 | ||
|
|
66839cf36a | ||
|
|
0c0581d04a | ||
|
|
276de5a49e | ||
|
|
16c41dff18 |
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user