Collapse get_reranker into a single guard and import-guarded dispatch
Replace the seven repeated `config.reranking.model and ... == provider` checks and six per-branch ImportError handlers with one None guard and one try/except around the provider dispatch. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
38079ff89a
commit
654cb2b94c
1 changed files with 23 additions and 47 deletions
|
|
@ -5,79 +5,55 @@ from haiku.rag.reranking.base import RerankerBase
|
||||||
|
|
||||||
|
|
||||||
def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
"""
|
"""Build the configured reranker, or None if reranking is disabled or its
|
||||||
Factory function to get the appropriate reranker based on the configuration.
|
optional dependency is not installed."""
|
||||||
Returns None if reranking is disabled.
|
model = config.reranking.model
|
||||||
|
if model is None:
|
||||||
|
return None
|
||||||
|
|
||||||
Args:
|
try:
|
||||||
config: Configuration to use. Defaults to global Config.
|
if model.provider == "mxbai":
|
||||||
|
|
||||||
Returns:
|
|
||||||
A reranker instance if configured, None otherwise.
|
|
||||||
"""
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "mxbai":
|
|
||||||
try:
|
|
||||||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||||
|
|
||||||
os.environ["TOKENIZERS_PARALLELISM"] = "true"
|
os.environ["TOKENIZERS_PARALLELISM"] = "true"
|
||||||
return MxBAIReranker()
|
return MxBAIReranker()
|
||||||
except ImportError: # pragma: no cover
|
|
||||||
return None
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "cohere":
|
if model.provider == "cohere":
|
||||||
try:
|
|
||||||
from haiku.rag.reranking.cohere import CohereReranker
|
from haiku.rag.reranking.cohere import CohereReranker
|
||||||
|
|
||||||
return CohereReranker()
|
return CohereReranker()
|
||||||
except ImportError: # pragma: no cover
|
|
||||||
return None
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "vllm":
|
if model.provider == "vllm":
|
||||||
try:
|
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
|
|
||||||
base_url = config.reranking.model.base_url
|
if not model.base_url:
|
||||||
if not base_url:
|
|
||||||
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
||||||
return VLLMReranker(config.reranking.model.name, base_url)
|
return VLLMReranker(model.name, model.base_url)
|
||||||
except ImportError: # pragma: no cover
|
|
||||||
return None
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "zeroentropy":
|
if model.provider == "zeroentropy":
|
||||||
try:
|
|
||||||
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
||||||
|
|
||||||
model = config.reranking.model.name or "zerank-1"
|
return ZeroEntropyReranker(model.name or "zerank-1")
|
||||||
return ZeroEntropyReranker(model)
|
|
||||||
except ImportError: # pragma: no cover
|
|
||||||
return None
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "jina":
|
if model.provider == "jina":
|
||||||
from haiku.rag.reranking.jina import JinaReranker
|
from haiku.rag.reranking.jina import JinaReranker
|
||||||
|
|
||||||
model = config.reranking.model.name or "jina-reranker-v3"
|
return JinaReranker(model.name or "jina-reranker-v3")
|
||||||
return JinaReranker(model)
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "jina-local":
|
if model.provider == "jina-local":
|
||||||
try:
|
|
||||||
from haiku.rag.reranking.jina_local import JinaLocalReranker
|
from haiku.rag.reranking.jina_local import JinaLocalReranker
|
||||||
|
|
||||||
model = config.reranking.model.name or "jinaai/jina-reranker-v3"
|
return JinaLocalReranker(model.name or "jinaai/jina-reranker-v3")
|
||||||
return JinaLocalReranker(model)
|
|
||||||
except ImportError: # pragma: no cover
|
|
||||||
return None
|
|
||||||
|
|
||||||
if config.reranking.model and config.reranking.model.provider == "cross-encoder":
|
if model.provider == "cross-encoder":
|
||||||
try:
|
|
||||||
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
|
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
|
||||||
|
|
||||||
name = config.reranking.model.name
|
if not model.name:
|
||||||
if not name:
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"cross-encoder reranker requires name in reranking.model"
|
"cross-encoder reranker requires name in reranking.model"
|
||||||
)
|
)
|
||||||
return CrossEncoderReranker(name)
|
return CrossEncoderReranker(model.name)
|
||||||
except ImportError: # pragma: no cover
|
except ImportError: # pragma: no cover
|
||||||
return None
|
return None
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue