Make reranker a global module object to avoid re-initialization
This commit is contained in:
parent
2343ab7751
commit
1437942edc
2 changed files with 13 additions and 2 deletions
|
|
@ -19,8 +19,8 @@ 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_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2"
|
|
||||||
RERANK_PROVIDER: str = "mxbai"
|
RERANK_PROVIDER: str = "mxbai"
|
||||||
|
RERANK_MODEL: str = "mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
|
|
||||||
QA_PROVIDER: str = "ollama"
|
QA_PROVIDER: str = "ollama"
|
||||||
QA_MODEL: str = "qwen3"
|
QA_MODEL: str = "qwen3"
|
||||||
|
|
|
||||||
|
|
@ -6,11 +6,21 @@ try:
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
_reranker: RerankerBase | None = None
|
||||||
|
|
||||||
|
|
||||||
def get_reranker() -> RerankerBase:
|
def get_reranker() -> RerankerBase:
|
||||||
"""
|
"""
|
||||||
Factory function to get the appropriate reranker based on the configuration.
|
Factory function to get the appropriate reranker based on the configuration.
|
||||||
"""
|
"""
|
||||||
|
global _reranker
|
||||||
|
if _reranker is not None:
|
||||||
|
return _reranker
|
||||||
|
if Config.RERANK_PROVIDER == "mxbai":
|
||||||
|
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||||
|
|
||||||
|
_reranker = MxBAIReranker()
|
||||||
|
return _reranker
|
||||||
|
|
||||||
if Config.RERANK_PROVIDER == "cohere":
|
if Config.RERANK_PROVIDER == "cohere":
|
||||||
try:
|
try:
|
||||||
|
|
@ -21,6 +31,7 @@ def get_reranker() -> RerankerBase:
|
||||||
"Please install haiku.rag with the 'cohere' extra:"
|
"Please install haiku.rag with the 'cohere' extra:"
|
||||||
"uv pip install haiku.rag --extra cohere"
|
"uv pip install haiku.rag --extra cohere"
|
||||||
)
|
)
|
||||||
return CohereReranker()
|
_reranker = CohereReranker()
|
||||||
|
return _reranker
|
||||||
|
|
||||||
raise ValueError(f"Unsupported reranker provider: {Config.RERANK_PROVIDER}")
|
raise ValueError(f"Unsupported reranker provider: {Config.RERANK_PROVIDER}")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue