diff --git a/docs/configuration.md b/docs/configuration.md index 00845f58..ae9dd7e0 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -97,6 +97,18 @@ QA_MODEL="claude-3-5-haiku-20241022" # or claude-3-5-sonnet-20241022, etc. ANTHROPIC_API_KEY="your-api-key" ``` +### vLLM + +For high-performance local inference, you can use vLLM to serve models with OpenAI-compatible APIs: + +```bash +QA_PROVIDER="vllm" +QA_MODEL="Qwen/Qwen3-4B" # Any model with tool support in vLLM +VLLM_QA_BASE_URL="http://localhost:8002" # vLLM server URL +``` + +**Note:** You need to run a vLLM server separately with a model that supports tool calling loaded. Consult the specific model's documentation for proper vLLM serving configuration. + ### Other Providers Any provider supported by Pydantic AI can be used. Examples include: diff --git a/docs/installation.md b/docs/installation.md index e3416913..eb1e0750 100644 --- a/docs/installation.md +++ b/docs/installation.md @@ -40,6 +40,9 @@ pip install vllm # Serve an embedding model vllm serve mixedbread-ai/mxbai-embed-large-v1 --port 8000 +# Serve a model for QA (requires tool calling support) +vllm serve Qwen/Qwen3-4B --port 8002 --enable-auto-tool-choice --tool-call-parser hermes + # Serve a model for reranking vllm serve mixedbread-ai/mxbai-rerank-base-v2 --hf_overrides '{"architectures": ["Qwen2ForSequenceClassification"],"classifier_from_token": ["0", "1"], "method": "from_2_way_softmax"}' --port 8001 ``` @@ -53,9 +56,14 @@ EMBEDDINGS_MODEL="mixedbread-ai/mxbai-embed-large-v1" EMBEDDINGS_VECTOR_DIM=512 VLLM_EMBEDDINGS_BASE_URL="http://localhost:8000" +# QA (optional) +QA_PROVIDER="vllm" +QA_MODEL="Qwen/Qwen3-4B" +VLLM_QA_BASE_URL="http://localhost:8002" + # Reranking (optional) RERANK_PROVIDER="vllm" -RERANK_MODEL="microsoft/DialoGPT-medium" +RERANK_MODEL="mixedbread-ai/mxbai-rerank-base-v2" VLLM_RERANK_BASE_URL="http://localhost:8001" ``` diff --git a/src/haiku/rag/qa/agent.py b/src/haiku/rag/qa/agent.py index 711106ba..594fad6e 100644 --- a/src/haiku/rag/qa/agent.py +++ b/src/haiku/rag/qa/agent.py @@ -2,6 +2,7 @@ from pydantic import BaseModel, Field from pydantic_ai import Agent, RunContext from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.ollama import OllamaProvider +from pydantic_ai.providers.openai import OpenAIProvider from haiku.rag.client import HaikuRAG from haiku.rag.config import Config @@ -65,6 +66,13 @@ class QuestionAnswerAgent: model_name=model, provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"), ) + elif provider == "vllm": + return OpenAIChatModel( + model_name=model, + provider=OpenAIProvider( + base_url=f"{Config.VLLM_QA_BASE_URL}/v1", api_key="none" + ), + ) else: # For all other providers, use the provider:model format return f"{provider}:{model}" diff --git a/tests/llm_judge.py b/tests/llm_judge.py index 68ffe8c8..435e517b 100644 --- a/tests/llm_judge.py +++ b/tests/llm_judge.py @@ -1,6 +1,6 @@ from pydantic import BaseModel from pydantic_ai import Agent -from pydantic_ai.models.openai import OpenAIModel +from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.ollama import OllamaProvider from haiku.rag.config import Config @@ -37,9 +37,9 @@ class LLMJudgeResponseSchema(BaseModel): class LLMJudge: """LLM-as-judge for evaluating answer equivalence using Pydantic AI.""" - def __init__(self, model: str = Config.QA_MODEL): + def __init__(self, model: str = "qwen3"): # Create Ollama model - ollama_model = OpenAIModel( + ollama_model = OpenAIChatModel( model_name=model, provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"), ) diff --git a/tests/test_qa.py b/tests/test_qa.py index 837c8f54..2fa2d417 100644 --- a/tests/test_qa.py +++ b/tests/test_qa.py @@ -9,6 +9,7 @@ from .llm_judge import LLMJudge OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY) ANTHROPIC_AVAILABLE = bool(Config.ANTHROPIC_API_KEY) +VLLM_QA_AVAILABLE = bool(Config.VLLM_QA_BASE_URL) @pytest.mark.asyncio @@ -80,3 +81,26 @@ 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) + qa = QuestionAnswerAgent(client, "vllm", "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}" + )