diff --git a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py index 679f289..63f8fb6 100644 --- a/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +++ b/fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py @@ -38,6 +38,23 @@ supported_onnx_models: list[BaseModelDescription] = [ sources=ModelSource(hf="BAAI/bge-reranker-base"), model_file="onnx/model.onnx", ), + BaseModelDescription( + model="BAAI/bge-reranker-v2-m3", + description="Multilingual BGE reranker based on BGE-M3.", + license="apache-2.0", + size_in_GB=2.27, + sources=ModelSource(hf="onnx-community/bge-reranker-v2-m3-ONNX"), + model_file="onnx/model.onnx", + additional_files=["onnx/model.onnx_data"], + ), + BaseModelDescription( + model="BAAI/bge-reranker-v2-m3-int8", + description="Multilingual BGE reranker based on BGE-M3 (INT8 ONNX).", + license="apache-2.0", + size_in_GB=0.57, + sources=ModelSource(hf="onnx-community/bge-reranker-v2-m3-ONNX"), + model_file="onnx/model_int8.onnx", + ), BaseModelDescription( model="jinaai/jina-reranker-v1-tiny-en", description="Designed for blazing-fast re-ranking with 8K context length and fewer parameters than jina-reranker-v1-turbo-en.", diff --git a/tests/test_text_cross_encoder.py b/tests/test_text_cross_encoder.py index 4d0d5b7..00499e1 100644 --- a/tests/test_text_cross_encoder.py +++ b/tests/test_text_cross_encoder.py @@ -11,6 +11,11 @@ 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]), + "BAAI/bge-reranker-v2-m3": np.array([8.78182220, -5.46485329]), + "BAAI/bge-reranker-v2-m3-int8": ( + np.array([8.59202957, -5.39362288]), # ONNX Runtime 1.27.0 + np.array([8.51154804, -5.56898642]), # ONNX Runtime 1.30.0 + ), "jinaai/jina-reranker-v1-tiny-en": np.array([2.5911, 0.1122]), "jinaai/jina-reranker-v1-turbo-en": np.array([1.8295, -2.8908]), "jinaai/jina-reranker-v2-base-multilingual": np.array([1.6533, -1.6455]), @@ -67,8 +72,11 @@ def test_rerank(model_cache, model_name: str) -> None: ), f"Model: {model_desc.model}, Scores: {scores}, Scores2: {scores2}" canonical_scores = CANONICAL_SCORE_VALUES[model_desc.model] - assert np.allclose( - scores, canonical_scores, atol=1e-3 + canonical_variants = ( + canonical_scores if isinstance(canonical_scores, tuple) else (canonical_scores,) + ) + assert any( + np.allclose(scores, expected, atol=1e-3) for expected in canonical_variants ), f"Model: {model_desc.model}, Scores: {scores}, Expected: {canonical_scores}"