Support vLLM for QA agents
This commit is contained in:
parent
f2847d0524
commit
bc28e60cf7
5 changed files with 56 additions and 4 deletions
|
|
@ -97,6 +97,18 @@ QA_MODEL="claude-3-5-haiku-20241022" # or claude-3-5-sonnet-20241022, etc.
|
||||||
ANTHROPIC_API_KEY="your-api-key"
|
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
|
### Other Providers
|
||||||
|
|
||||||
Any provider supported by Pydantic AI can be used. Examples include:
|
Any provider supported by Pydantic AI can be used. Examples include:
|
||||||
|
|
|
||||||
|
|
@ -40,6 +40,9 @@ pip install vllm
|
||||||
# Serve an embedding model
|
# Serve an embedding model
|
||||||
vllm serve mixedbread-ai/mxbai-embed-large-v1 --port 8000
|
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
|
# 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
|
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
|
EMBEDDINGS_VECTOR_DIM=512
|
||||||
VLLM_EMBEDDINGS_BASE_URL="http://localhost:8000"
|
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)
|
# Reranking (optional)
|
||||||
RERANK_PROVIDER="vllm"
|
RERANK_PROVIDER="vllm"
|
||||||
RERANK_MODEL="microsoft/DialoGPT-medium"
|
RERANK_MODEL="mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
VLLM_RERANK_BASE_URL="http://localhost:8001"
|
VLLM_RERANK_BASE_URL="http://localhost:8001"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ from pydantic import BaseModel, Field
|
||||||
from pydantic_ai import Agent, RunContext
|
from pydantic_ai import Agent, RunContext
|
||||||
from pydantic_ai.models.openai import OpenAIChatModel
|
from pydantic_ai.models.openai import OpenAIChatModel
|
||||||
from pydantic_ai.providers.ollama import OllamaProvider
|
from pydantic_ai.providers.ollama import OllamaProvider
|
||||||
|
from pydantic_ai.providers.openai import OpenAIProvider
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
|
|
@ -65,6 +66,13 @@ class QuestionAnswerAgent:
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
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:
|
else:
|
||||||
# For all other providers, use the provider:model format
|
# For all other providers, use the provider:model format
|
||||||
return f"{provider}:{model}"
|
return f"{provider}:{model}"
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from pydantic_ai import Agent
|
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 pydantic_ai.providers.ollama import OllamaProvider
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
|
|
@ -37,9 +37,9 @@ class LLMJudgeResponseSchema(BaseModel):
|
||||||
class LLMJudge:
|
class LLMJudge:
|
||||||
"""LLM-as-judge for evaluating answer equivalence using Pydantic AI."""
|
"""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
|
# Create Ollama model
|
||||||
ollama_model = OpenAIModel(
|
ollama_model = OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ from .llm_judge import LLMJudge
|
||||||
|
|
||||||
OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY)
|
OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY)
|
||||||
ANTHROPIC_AVAILABLE = bool(Config.ANTHROPIC_API_KEY)
|
ANTHROPIC_AVAILABLE = bool(Config.ANTHROPIC_API_KEY)
|
||||||
|
VLLM_QA_AVAILABLE = bool(Config.VLLM_QA_BASE_URL)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
@ -80,3 +81,26 @@ 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)
|
||||||
|
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}"
|
||||||
|
)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue