mirror of
https://github.com/qdrant/fastembed.git
synced 2026-10-03 19:37:46 -05:00
* feat: Added multi gpu support for text embedding * feat: Add support for multi-gpu for special text models * fix: Fix lazy_load to load the model to child processes when parallel is not none * feat: Added lazy_load and multi-gpu to colbert * feat: Add lazy_load and multi gpu to image models * feat: Support lazy_load and multi-gpu to sparse models (except BM25) * fix: Fixed BM25 not working * refactor: Remove redundant GPUParallelProcessor * refactor: Refactor _embed_*_parallel * feat: Add cuda argument refactor: Refactor how worker assign device * fix: Fix if providers and cuda are None * fix: Fix providers and cuda are none * WIP: Multi gpu support review (#361) * WIP: review * wip: review * refactor: refactor images * refactor: refactor sparse * refactor: refactor late interaction * add model loading * add tests * fix: uncomment models in tests * fix: fix variable declaration order * fix: fix device id assignment * tests: add multi gpu tests * fix: fix device id assignment for sparse embeddings * tests: update multi gpu tests --------- Co-authored-by: George Panchuk <george.panchuk@qdrant.tech> * refactor: remove redundant declarations * fix: rollback redundant changes * fix: remove num workers device ids dep, fix type hint * fix: fix post process for sparse models * fix: remove redundant model loading * new: add lazy load and new gpu support to cross encoders * fix: add rerankers to multi gpu tests * fix: unlock multilingual test * fix: fix gpu test with cross encoder --------- Co-authored-by: Andrey Vasnetsov <andrey@vasnetsov.com> Co-authored-by: George Panchuk <george.panchuk@qdrant.tech>
72 lines
2.4 KiB
Python
72 lines
2.4 KiB
Python
import os
|
|
|
|
import numpy as np
|
|
import pytest
|
|
import shutil
|
|
|
|
from fastembed.rerank.cross_encoder import TextCrossEncoder
|
|
|
|
CANONICAL_SCORE_VALUES = {
|
|
"Xenova/ms-marco-MiniLM-L-6-v2": np.array([8.500708, -2.541011]),
|
|
"Xenova/ms-marco-MiniLM-L-12-v2": np.array([9.330912, -2.0380247]),
|
|
"BAAI/bge-reranker-base": np.array([6.15733337, -3.65939403]),
|
|
}
|
|
|
|
|
|
def test_rerank():
|
|
is_ci = os.getenv("CI")
|
|
|
|
for model_desc in TextCrossEncoder.list_supported_models():
|
|
if not is_ci and model_desc["size_in_GB"] > 1:
|
|
continue
|
|
|
|
model_name = model_desc["model"]
|
|
model = TextCrossEncoder(model_name=model_name)
|
|
|
|
query = "What is the capital of France?"
|
|
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
|
|
scores = np.array(list(model.rerank(query, documents)))
|
|
|
|
canonical_scores = CANONICAL_SCORE_VALUES[model_name]
|
|
assert np.allclose(
|
|
scores, canonical_scores, atol=1e-3
|
|
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
|
|
if is_ci:
|
|
shutil.rmtree(model.model._model_dir)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"model_name",
|
|
["Xenova/ms-marco-MiniLM-L-6-v2", "Xenova/ms-marco-MiniLM-L-12-v2", "BAAI/bge-reranker-base"],
|
|
)
|
|
def test_batch_rerank(model_name):
|
|
is_ci = os.getenv("CI")
|
|
|
|
model = TextCrossEncoder(model_name=model_name)
|
|
|
|
query = "What is the capital of France?"
|
|
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."] * 50
|
|
scores = np.array(list(model.rerank(query, documents, batch_size=10)))
|
|
|
|
canonical_scores = np.tile(CANONICAL_SCORE_VALUES[model_name], 50)
|
|
|
|
assert scores.shape == canonical_scores.shape, f"Unexpected shape for model {model_name}"
|
|
assert np.allclose(
|
|
scores, canonical_scores, atol=1e-3
|
|
), f"Model: {model_name}, Scores: {scores}, Expected: {canonical_scores}"
|
|
if is_ci:
|
|
shutil.rmtree(model.model._model_dir)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"model_name",
|
|
["Xenova/ms-marco-MiniLM-L-6-v2"],
|
|
)
|
|
def test_lazy_load(model_name):
|
|
model = TextCrossEncoder(model_name=model_name, lazy_load=True)
|
|
assert not hasattr(model.model, "model")
|
|
query = "What is the capital of France?"
|
|
documents = ["Paris is the capital of France.", "Berlin is the capital of Germany."]
|
|
list(model.rerank(query, documents))
|
|
assert hasattr(model.model, "model")
|