Rename model to name under model
This commit is contained in:
parent
41ea675964
commit
a9701616c8
17 changed files with 84 additions and 91 deletions
|
|
@ -37,13 +37,13 @@ environment: production
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: qwen3-embedding:4b
|
name: qwen3-embedding:4b
|
||||||
vector_dim: 2560
|
vector_dim: 2560
|
||||||
|
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: gpt-oss
|
name: gpt-oss
|
||||||
enable_thinking: false
|
enable_thinking: false
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -72,18 +72,18 @@ lancedb:
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: qwen3-embedding:4b
|
name: qwen3-embedding:4b
|
||||||
vector_dim: 2560
|
vector_dim: 2560
|
||||||
|
|
||||||
reranking:
|
reranking:
|
||||||
model:
|
model:
|
||||||
provider: "" # Empty to disable, or mxbai, cohere, zeroentropy, vllm
|
provider: "" # Empty to disable, or mxbai, cohere, zeroentropy, vllm
|
||||||
model: ""
|
name: ""
|
||||||
|
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: gpt-oss
|
name: gpt-oss
|
||||||
enable_thinking: false
|
enable_thinking: false
|
||||||
max_sub_questions: 3
|
max_sub_questions: 3
|
||||||
max_iterations: 2
|
max_iterations: 2
|
||||||
|
|
@ -92,7 +92,7 @@ qa:
|
||||||
research:
|
research:
|
||||||
model:
|
model:
|
||||||
provider: "" # Empty to use qa settings
|
provider: "" # Empty to use qa settings
|
||||||
model: ""
|
name: ""
|
||||||
enable_thinking: true
|
enable_thinking: true
|
||||||
max_iterations: 3
|
max_iterations: 3
|
||||||
confidence_threshold: 0.8
|
confidence_threshold: 0.8
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,7 @@ Configure model behavior for `qa` and `research` workflows. These settings apply
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: gpt-oss
|
name: gpt-oss
|
||||||
temperature: 0.7
|
temperature: 0.7
|
||||||
max_tokens: 500
|
max_tokens: 500
|
||||||
```
|
```
|
||||||
|
|
@ -74,7 +74,7 @@ If you use Ollama, you can use any pulled model that supports embeddings.
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: mxbai-embed-large
|
name: mxbai-embed-large
|
||||||
vector_dim: 1024
|
vector_dim: 1024
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -106,7 +106,7 @@ uv pip install haiku.rag-slim[voyageai]
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
provider: voyageai
|
provider: voyageai
|
||||||
model: voyage-3.5
|
name: voyage-3.5
|
||||||
vector_dim: 1024
|
vector_dim: 1024
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -124,7 +124,7 @@ OpenAI embeddings are included in the default installation:
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
provider: openai
|
provider: openai
|
||||||
model: text-embedding-3-small # or text-embedding-3-large
|
name: text-embedding-3-small # or text-embedding-3-large
|
||||||
vector_dim: 1536
|
vector_dim: 1536
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -142,7 +142,7 @@ For high-performance local inference, you can use vLLM to serve embedding models
|
||||||
embeddings:
|
embeddings:
|
||||||
model:
|
model:
|
||||||
provider: vllm
|
provider: vllm
|
||||||
model: mixedbread-ai/mxbai-embed-large-v1
|
name: mixedbread-ai/mxbai-embed-large-v1
|
||||||
vector_dim: 512
|
vector_dim: 512
|
||||||
|
|
||||||
providers:
|
providers:
|
||||||
|
|
@ -162,7 +162,7 @@ Configure which LLM provider to use for question answering. Any provider and mod
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: gpt-oss
|
name: gpt-oss
|
||||||
```
|
```
|
||||||
|
|
||||||
The Ollama base URL can be configured via the `OLLAMA_BASE_URL` environment variable, config file, or defaults to `http://localhost:11434`:
|
The Ollama base URL can be configured via the `OLLAMA_BASE_URL` environment variable, config file, or defaults to `http://localhost:11434`:
|
||||||
|
|
@ -187,7 +187,7 @@ OpenAI QA is included in the default installation:
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: openai
|
provider: openai
|
||||||
model: gpt-4o-mini # or gpt-4, gpt-3.5-turbo, etc.
|
name: gpt-4o-mini # or gpt-4, gpt-3.5-turbo, etc.
|
||||||
```
|
```
|
||||||
|
|
||||||
Set your API key via environment variable:
|
Set your API key via environment variable:
|
||||||
|
|
@ -204,7 +204,7 @@ Anthropic QA is included in the default installation:
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: anthropic
|
provider: anthropic
|
||||||
model: claude-3-5-haiku-20241022 # or claude-3-5-sonnet-20241022, etc.
|
name: claude-3-5-haiku-20241022 # or claude-3-5-sonnet-20241022, etc.
|
||||||
```
|
```
|
||||||
|
|
||||||
Set your API key via environment variable:
|
Set your API key via environment variable:
|
||||||
|
|
@ -221,7 +221,7 @@ For high-performance local inference:
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: vllm
|
provider: vllm
|
||||||
model: Qwen/Qwen3-4B # Any model with tool support in vLLM
|
name: Qwen/Qwen3-4B # Any model with tool support in vLLM
|
||||||
|
|
||||||
providers:
|
providers:
|
||||||
vllm:
|
vllm:
|
||||||
|
|
@ -239,19 +239,19 @@ Any provider supported by Pydantic AI can be used. Examples:
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: gemini
|
provider: gemini
|
||||||
model: gemini-1.5-flash
|
name: gemini-1.5-flash
|
||||||
|
|
||||||
# Groq
|
# Groq
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: groq
|
provider: groq
|
||||||
model: llama-3.3-70b-versatile
|
name: llama-3.3-70b-versatile
|
||||||
|
|
||||||
# Mistral
|
# Mistral
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: mistral
|
provider: mistral
|
||||||
model: mistral-small-latest
|
name: mistral-small-latest
|
||||||
```
|
```
|
||||||
|
|
||||||
See the [Pydantic AI documentation](https://ai.pydantic.dev/models/) for the complete list of supported providers and models.
|
See the [Pydantic AI documentation](https://ai.pydantic.dev/models/) for the complete list of supported providers and models.
|
||||||
|
|
@ -276,7 +276,7 @@ Then configure:
|
||||||
reranking:
|
reranking:
|
||||||
model:
|
model:
|
||||||
provider: mxbai
|
provider: mxbai
|
||||||
model: mixedbread-ai/mxbai-rerank-base-v2
|
name: mixedbread-ai/mxbai-rerank-base-v2
|
||||||
```
|
```
|
||||||
|
|
||||||
### Cohere
|
### Cohere
|
||||||
|
|
@ -293,7 +293,7 @@ Then configure:
|
||||||
reranking:
|
reranking:
|
||||||
model:
|
model:
|
||||||
provider: cohere
|
provider: cohere
|
||||||
model: rerank-v3.5
|
name: rerank-v3.5
|
||||||
```
|
```
|
||||||
|
|
||||||
Set your API key via environment variable:
|
Set your API key via environment variable:
|
||||||
|
|
@ -316,7 +316,7 @@ Then configure:
|
||||||
reranking:
|
reranking:
|
||||||
model:
|
model:
|
||||||
provider: zeroentropy
|
provider: zeroentropy
|
||||||
model: zerank-1 # Currently the only available model
|
name: zerank-1 # Currently the only available model
|
||||||
```
|
```
|
||||||
|
|
||||||
Set your API key via environment variable:
|
Set your API key via environment variable:
|
||||||
|
|
@ -333,7 +333,7 @@ For high-performance local reranking using dedicated reranking models:
|
||||||
reranking:
|
reranking:
|
||||||
model:
|
model:
|
||||||
provider: vllm
|
provider: vllm
|
||||||
model: mixedbread-ai/mxbai-rerank-base-v2
|
name: mixedbread-ai/mxbai-rerank-base-v2
|
||||||
|
|
||||||
providers:
|
providers:
|
||||||
vllm:
|
vllm:
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,7 @@ Configure the QA workflow:
|
||||||
qa:
|
qa:
|
||||||
model:
|
model:
|
||||||
provider: ollama
|
provider: ollama
|
||||||
model: gpt-oss
|
name: gpt-oss
|
||||||
enable_thinking: false
|
enable_thinking: false
|
||||||
max_sub_questions: 3 # Maximum sub-questions for deep QA
|
max_sub_questions: 3 # Maximum sub-questions for deep QA
|
||||||
max_iterations: 2 # Maximum search iterations per sub-question
|
max_iterations: 2 # Maximum search iterations per sub-question
|
||||||
|
|
@ -30,7 +30,7 @@ Configure the multi-agent research workflow:
|
||||||
research:
|
research:
|
||||||
model:
|
model:
|
||||||
provider: "" # Empty to use qa settings
|
provider: "" # Empty to use qa settings
|
||||||
model: "" # Empty to use qa model
|
name: "" # Empty to use qa model
|
||||||
enable_thinking: true
|
enable_thinking: true
|
||||||
max_iterations: 3
|
max_iterations: 3
|
||||||
confidence_threshold: 0.8
|
confidence_threshold: 0.8
|
||||||
|
|
|
||||||
|
|
@ -44,18 +44,16 @@ def build_experiment_metadata(
|
||||||
"dataset": dataset_key,
|
"dataset": dataset_key,
|
||||||
"test_cases": test_cases,
|
"test_cases": test_cases,
|
||||||
"embedder_provider": config.embeddings.model.provider,
|
"embedder_provider": config.embeddings.model.provider,
|
||||||
"embedder_model": config.embeddings.model.model,
|
"embedder_model": config.embeddings.model.name,
|
||||||
"embedder_dim": config.embeddings.vector_dim,
|
"embedder_dim": config.embeddings.vector_dim,
|
||||||
"chunk_size": config.processing.chunk_size,
|
"chunk_size": config.processing.chunk_size,
|
||||||
"context_chunk_radius": config.processing.context_chunk_radius,
|
"context_chunk_radius": config.processing.context_chunk_radius,
|
||||||
"rerank_provider": config.reranking.model.provider
|
"rerank_provider": config.reranking.model.provider
|
||||||
if config.reranking.model
|
if config.reranking.model
|
||||||
else None,
|
else None,
|
||||||
"rerank_model": config.reranking.model.model
|
"rerank_model": config.reranking.model.name if config.reranking.model else None,
|
||||||
if config.reranking.model
|
|
||||||
else None,
|
|
||||||
"qa_provider": config.qa.model.provider,
|
"qa_provider": config.qa.model.provider,
|
||||||
"qa_model": config.qa.model.model,
|
"qa_model": config.qa.model.name,
|
||||||
"judge_provider": "ollama",
|
"judge_provider": "ollama",
|
||||||
"judge_model": judge_model,
|
"judge_model": judge_model,
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -47,7 +47,7 @@ if not db_path.exists():
|
||||||
|
|
||||||
logger.info(f"Initializing research assistant with database: {db_path}")
|
logger.info(f"Initializing research assistant with database: {db_path}")
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Research Provider: {Config.research.model.provider}, Model: {Config.research.model.model}"
|
f"Research Provider: {Config.research.model.provider}, Model: {Config.research.model.name}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Store client reference for proper lifecycle management
|
# Store client reference for proper lifecycle management
|
||||||
|
|
@ -154,7 +154,7 @@ async def health_check(_: Request) -> JSONResponse:
|
||||||
"status": "healthy",
|
"status": "healthy",
|
||||||
"agent_model": str(agent.model),
|
"agent_model": str(agent.model),
|
||||||
"research_provider": Config.research.model.provider,
|
"research_provider": Config.research.model.provider,
|
||||||
"research_model": Config.research.model.model,
|
"research_model": Config.research.model.name,
|
||||||
"db_path": str(db_path),
|
"db_path": str(db_path),
|
||||||
"db_exists": db_path.exists(),
|
"db_exists": db_path.exists(),
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -17,18 +17,21 @@ providers:
|
||||||
base_url: http://host.docker.internal:11434
|
base_url: http://host.docker.internal:11434
|
||||||
|
|
||||||
research:
|
research:
|
||||||
provider: ollama
|
model:
|
||||||
model: gpt-oss:latest
|
provider: ollama
|
||||||
|
name: gpt-oss:latest
|
||||||
max_iterations: 3
|
max_iterations: 3
|
||||||
confidence_threshold: 0.8
|
confidence_threshold: 0.8
|
||||||
max_concurrency: 1
|
max_concurrency: 1
|
||||||
|
|
||||||
# For OpenAI:
|
# For OpenAI:
|
||||||
# research:
|
# research:
|
||||||
# provider: openai
|
# model:
|
||||||
# model: gpt-4o-mini
|
# provider: openai
|
||||||
|
# name: gpt-4o-mini
|
||||||
|
|
||||||
# For Anthropic:
|
# For Anthropic:
|
||||||
# research:
|
# research:
|
||||||
# provider: anthropic
|
# model:
|
||||||
# model: claude-3-5-haiku-20241022
|
# provider: anthropic
|
||||||
|
# name: claude-3-5-haiku-20241022
|
||||||
|
|
|
||||||
|
|
@ -11,14 +11,14 @@ class ModelConfig(BaseModel):
|
||||||
|
|
||||||
Attributes:
|
Attributes:
|
||||||
provider: Model provider (ollama, openai, anthropic, etc.)
|
provider: Model provider (ollama, openai, anthropic, etc.)
|
||||||
model: Model name/identifier
|
name: Model name/identifier
|
||||||
enable_thinking: Control reasoning behavior (true/false/None for default)
|
enable_thinking: Control reasoning behavior (true/false/None for default)
|
||||||
temperature: Sampling temperature (0.0 to 1.0+)
|
temperature: Sampling temperature (0.0 to 1.0+)
|
||||||
max_tokens: Maximum tokens to generate
|
max_tokens: Maximum tokens to generate
|
||||||
"""
|
"""
|
||||||
|
|
||||||
provider: str = "ollama"
|
provider: str = "ollama"
|
||||||
model: str = "gpt-oss"
|
name: str = "gpt-oss"
|
||||||
|
|
||||||
enable_thinking: bool | None = None
|
enable_thinking: bool | None = None
|
||||||
temperature: float | None = None
|
temperature: float | None = None
|
||||||
|
|
@ -47,7 +47,7 @@ class EmbeddingsConfig(BaseModel):
|
||||||
model: ModelConfig = Field(
|
model: ModelConfig = Field(
|
||||||
default_factory=lambda: ModelConfig(
|
default_factory=lambda: ModelConfig(
|
||||||
provider="ollama",
|
provider="ollama",
|
||||||
model="qwen3-embedding:4b",
|
name="qwen3-embedding:4b",
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
vector_dim: int = 2560
|
vector_dim: int = 2560
|
||||||
|
|
@ -61,7 +61,7 @@ class QAConfig(BaseModel):
|
||||||
model: ModelConfig = Field(
|
model: ModelConfig = Field(
|
||||||
default_factory=lambda: ModelConfig(
|
default_factory=lambda: ModelConfig(
|
||||||
provider="ollama",
|
provider="ollama",
|
||||||
model="gpt-oss",
|
name="gpt-oss",
|
||||||
enable_thinking=False,
|
enable_thinking=False,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
@ -74,7 +74,7 @@ class ResearchConfig(BaseModel):
|
||||||
model: ModelConfig = Field(
|
model: ModelConfig = Field(
|
||||||
default_factory=lambda: ModelConfig(
|
default_factory=lambda: ModelConfig(
|
||||||
provider="ollama",
|
provider="ollama",
|
||||||
model="gpt-oss",
|
name="gpt-oss",
|
||||||
enable_thinking=True,
|
enable_thinking=True,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -16,7 +16,7 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
|
||||||
|
|
||||||
if config.embeddings.model.provider == "ollama":
|
if config.embeddings.model.provider == "ollama":
|
||||||
return OllamaEmbedder(
|
return OllamaEmbedder(
|
||||||
config.embeddings.model.model, config.embeddings.vector_dim, config
|
config.embeddings.model.name, config.embeddings.vector_dim, config
|
||||||
)
|
)
|
||||||
|
|
||||||
if config.embeddings.model.provider == "voyageai":
|
if config.embeddings.model.provider == "voyageai":
|
||||||
|
|
@ -29,21 +29,21 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
|
||||||
"uv pip install haiku.rag[voyageai]"
|
"uv pip install haiku.rag[voyageai]"
|
||||||
)
|
)
|
||||||
return VoyageAIEmbedder(
|
return VoyageAIEmbedder(
|
||||||
config.embeddings.model.model, config.embeddings.vector_dim, config
|
config.embeddings.model.name, config.embeddings.vector_dim, config
|
||||||
)
|
)
|
||||||
|
|
||||||
if config.embeddings.model.provider == "openai":
|
if config.embeddings.model.provider == "openai":
|
||||||
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
||||||
|
|
||||||
return OpenAIEmbedder(
|
return OpenAIEmbedder(
|
||||||
config.embeddings.model.model, config.embeddings.vector_dim, config
|
config.embeddings.model.name, config.embeddings.vector_dim, config
|
||||||
)
|
)
|
||||||
|
|
||||||
if config.embeddings.model.provider == "vllm":
|
if config.embeddings.model.provider == "vllm":
|
||||||
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
|
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
|
||||||
|
|
||||||
return VllmEmbedder(
|
return VllmEmbedder(
|
||||||
config.embeddings.model.model, config.embeddings.vector_dim, config
|
config.embeddings.model.name, config.embeddings.vector_dim, config
|
||||||
)
|
)
|
||||||
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|
|
||||||
|
|
@ -4,7 +4,7 @@ from haiku.rag.config import AppConfig, Config
|
||||||
|
|
||||||
|
|
||||||
class EmbedderBase:
|
class EmbedderBase:
|
||||||
_model: str = Config.embeddings.model.model
|
_model: str = Config.embeddings.model.name
|
||||||
_vector_dim: int = Config.embeddings.vector_dim
|
_vector_dim: int = Config.embeddings.vector_dim
|
||||||
_config: AppConfig = Config
|
_config: AppConfig = Config
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -45,7 +45,7 @@ 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.model)
|
reranker = VLLMReranker(config.reranking.model.name)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
reranker = None
|
reranker = None
|
||||||
|
|
||||||
|
|
@ -54,7 +54,7 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
||||||
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
||||||
|
|
||||||
# Use configured model or default to zerank-1
|
# Use configured model or default to zerank-1
|
||||||
model = config.reranking.model.model or "zerank-1"
|
model = config.reranking.model.name or "zerank-1"
|
||||||
reranker = ZeroEntropyReranker(model)
|
reranker = ZeroEntropyReranker(model)
|
||||||
except ImportError:
|
except ImportError:
|
||||||
reranker = None
|
reranker = None
|
||||||
|
|
|
||||||
|
|
@ -3,9 +3,7 @@ from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
|
||||||
class RerankerBase:
|
class RerankerBase:
|
||||||
_model: str | None = (
|
_model: str | None = Config.reranking.model.name if Config.reranking.model else None
|
||||||
Config.reranking.model.model if Config.reranking.model else None
|
|
||||||
)
|
|
||||||
|
|
||||||
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
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,7 @@ from haiku.rag.store.models.chunk import Chunk
|
||||||
class MxBAIReranker(RerankerBase):
|
class MxBAIReranker(RerankerBase):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
model_name = (
|
model_name = (
|
||||||
Config.reranking.model.model
|
Config.reranking.model.name
|
||||||
if Config.reranking.model
|
if Config.reranking.model
|
||||||
else "mxbai-rerank-base-v2"
|
else "mxbai-rerank-base-v2"
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -65,7 +65,7 @@ def get_model(
|
||||||
app_config = Config
|
app_config = Config
|
||||||
|
|
||||||
provider = model_config.provider
|
provider = model_config.provider
|
||||||
model = model_config.model
|
model = model_config.name
|
||||||
|
|
||||||
if provider == "ollama":
|
if provider == "ollama":
|
||||||
model_settings = None
|
model_settings = None
|
||||||
|
|
@ -366,13 +366,13 @@ def prefetch_models():
|
||||||
# Collect Ollama models from config
|
# Collect Ollama models from config
|
||||||
required_models: set[str] = set()
|
required_models: set[str] = set()
|
||||||
if Config.embeddings.model.provider == "ollama":
|
if Config.embeddings.model.provider == "ollama":
|
||||||
required_models.add(Config.embeddings.model.model)
|
required_models.add(Config.embeddings.model.name)
|
||||||
if Config.qa.model.provider == "ollama":
|
if Config.qa.model.provider == "ollama":
|
||||||
required_models.add(Config.qa.model.model)
|
required_models.add(Config.qa.model.name)
|
||||||
if Config.research.model.provider == "ollama":
|
if Config.research.model.provider == "ollama":
|
||||||
required_models.add(Config.research.model.model)
|
required_models.add(Config.research.model.name)
|
||||||
if Config.reranking.model and Config.reranking.model.provider == "ollama":
|
if Config.reranking.model and Config.reranking.model.provider == "ollama":
|
||||||
required_models.add(Config.reranking.model.model)
|
required_models.add(Config.reranking.model.name)
|
||||||
|
|
||||||
if not required_models:
|
if not required_models:
|
||||||
return
|
return
|
||||||
|
|
|
||||||
|
|
@ -17,7 +17,7 @@ def test_embedder_uses_config_from_get_embedder():
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=ModelConfig(
|
model=ModelConfig(
|
||||||
provider="ollama",
|
provider="ollama",
|
||||||
model="custom-model",
|
name="custom-model",
|
||||||
),
|
),
|
||||||
vector_dim=512,
|
vector_dim=512,
|
||||||
),
|
),
|
||||||
|
|
@ -42,7 +42,7 @@ def test_vllm_embedder_uses_config():
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=ModelConfig(
|
model=ModelConfig(
|
||||||
provider="vllm",
|
provider="vllm",
|
||||||
model="custom-vllm-model",
|
name="custom-vllm-model",
|
||||||
),
|
),
|
||||||
vector_dim=768,
|
vector_dim=768,
|
||||||
),
|
),
|
||||||
|
|
@ -68,7 +68,7 @@ def test_openai_embedder_uses_config():
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=ModelConfig(
|
model=ModelConfig(
|
||||||
provider="openai",
|
provider="openai",
|
||||||
model="text-embedding-3-large",
|
name="text-embedding-3-large",
|
||||||
),
|
),
|
||||||
vector_dim=3072,
|
vector_dim=3072,
|
||||||
),
|
),
|
||||||
|
|
@ -90,7 +90,7 @@ def test_voyageai_embedder_uses_config():
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=ModelConfig(
|
model=ModelConfig(
|
||||||
provider="voyageai",
|
provider="voyageai",
|
||||||
model="voyage-large-2",
|
name="voyage-large-2",
|
||||||
),
|
),
|
||||||
vector_dim=1536,
|
vector_dim=1536,
|
||||||
),
|
),
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,7 @@ async def test_qa_ollama(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test Ollama QA with LLM judge."""
|
"""Test Ollama QA with LLM judge."""
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
qa = QuestionAnswerAgent(
|
qa = QuestionAnswerAgent(
|
||||||
client, ModelConfig(provider="ollama", model="gpt-oss", enable_thinking=False)
|
client, ModelConfig(provider="ollama", name="gpt-oss", enable_thinking=False)
|
||||||
)
|
)
|
||||||
llm_judge = LLMJudge()
|
llm_judge = LLMJudge()
|
||||||
|
|
||||||
|
|
@ -44,9 +44,7 @@ async def test_qa_ollama(qa_corpus: Dataset, temp_db_path):
|
||||||
async def test_qa_openai(qa_corpus: Dataset, temp_db_path):
|
async def test_qa_openai(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test OpenAI QA with LLM judge."""
|
"""Test OpenAI QA with LLM judge."""
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
qa = QuestionAnswerAgent(
|
qa = QuestionAnswerAgent(client, ModelConfig(provider="openai", name="gpt-4o-mini"))
|
||||||
client, ModelConfig(provider="openai", model="gpt-4o-mini")
|
|
||||||
)
|
|
||||||
llm_judge = LLMJudge()
|
llm_judge = LLMJudge()
|
||||||
|
|
||||||
doc = qa_corpus[1]
|
doc = qa_corpus[1]
|
||||||
|
|
@ -71,7 +69,7 @@ async def test_qa_anthropic(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test Anthropic QA with LLM judge."""
|
"""Test Anthropic QA with LLM judge."""
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
qa = QuestionAnswerAgent(
|
qa = QuestionAnswerAgent(
|
||||||
client, ModelConfig(provider="anthropic", model="claude-3-5-haiku-20241022")
|
client, ModelConfig(provider="anthropic", name="claude-3-5-haiku-20241022")
|
||||||
)
|
)
|
||||||
llm_judge = LLMJudge()
|
llm_judge = LLMJudge()
|
||||||
|
|
||||||
|
|
@ -96,9 +94,7 @@ async def test_qa_anthropic(qa_corpus: Dataset, temp_db_path):
|
||||||
async def test_qa_vllm(qa_corpus: Dataset, temp_db_path):
|
async def test_qa_vllm(qa_corpus: Dataset, temp_db_path):
|
||||||
"""Test vLLM QA with LLM judge."""
|
"""Test vLLM QA with LLM judge."""
|
||||||
client = HaikuRAG(temp_db_path)
|
client = HaikuRAG(temp_db_path)
|
||||||
qa = QuestionAnswerAgent(
|
qa = QuestionAnswerAgent(client, ModelConfig(provider="vllm", name="Qwen/Qwen3-4B"))
|
||||||
client, ModelConfig(provider="vllm", model="Qwen/Qwen3-4B")
|
|
||||||
)
|
|
||||||
llm_judge = LLMJudge()
|
llm_judge = LLMJudge()
|
||||||
|
|
||||||
doc = qa_corpus[1]
|
doc = qa_corpus[1]
|
||||||
|
|
|
||||||
|
|
@ -44,7 +44,7 @@ async def test_mxbai_reranker():
|
||||||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||||
|
|
||||||
Config.reranking.model = ModelConfig(
|
Config.reranking.model = ModelConfig(
|
||||||
provider="mxbai", model="mixedbread-ai/mxbai-rerank-base-v2"
|
provider="mxbai", name="mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
)
|
)
|
||||||
reranker = MxBAIReranker()
|
reranker = MxBAIReranker()
|
||||||
reranked = await reranker.rerank(
|
reranked = await reranker.rerank(
|
||||||
|
|
|
||||||
|
|
@ -136,16 +136,14 @@ Emoji test: 🚀 ✅ 📝"""
|
||||||
|
|
||||||
def test_get_model_ollama():
|
def test_get_model_ollama():
|
||||||
"""Test get_model returns OpenAIChatModel for Ollama."""
|
"""Test get_model returns OpenAIChatModel for Ollama."""
|
||||||
model_config = ModelConfig(provider="ollama", model="llama3")
|
model_config = ModelConfig(provider="ollama", name="llama3")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_ollama_with_thinking():
|
def test_get_model_ollama_without_thinking():
|
||||||
"""Test get_model configures thinking for gpt-oss on Ollama."""
|
"""Test get_model configures thinking for gpt-oss on Ollama."""
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(provider="ollama", name="gpt-oss", enable_thinking=False)
|
||||||
provider="ollama", model="gpt-oss", enable_thinking=False
|
|
||||||
)
|
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
@ -153,7 +151,7 @@ def test_get_model_ollama_with_thinking():
|
||||||
def test_get_model_ollama_with_settings():
|
def test_get_model_ollama_with_settings():
|
||||||
"""Test get_model applies temperature and max_tokens for Ollama."""
|
"""Test get_model applies temperature and max_tokens for Ollama."""
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="ollama", model="llama3", temperature=0.5, max_tokens=100
|
provider="ollama", name="llama3", temperature=0.5, max_tokens=100
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
@ -161,14 +159,14 @@ def test_get_model_ollama_with_settings():
|
||||||
|
|
||||||
def test_get_model_openai():
|
def test_get_model_openai():
|
||||||
"""Test get_model returns OpenAIChatModel for OpenAI."""
|
"""Test get_model returns OpenAIChatModel for OpenAI."""
|
||||||
model_config = ModelConfig(provider="openai", model="gpt-4o")
|
model_config = ModelConfig(provider="openai", name="gpt-4o")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_openai_with_thinking():
|
def test_get_model_openai_with_thinking():
|
||||||
"""Test get_model configures thinking for OpenAI reasoning models."""
|
"""Test get_model configures thinking for OpenAI reasoning models."""
|
||||||
model_config = ModelConfig(provider="openai", model="o1", enable_thinking=True)
|
model_config = ModelConfig(provider="openai", name="o1", enable_thinking=True)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
@ -178,7 +176,7 @@ def test_get_model_anthropic():
|
||||||
"""Test get_model returns AnthropicModel for Anthropic."""
|
"""Test get_model returns AnthropicModel for Anthropic."""
|
||||||
from pydantic_ai.models.anthropic import AnthropicModel
|
from pydantic_ai.models.anthropic import AnthropicModel
|
||||||
|
|
||||||
model_config = ModelConfig(provider="anthropic", model="claude-3-5-sonnet-20241022")
|
model_config = ModelConfig(provider="anthropic", name="claude-3-5-sonnet-20241022")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, AnthropicModel)
|
assert isinstance(result, AnthropicModel)
|
||||||
|
|
||||||
|
|
@ -190,7 +188,7 @@ def test_get_model_anthropic_with_thinking():
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="anthropic",
|
provider="anthropic",
|
||||||
model="claude-3-5-sonnet-20241022",
|
name="claude-3-5-sonnet-20241022",
|
||||||
enable_thinking=True,
|
enable_thinking=True,
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
|
|
@ -202,7 +200,7 @@ def test_get_model_gemini():
|
||||||
"""Test get_model returns GoogleModel for Gemini."""
|
"""Test get_model returns GoogleModel for Gemini."""
|
||||||
from pydantic_ai.models.google import GoogleModel
|
from pydantic_ai.models.google import GoogleModel
|
||||||
|
|
||||||
model_config = ModelConfig(provider="gemini", model="gemini-2.0-flash-exp")
|
model_config = ModelConfig(provider="gemini", name="gemini-2.0-flash-exp")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, GoogleModel)
|
assert isinstance(result, GoogleModel)
|
||||||
|
|
||||||
|
|
@ -213,7 +211,7 @@ def test_get_model_gemini_with_thinking():
|
||||||
from pydantic_ai.models.google import GoogleModel
|
from pydantic_ai.models.google import GoogleModel
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="gemini", model="gemini-2.0-flash-thinking-exp", enable_thinking=True
|
provider="gemini", name="gemini-2.0-flash-thinking-exp", enable_thinking=True
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, GoogleModel)
|
assert isinstance(result, GoogleModel)
|
||||||
|
|
@ -224,7 +222,7 @@ def test_get_model_groq():
|
||||||
"""Test get_model returns GroqModel for Groq."""
|
"""Test get_model returns GroqModel for Groq."""
|
||||||
from pydantic_ai.models.groq import GroqModel
|
from pydantic_ai.models.groq import GroqModel
|
||||||
|
|
||||||
model_config = ModelConfig(provider="groq", model="llama-3.3-70b-versatile")
|
model_config = ModelConfig(provider="groq", name="llama-3.3-70b-versatile")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, GroqModel)
|
assert isinstance(result, GroqModel)
|
||||||
|
|
||||||
|
|
@ -235,7 +233,7 @@ def test_get_model_groq_with_thinking():
|
||||||
from pydantic_ai.models.groq import GroqModel
|
from pydantic_ai.models.groq import GroqModel
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="groq", model="llama-3.3-70b-versatile", enable_thinking=False
|
provider="groq", name="llama-3.3-70b-versatile", enable_thinking=False
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, GroqModel)
|
assert isinstance(result, GroqModel)
|
||||||
|
|
@ -247,7 +245,7 @@ def test_get_model_bedrock():
|
||||||
from pydantic_ai.models.bedrock import BedrockConverseModel
|
from pydantic_ai.models.bedrock import BedrockConverseModel
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="bedrock", model="anthropic.claude-3-5-sonnet-20241022-v2:0"
|
provider="bedrock", name="anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, BedrockConverseModel)
|
assert isinstance(result, BedrockConverseModel)
|
||||||
|
|
@ -260,7 +258,7 @@ def test_get_model_bedrock_with_thinking():
|
||||||
|
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="bedrock",
|
provider="bedrock",
|
||||||
model="anthropic.claude-3-5-sonnet-20241022-v2:0",
|
name="anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||||
enable_thinking=True,
|
enable_thinking=True,
|
||||||
)
|
)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
|
|
@ -269,21 +267,21 @@ def test_get_model_bedrock_with_thinking():
|
||||||
|
|
||||||
def test_get_model_vllm():
|
def test_get_model_vllm():
|
||||||
"""Test get_model returns OpenAIChatModel for vLLM."""
|
"""Test get_model returns OpenAIChatModel for vLLM."""
|
||||||
model_config = ModelConfig(provider="vllm", model="Qwen/Qwen3-4B")
|
model_config = ModelConfig(provider="vllm", name="Qwen/Qwen3-4B")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_vllm_with_thinking():
|
def test_get_model_vllm_with_thinking():
|
||||||
"""Test get_model configures thinking for gpt-oss on vLLM."""
|
"""Test get_model configures thinking for gpt-oss on vLLM."""
|
||||||
model_config = ModelConfig(provider="vllm", model="gpt-oss", enable_thinking=False)
|
model_config = ModelConfig(provider="vllm", name="gpt-oss", enable_thinking=False)
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, OpenAIChatModel)
|
assert isinstance(result, OpenAIChatModel)
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_unknown_provider():
|
def test_get_model_unknown_provider():
|
||||||
"""Test get_model returns string format for unknown providers."""
|
"""Test get_model returns string format for unknown providers."""
|
||||||
model_config = ModelConfig(provider="mistral", model="mistral-large-latest")
|
model_config = ModelConfig(provider="mistral", name="mistral-large-latest")
|
||||||
result = get_model(model_config)
|
result = get_model(model_config)
|
||||||
assert isinstance(result, str)
|
assert isinstance(result, str)
|
||||||
assert result == "mistral:mistral-large-latest"
|
assert result == "mistral:mistral-large-latest"
|
||||||
|
|
@ -293,7 +291,7 @@ def test_get_model_with_all_settings():
|
||||||
"""Test get_model applies all settings together."""
|
"""Test get_model applies all settings together."""
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(
|
||||||
provider="openai",
|
provider="openai",
|
||||||
model="gpt-4o",
|
name="gpt-4o",
|
||||||
enable_thinking=False,
|
enable_thinking=False,
|
||||||
temperature=0.7,
|
temperature=0.7,
|
||||||
max_tokens=500,
|
max_tokens=500,
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue