57 lines
2.1 KiB
Python
57 lines
2.1 KiB
Python
from haiku.rag.config import AppConfig, get_config
|
|
from haiku.rag.reranking.base import RerankerBase
|
|
from haiku.rag.utils import check_api_key_supported
|
|
|
|
|
|
def get_reranker(config: AppConfig | None = None) -> RerankerBase | None:
|
|
"""Build the configured reranker, or None if reranking is disabled.
|
|
|
|
A configured reranker whose optional dependency is missing raises: the
|
|
provider modules import their dependency at module scope and name the
|
|
extra to install.
|
|
"""
|
|
config = config if config is not None else get_config()
|
|
model = config.reranking.model
|
|
if model is None:
|
|
return None
|
|
|
|
check_api_key_supported(model, {"vllm"})
|
|
|
|
if config.reranking.multimodal and model.provider != "vllm":
|
|
raise ValueError("reranking.multimodal is only supported on the vllm provider")
|
|
|
|
if model.provider == "cohere":
|
|
from haiku.rag.reranking.cohere import CohereReranker
|
|
|
|
return CohereReranker(model.name)
|
|
|
|
if model.provider == "vllm":
|
|
if not model.base_url:
|
|
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
|
from haiku.rag.reranking.vllm import VLLMReranker
|
|
|
|
return VLLMReranker(model.name, model.base_url, api_key=model.api_key)
|
|
|
|
if model.provider == "zeroentropy":
|
|
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
|
|
|
return ZeroEntropyReranker(model.name or "zerank-1")
|
|
|
|
if model.provider == "jina":
|
|
from haiku.rag.reranking.jina import JinaReranker
|
|
|
|
return JinaReranker(model.name or "jina-reranker-v3")
|
|
|
|
if model.provider == "jina-local":
|
|
from haiku.rag.reranking.jina_local import JinaLocalReranker
|
|
|
|
return JinaLocalReranker(model.name or "jinaai/jina-reranker-v3")
|
|
|
|
if model.provider == "cross-encoder":
|
|
if not model.name:
|
|
raise ValueError("cross-encoder reranker requires name in reranking.model")
|
|
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
|
|
|
|
return CrossEncoderReranker(model.name)
|
|
|
|
raise ValueError(f"Unknown reranking provider: {model.provider}")
|