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
|
### Removed
|
||||||
|
|
||||||
- Deleted obsolete embedder implementations: `ollama.py`, `openai.py`, `vllm.py`, `lm_studio.py`, `base.py`
|
- 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
|
## [0.22.0] - 2025-12-19
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -378,10 +378,7 @@ reranking:
|
||||||
model:
|
model:
|
||||||
provider: vllm
|
provider: vllm
|
||||||
name: mixedbread-ai/mxbai-rerank-base-v2
|
name: mixedbread-ai/mxbai-rerank-base-v2
|
||||||
|
base_url: http://localhost:8001
|
||||||
providers:
|
|
||||||
vllm:
|
|
||||||
rerank_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,
|
EmbeddingModelConfig,
|
||||||
EmbeddingsConfig,
|
EmbeddingsConfig,
|
||||||
LanceDBConfig,
|
LanceDBConfig,
|
||||||
LMStudioConfig,
|
|
||||||
ModelConfig,
|
ModelConfig,
|
||||||
MonitorConfig,
|
MonitorConfig,
|
||||||
OllamaConfig,
|
OllamaConfig,
|
||||||
|
|
@ -22,7 +21,6 @@ from haiku.rag.config.models import (
|
||||||
RerankingConfig,
|
RerankingConfig,
|
||||||
ResearchConfig,
|
ResearchConfig,
|
||||||
StorageConfig,
|
StorageConfig,
|
||||||
VLLMConfig,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
|
@ -33,7 +31,6 @@ __all__ = [
|
||||||
"EmbeddingModelConfig",
|
"EmbeddingModelConfig",
|
||||||
"EmbeddingsConfig",
|
"EmbeddingsConfig",
|
||||||
"LanceDBConfig",
|
"LanceDBConfig",
|
||||||
"LMStudioConfig",
|
|
||||||
"ModelConfig",
|
"ModelConfig",
|
||||||
"MonitorConfig",
|
"MonitorConfig",
|
||||||
"OllamaConfig",
|
"OllamaConfig",
|
||||||
|
|
@ -43,7 +40,6 @@ __all__ = [
|
||||||
"RerankingConfig",
|
"RerankingConfig",
|
||||||
"ResearchConfig",
|
"ResearchConfig",
|
||||||
"StorageConfig",
|
"StorageConfig",
|
||||||
"VLLMConfig",
|
|
||||||
"find_config_file",
|
"find_config_file",
|
||||||
"generate_default_config",
|
"generate_default_config",
|
||||||
"get_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):
|
class DoclingServeConfig(BaseModel):
|
||||||
base_url: str = "http://localhost:5001"
|
base_url: str = "http://localhost:5001"
|
||||||
api_key: str = ""
|
api_key: str = ""
|
||||||
timeout: int = 300
|
timeout: int = 300
|
||||||
|
|
||||||
|
|
||||||
class LMStudioConfig(BaseModel):
|
|
||||||
base_url: str = "http://localhost:1234"
|
|
||||||
|
|
||||||
|
|
||||||
class ProvidersConfig(BaseModel):
|
class ProvidersConfig(BaseModel):
|
||||||
ollama: OllamaConfig = Field(default_factory=OllamaConfig)
|
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)
|
docling_serve: DoclingServeConfig = Field(default_factory=DoclingServeConfig)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -45,7 +45,10 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
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:
|
except ImportError:
|
||||||
reranker = None
|
reranker = None
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,14 +1,13 @@
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
|
||||||
class VLLMReranker(RerankerBase): # pragma: no cover
|
class VLLMReranker(RerankerBase): # pragma: no cover
|
||||||
def __init__(self, model: str):
|
def __init__(self, model: str, base_url: str):
|
||||||
self._model = model
|
self._model = model
|
||||||
self._base_url = Config.providers.vllm.rerank_base_url
|
self._base_url = base_url
|
||||||
|
|
||||||
async def rerank(
|
async def rerank(
|
||||||
self, query: str, chunks: list[Chunk], top_n: int = 10
|
self, query: str, chunks: list[Chunk], top_n: int = 10
|
||||||
|
|
|
||||||
|
|
@ -5,13 +5,11 @@ from datasets import Dataset
|
||||||
from evaluations.evaluators import LLMJudge
|
from evaluations.evaluators import LLMJudge
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import Config
|
|
||||||
from haiku.rag.config.models import ModelConfig
|
from haiku.rag.config.models import ModelConfig
|
||||||
from haiku.rag.qa.agent import QuestionAnswerAgent
|
from haiku.rag.qa.agent import QuestionAnswerAgent
|
||||||
|
|
||||||
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
||||||
ANTHROPIC_AVAILABLE = bool(os.getenv("ANTHROPIC_API_KEY"))
|
ANTHROPIC_AVAILABLE = bool(os.getenv("ANTHROPIC_API_KEY"))
|
||||||
VLLM_QA_AVAILABLE = bool(Config.providers.vllm.qa_base_url)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -87,26 +85,3 @@ async def test_qa_anthropic(qa_corpus: Dataset, temp_db_path):
|
||||||
assert is_equivalent, (
|
assert is_equivalent, (
|
||||||
f"Generated answer not equivalent to expected answer.\nQuestion: {question}\nGenerated: {answer}\nExpected: {expected_answer}"
|
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
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
COHERE_AVAILABLE = bool(os.getenv("CO_API_KEY"))
|
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"))
|
ZEROENTROPY_AVAILABLE = bool(os.getenv("ZEROENTROPY_API_KEY"))
|
||||||
|
|
||||||
chunks = [
|
chunks = [
|
||||||
|
|
@ -41,6 +40,7 @@ async def test_reranker_base():
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_mxbai_reranker():
|
async def test_mxbai_reranker():
|
||||||
try:
|
try:
|
||||||
|
from haiku.rag.config import Config
|
||||||
from haiku.rag.config.models import ModelConfig
|
from haiku.rag.config.models import ModelConfig
|
||||||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||||
|
|
||||||
|
|
@ -80,11 +80,13 @@ async def test_cohere_reranker():
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.skipif(
|
@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():
|
async def test_vllm_reranker():
|
||||||
try:
|
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(
|
reranked = await reranker.rerank(
|
||||||
"Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2
|
"Who wrote 'To Kill a Mockingbird'?", chunks, top_n=2
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue