Revert embeddings to use flat config

This commit is contained in:
Yiorgis Gozadinos 2025-11-25 12:51:58 +02:00
parent ba4dddeeaf
commit e83978f88b
No known key found for this signature in database
10 changed files with 42 additions and 64 deletions

View file

@ -30,8 +30,9 @@ embeddings:
vector_dim: 768
qa:
provider: ollama
model: qwen3
model:
provider: ollama
name: qwen3
```
See [Configuration docs](https://ggozad.github.io/haiku.rag/configuration/) for all available options.

View file

@ -35,9 +35,8 @@ A minimal configuration file with defaults:
environment: production
embeddings:
model:
provider: ollama
name: qwen3-embedding:4b
provider: ollama
model: qwen3-embedding:4b
vector_dim: 2560
qa:
@ -70,9 +69,8 @@ lancedb:
region: ""
embeddings:
model:
provider: ollama
name: qwen3-embedding:4b
provider: ollama
model: qwen3-embedding:4b
vector_dim: 2560
reranking:

View file

@ -72,9 +72,8 @@ If you use Ollama, you can use any pulled model that supports embeddings.
```yaml
embeddings:
model:
provider: ollama
name: mxbai-embed-large
provider: ollama
model: mxbai-embed-large
vector_dim: 1024
```
@ -104,9 +103,8 @@ uv pip install haiku.rag-slim[voyageai]
```yaml
embeddings:
model:
provider: voyageai
name: voyage-3.5
provider: voyageai
model: voyage-3.5
vector_dim: 1024
```
@ -122,9 +120,8 @@ OpenAI embeddings are included in the default installation:
```yaml
embeddings:
model:
provider: openai
name: text-embedding-3-small # or text-embedding-3-large
provider: openai
model: text-embedding-3-small # or text-embedding-3-large
vector_dim: 1536
```
@ -140,9 +137,8 @@ For high-performance local inference, you can use vLLM to serve embedding models
```yaml
embeddings:
model:
provider: vllm
name: mixedbread-ai/mxbai-embed-large-v1
provider: vllm
model: mixedbread-ai/mxbai-embed-large-v1
vector_dim: 512
providers:

View file

@ -38,8 +38,9 @@ embeddings:
vector_dim: 1536
qa:
provider: openai
model: gpt-4o-mini # or gpt-4o, gpt-4, etc.
model:
provider: openai
name: gpt-4o-mini # or gpt-4o, gpt-4, etc.
```
Set your OpenAI API key as an environment variable (API keys should not be stored in the YAML file):

View file

@ -43,8 +43,8 @@ def build_experiment_metadata(
return {
"dataset": dataset_key,
"test_cases": test_cases,
"embedder_provider": config.embeddings.model.provider,
"embedder_model": config.embeddings.model.name,
"embedder_provider": config.embeddings.provider,
"embedder_model": config.embeddings.model,
"embedder_dim": config.embeddings.vector_dim,
"chunk_size": config.processing.chunk_size,
"context_chunk_radius": config.processing.context_chunk_radius,

View file

@ -44,12 +44,8 @@ class LanceDBConfig(BaseModel):
class EmbeddingsConfig(BaseModel):
model: ModelConfig = Field(
default_factory=lambda: ModelConfig(
provider="ollama",
name="qwen3-embedding:4b",
)
)
provider: str = "ollama"
model: str = "qwen3-embedding:4b"
vector_dim: int = 2560

View file

@ -14,12 +14,12 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
An embedder instance configured according to the config.
"""
if config.embeddings.model.provider == "ollama":
if config.embeddings.provider == "ollama":
return OllamaEmbedder(
config.embeddings.model.name, config.embeddings.vector_dim, config
config.embeddings.model, config.embeddings.vector_dim, config
)
if config.embeddings.model.provider == "voyageai":
if config.embeddings.provider == "voyageai":
try:
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
except ImportError:
@ -29,23 +29,21 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
"uv pip install haiku.rag[voyageai]"
)
return VoyageAIEmbedder(
config.embeddings.model.name, config.embeddings.vector_dim, config
config.embeddings.model, config.embeddings.vector_dim, config
)
if config.embeddings.model.provider == "openai":
if config.embeddings.provider == "openai":
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
return OpenAIEmbedder(
config.embeddings.model.name, config.embeddings.vector_dim, config
config.embeddings.model, config.embeddings.vector_dim, config
)
if config.embeddings.model.provider == "vllm":
if config.embeddings.provider == "vllm":
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
return VllmEmbedder(
config.embeddings.model.name, config.embeddings.vector_dim, config
config.embeddings.model, config.embeddings.vector_dim, config
)
raise ValueError(
f"Unsupported embedding provider: {config.embeddings.model.provider}"
)
raise ValueError(f"Unsupported embedding provider: {config.embeddings.provider}")

View file

@ -4,7 +4,7 @@ from haiku.rag.config import AppConfig, Config
class EmbedderBase:
_model: str = Config.embeddings.model.name
_model: str = Config.embeddings.model
_vector_dim: int = Config.embeddings.vector_dim
_config: AppConfig = Config

View file

@ -365,8 +365,8 @@ def prefetch_models():
# Collect Ollama models from config
required_models: set[str] = set()
if Config.embeddings.model.provider == "ollama":
required_models.add(Config.embeddings.model.name)
if Config.embeddings.provider == "ollama":
required_models.add(Config.embeddings.model)
if Config.qa.model.provider == "ollama":
required_models.add(Config.qa.model.name)
if Config.research.model.provider == "ollama":

View file

@ -7,7 +7,6 @@ from haiku.rag.config import (
ProvidersConfig,
VLLMConfig,
)
from haiku.rag.config.models import ModelConfig
from haiku.rag.embeddings import get_embedder
@ -15,10 +14,8 @@ def test_embedder_uses_config_from_get_embedder():
"""Test that embedders use the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
model=ModelConfig(
provider="ollama",
name="custom-model",
),
provider="ollama",
model="custom-model",
vector_dim=512,
),
providers=ProvidersConfig(
@ -36,14 +33,10 @@ def test_embedder_uses_config_from_get_embedder():
def test_vllm_embedder_uses_config():
"""Test that vllm embedder uses the config passed to get_embedder."""
from haiku.rag.config.models import ModelConfig
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
model=ModelConfig(
provider="vllm",
name="custom-vllm-model",
),
provider="vllm",
model="custom-vllm-model",
vector_dim=768,
),
providers=ProvidersConfig(
@ -62,14 +55,11 @@ def test_vllm_embedder_uses_config():
def test_openai_embedder_uses_config():
"""Test that openai embedder uses the config passed to get_embedder."""
from haiku.rag.config.models import ModelConfig
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
model=ModelConfig(
provider="openai",
name="text-embedding-3-large",
),
provider="openai",
model="text-embedding-3-large",
vector_dim=3072,
),
)
@ -88,10 +78,8 @@ def test_voyageai_embedder_uses_config():
"""Test that voyageai embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
model=ModelConfig(
provider="voyageai",
name="voyage-large-2",
),
provider="voyageai",
model="voyage-large-2",
vector_dim=1536,
),
)