Remove vllm and lmstudio custom configs, they can now use the openai base
This commit is contained in:
parent
d4861b6408
commit
958fe43f2a
8 changed files with 15 additions and 55 deletions
|
|
@ -21,6 +21,7 @@
|
|||
### Removed
|
||||
|
||||
- Deleted obsolete embedder implementations: `ollama.py`, `openai.py`, `vllm.py`, `lm_studio.py`, `base.py`
|
||||
- Removed `VLLMConfig` and `LMStudioConfig` from configuration (use `base_url` in model config instead)
|
||||
|
||||
## [0.22.0] - 2025-12-19
|
||||
|
||||
|
|
|
|||
|
|
@ -378,10 +378,7 @@ reranking:
|
|||
model:
|
||||
provider: vllm
|
||||
name: mixedbread-ai/mxbai-rerank-base-v2
|
||||
|
||||
providers:
|
||||
vllm:
|
||||
rerank_base_url: http://localhost:8001
|
||||
base_url: http://localhost:8001
|
||||
```
|
||||
|
||||
**Note:** vLLM reranking uses the `/rerank` API endpoint. You need to run a vLLM server separately with a reranking model loaded. Consult the specific model's documentation for proper vLLM serving configuration.
|
||||
**Note:** vLLM reranking uses the `/v1/rerank` API endpoint. You need to run a vLLM server separately with a reranking model loaded.
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ from haiku.rag.config.models import (
|
|||
EmbeddingModelConfig,
|
||||
EmbeddingsConfig,
|
||||
LanceDBConfig,
|
||||
LMStudioConfig,
|
||||
ModelConfig,
|
||||
MonitorConfig,
|
||||
OllamaConfig,
|
||||
|
|
@ -22,7 +21,6 @@ from haiku.rag.config.models import (
|
|||
RerankingConfig,
|
||||
ResearchConfig,
|
||||
StorageConfig,
|
||||
VLLMConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
|
|
@ -33,7 +31,6 @@ __all__ = [
|
|||
"EmbeddingModelConfig",
|
||||
"EmbeddingsConfig",
|
||||
"LanceDBConfig",
|
||||
"LMStudioConfig",
|
||||
"ModelConfig",
|
||||
"MonitorConfig",
|
||||
"OllamaConfig",
|
||||
|
|
@ -43,7 +40,6 @@ __all__ = [
|
|||
"RerankingConfig",
|
||||
"ResearchConfig",
|
||||
"StorageConfig",
|
||||
"VLLMConfig",
|
||||
"find_config_file",
|
||||
"generate_default_config",
|
||||
"get_config",
|
||||
|
|
|
|||
|
|
@ -142,27 +142,14 @@ class OllamaConfig(BaseModel):
|
|||
)
|
||||
|
||||
|
||||
class VLLMConfig(BaseModel):
|
||||
embeddings_base_url: str = ""
|
||||
rerank_base_url: str = ""
|
||||
qa_base_url: str = ""
|
||||
research_base_url: str = ""
|
||||
|
||||
|
||||
class DoclingServeConfig(BaseModel):
|
||||
base_url: str = "http://localhost:5001"
|
||||
api_key: str = ""
|
||||
timeout: int = 300
|
||||
|
||||
|
||||
class LMStudioConfig(BaseModel):
|
||||
base_url: str = "http://localhost:1234"
|
||||
|
||||
|
||||
class ProvidersConfig(BaseModel):
|
||||
ollama: OllamaConfig = Field(default_factory=OllamaConfig)
|
||||
vllm: VLLMConfig = Field(default_factory=VLLMConfig)
|
||||
lm_studio: LMStudioConfig = Field(default_factory=LMStudioConfig)
|
||||
docling_serve: DoclingServeConfig = Field(default_factory=DoclingServeConfig)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,10 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
|||
try:
|
||||
from haiku.rag.reranking.vllm import VLLMReranker
|
||||
|
||||
reranker = VLLMReranker(config.reranking.model.name)
|
||||
base_url = config.reranking.model.base_url
|
||||
if not base_url:
|
||||
raise ValueError("vLLM reranker requires base_url in reranking.model")
|
||||
reranker = VLLMReranker(config.reranking.model.name, base_url)
|
||||
except ImportError:
|
||||
reranker = None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,14 +1,13 @@
|
|||
import httpx
|
||||
|
||||
from haiku.rag.config import Config
|
||||
from haiku.rag.reranking.base import RerankerBase
|
||||
from haiku.rag.store.models.chunk import Chunk
|
||||
|
||||
|
||||
class VLLMReranker(RerankerBase): # pragma: no cover
|
||||
def __init__(self, model: str):
|
||||
def __init__(self, model: str, base_url: str):
|
||||
self._model = model
|
||||
self._base_url = Config.providers.vllm.rerank_base_url
|
||||
self._base_url = base_url
|
||||
|
||||
async def rerank(
|
||||
self, query: str, chunks: list[Chunk], top_n: int = 10
|
||||
|
|
|
|||
|
|
@ -5,13 +5,11 @@ from datasets import Dataset
|
|||
from evaluations.evaluators import LLMJudge
|
||||
|
||||
from haiku.rag.client import HaikuRAG
|
||||
from haiku.rag.config import Config
|
||||
from haiku.rag.config.models import ModelConfig
|
||||
from haiku.rag.qa.agent import QuestionAnswerAgent
|
||||
|
||||
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
||||
ANTHROPIC_AVAILABLE = bool(os.getenv("ANTHROPIC_API_KEY"))
|
||||
VLLM_QA_AVAILABLE = bool(Config.providers.vllm.qa_base_url)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -87,26 +85,3 @@ async def test_qa_anthropic(qa_corpus: Dataset, temp_db_path):
|
|||
assert is_equivalent, (
|
||||
f"Generated answer not equivalent to expected answer.\nQuestion: {question}\nGenerated: {answer}\nExpected: {expected_answer}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(not VLLM_QA_AVAILABLE, reason="vLLM QA server not configured")
|
||||
async def test_qa_vllm(qa_corpus: Dataset, temp_db_path):
|
||||
"""Test vLLM QA with LLM judge."""
|
||||
client = HaikuRAG(temp_db_path, create=True)
|
||||
qa = QuestionAnswerAgent(client, ModelConfig(provider="vllm", name="Qwen/Qwen3-4B"))
|
||||
llm_judge = LLMJudge()
|
||||
|
||||
doc = qa_corpus[1]
|
||||
await client.create_document(
|
||||
content=doc["document_extracted"], uri=doc["document_id"]
|
||||
)
|
||||
|
||||
question = doc["question"]
|
||||
expected_answer = doc["answer"]
|
||||
answer, _ = await qa.answer(question)
|
||||
is_equivalent = await llm_judge.judge_answers(question, answer, expected_answer)
|
||||
|
||||
assert is_equivalent, (
|
||||
f"Generated answer not equivalent to expected answer.\nQuestion: {question}\nGenerated: {answer}\nExpected: {expected_answer}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,13 +2,12 @@ import os
|
|||
|
||||
import pytest
|
||||
|
||||
from haiku.rag.config import Config
|
||||
from haiku.rag.reranking.base import RerankerBase
|
||||
from haiku.rag.reranking.vllm import VLLMReranker
|
||||
from haiku.rag.store.models.chunk import Chunk
|
||||
|
||||
COHERE_AVAILABLE = bool(os.getenv("CO_API_KEY"))
|
||||
VLLM_RERANK_AVAILABLE = bool(Config.providers.vllm.rerank_base_url)
|
||||
VLLM_RERANK_BASE_URL = os.getenv("VLLM_RERANK_BASE_URL", "")
|
||||
ZEROENTROPY_AVAILABLE = bool(os.getenv("ZEROENTROPY_API_KEY"))
|
||||
|
||||
chunks = [
|
||||
|
|
@ -41,6 +40,7 @@ async def test_reranker_base():
|
|||
@pytest.mark.asyncio
|
||||
async def test_mxbai_reranker():
|
||||
try:
|
||||
from haiku.rag.config import Config
|
||||
from haiku.rag.config.models import ModelConfig
|
||||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||
|
||||
|
|
@ -80,11 +80,13 @@ async def test_cohere_reranker():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.skipif(
|
||||
not VLLM_RERANK_AVAILABLE, reason="vLLM rerank server not configured"
|
||||
not VLLM_RERANK_BASE_URL, reason="vLLM rerank server not configured"
|
||||
)
|
||||
async def test_vllm_reranker():
|
||||
try:
|
||||
reranker = VLLMReranker("mixedbread-ai/mxbai-rerank-base-v2")
|
||||
reranker = VLLMReranker(
|
||||
"mixedbread-ai/mxbai-rerank-base-v2", VLLM_RERANK_BASE_URL
|
||||
)
|
||||
|
||||
reranked = await reranker.rerank(
|
||||
"Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2
|
||||
|
|
|
|||
Loading…
Reference in a new issue