Properly test cohere reranker
This commit is contained in:
parent
044e1b0a06
commit
f25416b0ff
2 changed files with 21 additions and 76 deletions
|
|
@ -19,8 +19,6 @@ class AppConfig(BaseModel):
|
||||||
EMBEDDINGS_MODEL: str = "mxbai-embed-large"
|
EMBEDDINGS_MODEL: str = "mxbai-embed-large"
|
||||||
EMBEDDINGS_VECTOR_DIM: int = 1024
|
EMBEDDINGS_VECTOR_DIM: int = 1024
|
||||||
|
|
||||||
# RERANK_PROVIDER: str = "cohere"
|
|
||||||
# RERANK_MODEL: str = "rerank-v3.5"
|
|
||||||
RERANK_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2"
|
RERANK_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
RERANK_PROVIDER: str = "mxbai"
|
RERANK_PROVIDER: str = "mxbai"
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,23 +1,9 @@
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.reranking import get_reranker
|
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_reranker_base():
|
|
||||||
reranker = RerankerBase()
|
|
||||||
assert reranker._model == "mixedbread-ai/mxbai-rerank-base-v2"
|
|
||||||
|
|
||||||
with pytest.raises(NotImplementedError):
|
|
||||||
await reranker.rerank("query", [])
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_mxbai_reranker():
|
|
||||||
reranker = MxBAIReranker()
|
|
||||||
chunks = [
|
chunks = [
|
||||||
Chunk(content=content, document_id=i)
|
Chunk(content=content, document_id=i)
|
||||||
for i, content in enumerate(
|
for i, content in enumerate(
|
||||||
|
|
@ -32,6 +18,19 @@ async def test_mxbai_reranker():
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_reranker_base():
|
||||||
|
reranker = RerankerBase()
|
||||||
|
assert reranker._model == "mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
|
|
||||||
|
with pytest.raises(NotImplementedError):
|
||||||
|
await reranker.rerank("query", [])
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mxbai_reranker():
|
||||||
|
reranker = MxBAIReranker()
|
||||||
reranked = await reranker.rerank(
|
reranked = await reranker.rerank(
|
||||||
"Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2
|
"Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2
|
||||||
)
|
)
|
||||||
|
|
@ -40,68 +39,16 @@ async def test_mxbai_reranker():
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cohere_reranker():
|
async def test_cohere_reranker():
|
||||||
try:
|
|
||||||
# Mock the client
|
|
||||||
class MockResult:
|
|
||||||
def __init__(self, index):
|
|
||||||
self.index = index
|
|
||||||
|
|
||||||
class MockResponse:
|
|
||||||
def __init__(self, results):
|
|
||||||
self.results = results
|
|
||||||
|
|
||||||
class MockClient:
|
|
||||||
def __init__(self, api_key=None):
|
|
||||||
pass
|
|
||||||
|
|
||||||
def rerank(self, model, query, documents, top_n):
|
|
||||||
return MockResponse([MockResult(1), MockResult(0)])
|
|
||||||
|
|
||||||
import haiku.rag.reranking.cohere
|
|
||||||
|
|
||||||
original_client = haiku.rag.reranking.cohere.cohere.ClientV2
|
|
||||||
haiku.rag.reranking.cohere.cohere.ClientV2 = MockClient
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.cohere import CohereReranker
|
from haiku.rag.reranking.cohere import CohereReranker
|
||||||
|
|
||||||
reranker = CohereReranker()
|
reranker = CohereReranker()
|
||||||
assert reranker._model == "rerank-v3.5"
|
assert reranker._model == "rerank-v3.5"
|
||||||
|
|
||||||
chunks = [
|
reranked = await reranker.rerank(
|
||||||
Chunk(id=1, content="First chunk", document_id=1),
|
"Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2
|
||||||
Chunk(id=2, content="Second chunk", document_id=1),
|
)
|
||||||
]
|
assert [r.document_id for r in reranked] == [0, 2]
|
||||||
|
|
||||||
result = await reranker.rerank("test query", chunks)
|
|
||||||
assert len(result) == 2
|
|
||||||
assert result[0] == chunks[1] # Should return chunk at index 1 first
|
|
||||||
assert result[1] == chunks[0] # Should return chunk at index 0 second
|
|
||||||
finally:
|
|
||||||
haiku.rag.reranking.cohere.cohere.ClientV2 = original_client
|
|
||||||
|
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pytest.skip("Cohere package not installed")
|
pytest.skip("Cohere package not installed")
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_get_reranker():
|
|
||||||
try:
|
|
||||||
|
|
||||||
class MockClient:
|
|
||||||
def __init__(self, api_key=None):
|
|
||||||
pass
|
|
||||||
|
|
||||||
import haiku.rag.reranking.cohere
|
|
||||||
|
|
||||||
original_client = haiku.rag.reranking.cohere.cohere.ClientV2
|
|
||||||
haiku.rag.reranking.cohere.cohere.ClientV2 = MockClient
|
|
||||||
|
|
||||||
try:
|
|
||||||
reranker = get_reranker()
|
|
||||||
assert reranker._model == "rerank-v3.5"
|
|
||||||
assert hasattr(reranker, "rerank")
|
|
||||||
finally:
|
|
||||||
haiku.rag.reranking.cohere.cohere.ClientV2 = original_client
|
|
||||||
except ImportError:
|
|
||||||
pytest.skip("Cohere package not installed")
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue