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 vector_dim: 768
qa: qa:
provider: ollama model:
model: qwen3 provider: ollama
name: qwen3
``` ```
See [Configuration docs](https://ggozad.github.io/haiku.rag/configuration/) for all available options. 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 environment: production
embeddings: embeddings:
model: provider: ollama
provider: ollama model: qwen3-embedding:4b
name: qwen3-embedding:4b
vector_dim: 2560 vector_dim: 2560
qa: qa:
@ -70,9 +69,8 @@ lancedb:
region: "" region: ""
embeddings: embeddings:
model: provider: ollama
provider: ollama model: qwen3-embedding:4b
name: qwen3-embedding:4b
vector_dim: 2560 vector_dim: 2560
reranking: reranking:

View file

@ -72,9 +72,8 @@ If you use Ollama, you can use any pulled model that supports embeddings.
```yaml ```yaml
embeddings: embeddings:
model: provider: ollama
provider: ollama model: mxbai-embed-large
name: mxbai-embed-large
vector_dim: 1024 vector_dim: 1024
``` ```
@ -104,9 +103,8 @@ uv pip install haiku.rag-slim[voyageai]
```yaml ```yaml
embeddings: embeddings:
model: provider: voyageai
provider: voyageai model: voyage-3.5
name: voyage-3.5
vector_dim: 1024 vector_dim: 1024
``` ```
@ -122,9 +120,8 @@ OpenAI embeddings are included in the default installation:
```yaml ```yaml
embeddings: embeddings:
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
``` ```
@ -140,9 +137,8 @@ For high-performance local inference, you can use vLLM to serve embedding models
```yaml ```yaml
embeddings: embeddings:
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:

View file

@ -38,8 +38,9 @@ embeddings:
vector_dim: 1536 vector_dim: 1536
qa: qa:
provider: openai model:
model: gpt-4o-mini # or gpt-4o, gpt-4, etc. 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): 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 { return {
"dataset": dataset_key, "dataset": dataset_key,
"test_cases": test_cases, "test_cases": test_cases,
"embedder_provider": config.embeddings.model.provider, "embedder_provider": config.embeddings.provider,
"embedder_model": config.embeddings.model.name, "embedder_model": config.embeddings.model,
"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,

View file

@ -44,12 +44,8 @@ class LanceDBConfig(BaseModel):
class EmbeddingsConfig(BaseModel): class EmbeddingsConfig(BaseModel):
model: ModelConfig = Field( provider: str = "ollama"
default_factory=lambda: ModelConfig( model: str = "qwen3-embedding:4b"
provider="ollama",
name="qwen3-embedding:4b",
)
)
vector_dim: int = 2560 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. An embedder instance configured according to the config.
""" """
if config.embeddings.model.provider == "ollama": if config.embeddings.provider == "ollama":
return OllamaEmbedder( 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: try:
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
except ImportError: except ImportError:
@ -29,23 +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.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 from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
return 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 from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
return VllmEmbedder( return VllmEmbedder(
config.embeddings.model.name, config.embeddings.vector_dim, config config.embeddings.model, config.embeddings.vector_dim, config
) )
raise ValueError( raise ValueError(f"Unsupported embedding provider: {config.embeddings.provider}")
f"Unsupported embedding provider: {config.embeddings.model.provider}"
)

View file

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

View file

@ -365,8 +365,8 @@ 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.provider == "ollama":
required_models.add(Config.embeddings.model.name) required_models.add(Config.embeddings.model)
if Config.qa.model.provider == "ollama": if Config.qa.model.provider == "ollama":
required_models.add(Config.qa.model.name) required_models.add(Config.qa.model.name)
if Config.research.model.provider == "ollama": if Config.research.model.provider == "ollama":

View file

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