Validate reranker config before importing the provider package
This commit is contained in:
parent
ac64aa9c17
commit
7f50d3698f
1 changed files with 4 additions and 4 deletions
|
|
@ -24,10 +24,10 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
return CohereReranker()
|
return CohereReranker()
|
||||||
|
|
||||||
if model.provider == "vllm":
|
if model.provider == "vllm":
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
|
||||||
|
|
||||||
if not model.base_url:
|
if not model.base_url:
|
||||||
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
||||||
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
|
|
||||||
return VLLMReranker(model.name, model.base_url)
|
return VLLMReranker(model.name, model.base_url)
|
||||||
|
|
||||||
if model.provider == "zeroentropy":
|
if model.provider == "zeroentropy":
|
||||||
|
|
@ -46,12 +46,12 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
return JinaLocalReranker(model.name or "jinaai/jina-reranker-v3")
|
return JinaLocalReranker(model.name or "jinaai/jina-reranker-v3")
|
||||||
|
|
||||||
if model.provider == "cross-encoder":
|
if model.provider == "cross-encoder":
|
||||||
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
|
|
||||||
|
|
||||||
if not model.name:
|
if not model.name:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"cross-encoder reranker requires name in reranking.model"
|
"cross-encoder reranker requires name in reranking.model"
|
||||||
)
|
)
|
||||||
|
from haiku.rag.reranking.cross_encoder import CrossEncoderReranker
|
||||||
|
|
||||||
return CrossEncoderReranker(model.name)
|
return CrossEncoderReranker(model.name)
|
||||||
except ImportError: # pragma: no cover
|
except ImportError: # pragma: no cover
|
||||||
return None
|
return None
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue