mirror of
https://github.com/qdrant/fastembed.git
synced 2026-07-23 11:20:51 -05:00
Fix spladepp parallelism (#169)
* fix: add get_worker_class implementation to spladepp * fix: add tests for parallel embed for spladepp
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
|
||||
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union, Type
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -114,6 +114,10 @@ class SpladePP(SparseTextEmbeddingBase, OnnxModel[SparseEmbedding]):
|
||||
parallel=parallel,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_worker_class(cls) -> Type[EmbeddingWorker]:
|
||||
return SpladePPEmbeddingWorker
|
||||
|
||||
|
||||
class SpladePPEmbeddingWorker(EmbeddingWorker):
|
||||
def init_embedding(
|
||||
|
||||
@@ -71,3 +71,28 @@ def test_single_embedding():
|
||||
|
||||
for i, value in enumerate(result.values):
|
||||
assert pytest.approx(value, abs=0.001) == expected_result["values"][i]
|
||||
|
||||
|
||||
def test_parallel_processing():
|
||||
import numpy as np
|
||||
|
||||
model = SparseTextEmbedding(
|
||||
model_name="prithivida/Splade_PP_en_v1",
|
||||
)
|
||||
docs = ["hello world", "flag embedding"] * 30
|
||||
sparse_embeddings_duo = list(model.embed(docs, batch_size=10, parallel=2))
|
||||
sparse_embeddings_all = list(model.embed(docs, batch_size=10, parallel=0))
|
||||
sparse_embeddings = list(model.embed(docs, batch_size=10, parallel=None))
|
||||
|
||||
assert len(sparse_embeddings) == len(sparse_embeddings_duo) == len(sparse_embeddings_all) == len(docs)
|
||||
|
||||
for sparse_embedding, sparse_embedding_duo, sparse_embedding_all in zip(
|
||||
sparse_embeddings, sparse_embeddings_duo, sparse_embeddings_all
|
||||
):
|
||||
assert (
|
||||
sparse_embedding.indices.tolist()
|
||||
== sparse_embedding_duo.indices.tolist()
|
||||
== sparse_embedding_all.indices.tolist()
|
||||
)
|
||||
assert np.allclose(sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3)
|
||||
assert np.allclose(sparse_embedding.values, sparse_embedding_all.values, atol=1e-3)
|
||||
|
||||
Reference in New Issue
Block a user