Use EmbeddingModelConfig similar to ModelConfig for embeddings

This commit is contained in:
Yiorgis Gozadinos 2025-12-02 11:13:04 +02:00
parent f07590bd5b
commit 3525fae625
No known key found for this signature in database
19 changed files with 165 additions and 89 deletions

View file

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

View file

@ -38,9 +38,7 @@ class LLMJudge:
def __init__(self, model: str = "gpt-oss"):
# Create model using get_model with thinking disabled
model_config = ModelConfig(
provider="ollama", model=model, enable_thinking=False
)
model_config = ModelConfig(provider="ollama", name=model, enable_thinking=False)
model_obj = get_model(model_config, Config)
# Create Pydantic AI agent

View file

@ -2,7 +2,7 @@
name = "haiku.rag-evals"
description = "Benchmarking and evaluation scripts for haiku.rag"
version = "0.19.5"
version = "0.19.6"
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
license = { text = "MIT" }
requires-python = ">=3.12"

View file

@ -80,9 +80,10 @@ class HaikuRAGApp:
data = json.loads(raw) if isinstance(raw, str) else (raw or {})
stored_version = str(data.get("version", stored_version))
embeddings = data.get("embeddings", {})
embed_provider = embeddings.get("provider")
embed_model = embeddings.get("model")
vector_dim = embeddings.get("vector_dim")
embed_model_obj = embeddings.get("model", {})
embed_provider = embed_model_obj.get("provider")
embed_model = embed_model_obj.get("name")
vector_dim = embed_model_obj.get("vector_dim")
# Get comprehensive table statistics
from haiku.rag.store.engine import Store

View file

@ -9,9 +9,11 @@ from haiku.rag.config.models import (
AGUIConfig,
AppConfig,
ConversionOptions,
EmbeddingModelConfig,
EmbeddingsConfig,
LanceDBConfig,
LMStudioConfig,
ModelConfig,
MonitorConfig,
OllamaConfig,
ProcessingConfig,
@ -28,22 +30,24 @@ __all__ = [
"AGUIConfig",
"AppConfig",
"ConversionOptions",
"StorageConfig",
"MonitorConfig",
"LanceDBConfig",
"EmbeddingModelConfig",
"EmbeddingsConfig",
"RerankingConfig",
"QAConfig",
"ResearchConfig",
"ProcessingConfig",
"OllamaConfig",
"LanceDBConfig",
"LMStudioConfig",
"VLLMConfig",
"ModelConfig",
"MonitorConfig",
"OllamaConfig",
"ProcessingConfig",
"ProvidersConfig",
"QAConfig",
"RerankingConfig",
"ResearchConfig",
"StorageConfig",
"VLLMConfig",
"find_config_file",
"load_yaml_config",
"generate_default_config",
"get_config",
"load_yaml_config",
"set_config",
]

View file

@ -25,6 +25,20 @@ class ModelConfig(BaseModel):
max_tokens: int | None = None
class EmbeddingModelConfig(BaseModel):
"""Configuration for an embedding model.
Attributes:
provider: Model provider (ollama, openai, voyageai, vllm, lm_studio)
name: Model name/identifier
vector_dim: Vector dimensions produced by the model
"""
provider: str = "ollama"
name: str = "qwen3-embedding:4b"
vector_dim: int = 2560
class StorageConfig(BaseModel):
data_dir: Path = Field(default_factory=get_default_data_dir)
vacuum_retention_seconds: int = 86400
@ -44,9 +58,7 @@ class LanceDBConfig(BaseModel):
class EmbeddingsConfig(BaseModel):
provider: str = "ollama"
model: str = "qwen3-embedding:4b"
vector_dim: int = 2560
model: EmbeddingModelConfig = Field(default_factory=EmbeddingModelConfig)
class RerankingConfig(BaseModel):

View file

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

View file

@ -4,8 +4,8 @@ from haiku.rag.config import AppConfig, Config
class EmbedderBase:
_model: str = Config.embeddings.model
_vector_dim: int = Config.embeddings.vector_dim
_model: str = Config.embeddings.model.name
_vector_dim: int = Config.embeddings.model.vector_dim
_config: AppConfig = Config
def __init__(self, model: str, vector_dim: int, config: AppConfig = Config):

View file

@ -118,25 +118,21 @@ class SettingsRepository:
current_config = self.store._config.model_dump(mode="json")
# Check if embedding provider or model has changed
# Support both old flat structure and new nested structure for backward compatibility
# Both stored and current use nested structure: embeddings.model.{provider,name,vector_dim}
stored_embeddings = stored_settings.get("embeddings", {})
current_embeddings = current_config.get("embeddings", {})
# Try nested structure first, fall back to flat for old databases
stored_provider = stored_embeddings.get("provider") or stored_settings.get(
"EMBEDDINGS_PROVIDER"
)
current_provider = current_embeddings.get("provider")
stored_model_obj = stored_embeddings.get("model", {})
current_model_obj = current_embeddings.get("model", {})
stored_model = stored_embeddings.get("model") or stored_settings.get(
"EMBEDDINGS_MODEL"
)
current_model = current_embeddings.get("model")
stored_provider = stored_model_obj.get("provider")
current_provider = current_model_obj.get("provider")
stored_vector_dim = stored_embeddings.get("vector_dim") or stored_settings.get(
"EMBEDDINGS_VECTOR_DIM"
)
current_vector_dim = current_embeddings.get("vector_dim")
stored_model = stored_model_obj.get("name")
current_model = current_model_obj.get("name")
stored_vector_dim = stored_model_obj.get("vector_dim")
current_vector_dim = current_model_obj.get("vector_dim")
# Check for incompatible changes
incompatible_changes = []

View file

@ -56,7 +56,11 @@ def run_pending_upgrades(store: Store, from_version: str, to_version: str) -> No
from .v0_9_3 import upgrade_fts_phrase as upgrade_0_9_3_fts # noqa: E402
from .v0_9_3 import upgrade_order as upgrade_0_9_3_order # noqa: E402
from .v0_10_1 import upgrade_add_title as upgrade_0_10_1_add_title # noqa: E402
from .v0_19_6 import ( # noqa: E402
upgrade_embeddings_model_config as upgrade_0_19_6_embeddings,
)
upgrades.append(upgrade_0_9_3_order)
upgrades.append(upgrade_0_9_3_fts)
upgrades.append(upgrade_0_10_1_add_title)
upgrades.append(upgrade_0_19_6_embeddings)

View file

@ -0,0 +1,65 @@
import json
import logging
from haiku.rag.store.engine import SettingsRecord, Store
from haiku.rag.store.upgrades import Upgrade
logger = logging.getLogger(__name__)
def _apply_embeddings_model_config(store: Store) -> None:
"""Migrate embeddings config from flat to nested EmbeddingModelConfig structure."""
results = list(
store.settings_table.search()
.where("id = 'settings'")
.limit(1)
.to_pydantic(SettingsRecord)
)
if not results or not results[0].settings:
return
settings = json.loads(results[0].settings)
embeddings = settings.get("embeddings", {})
# Check if already migrated (model is a dict with nested structure)
if isinstance(embeddings.get("model"), dict):
return
# Migrate from flat structure to nested EmbeddingModelConfig
old_provider = embeddings.get("provider", "ollama")
old_model = embeddings.get("model", "qwen3-embedding:4b")
old_vector_dim = embeddings.get("vector_dim", 2560)
logger.warning(
"Migrating embeddings config to new nested structure: "
"embeddings.{provider,model,vector_dim} -> embeddings.model.{provider,name,vector_dim}"
)
# Create new nested structure
settings["embeddings"] = {
"model": {
"provider": old_provider,
"name": old_model,
"vector_dim": old_vector_dim,
}
}
store.settings_table.update(
where="id = 'settings'",
values={"settings": json.dumps(settings)},
)
logger.warning(
"Embeddings config migrated: provider=%s, name=%s, vector_dim=%d",
old_provider,
old_model,
old_vector_dim,
)
upgrade_embeddings_model_config = Upgrade(
version="0.19.6",
apply=_apply_embeddings_model_config,
description="Migrate embeddings config to nested EmbeddingModelConfig structure",
)

View file

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

View file

@ -2,7 +2,7 @@
name = "haiku.rag-slim"
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling - Minimal dependencies"
version = "0.19.5"
version = "0.19.6"
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
license = { text = "MIT" }
readme = { file = "README.md", content-type = "text/markdown" }

View file

@ -2,7 +2,7 @@
name = "haiku.rag"
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling"
version = "0.19.5"
version = "0.19.6"
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
license = { text = "MIT" }
readme = { file = "README.md", content-type = "text/markdown" }
@ -22,7 +22,7 @@ classifiers = [
]
dependencies = [
"haiku.rag-slim[docling,voyageai,mxbai,cohere,zeroentropy,inspector]==0.19.5",
"haiku.rag-slim[docling,voyageai,mxbai,cohere,zeroentropy,inspector]==0.19.6",
]
[project.scripts]

View file

@ -699,7 +699,7 @@ async def test_client_create_document_with_custom_chunks(temp_db_path):
Chunk(
content="This is the second chunk",
metadata={"custom": "metadata2"},
embedding=[0.1] * Config.embeddings.vector_dim,
embedding=[0.1] * Config.embeddings.model.vector_dim,
order=1,
), # With embedding
Chunk(

View file

@ -15,16 +15,17 @@ def test_load_yaml_config(tmp_path):
config_file.write_text("""
environment: production
embeddings:
provider: ollama
model: test-model
vector_dim: 1024
model:
provider: ollama
name: test-model
vector_dim: 1024
""")
config = load_yaml_config(config_file)
assert config["environment"] == "production"
assert config["embeddings"]["provider"] == "ollama"
assert config["embeddings"]["model"] == "test-model"
assert config["embeddings"]["vector_dim"] == 1024
assert config["embeddings"]["model"]["provider"] == "ollama"
assert config["embeddings"]["model"]["name"] == "test-model"
assert config["embeddings"]["model"]["vector_dim"] == 1024
def test_find_config_file_cwd(tmp_path, monkeypatch):
@ -179,7 +180,7 @@ def test_generate_default_config_completeness():
# Verify config validates successfully
assert config.environment == "production"
assert config.embeddings.provider == "ollama"
assert config.embeddings.model.provider == "ollama"
assert config.qa.model.provider == "ollama"
assert config.research.model.provider == "ollama"
assert config.reranking.model is None

View file

@ -2,6 +2,7 @@ import pytest
from haiku.rag.config import (
AppConfig,
EmbeddingModelConfig,
EmbeddingsConfig,
LMStudioConfig,
OllamaConfig,
@ -15,9 +16,9 @@ def test_embedder_uses_config_from_get_embedder():
"""Test that embedders use the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
provider="ollama",
model="custom-model",
vector_dim=512,
model=EmbeddingModelConfig(
provider="ollama", name="custom-model", vector_dim=512
),
),
providers=ProvidersConfig(
ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
@ -36,9 +37,9 @@ def test_vllm_embedder_uses_config():
"""Test that vllm embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
provider="vllm",
model="custom-vllm-model",
vector_dim=768,
model=EmbeddingModelConfig(
provider="vllm", name="custom-vllm-model", vector_dim=768
),
),
providers=ProvidersConfig(
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9001"),
@ -56,12 +57,11 @@ def test_vllm_embedder_uses_config():
def test_openai_embedder_uses_config():
"""Test that openai embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
provider="openai",
model="text-embedding-3-large",
vector_dim=3072,
model=EmbeddingModelConfig(
provider="openai", name="text-embedding-3-large", vector_dim=3072
),
),
)
@ -76,9 +76,9 @@ def test_lm_studio_embedder_uses_config():
"""Test that lm_studio embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
provider="lm_studio",
model="custom-lm-studio-model",
vector_dim=1024,
model=EmbeddingModelConfig(
provider="lm_studio", name="custom-lm-studio-model", vector_dim=1024
),
),
providers=ProvidersConfig(
lm_studio=LMStudioConfig(base_url="http://custom-lmstudio:5678"),
@ -101,9 +101,9 @@ def test_voyageai_embedder_uses_config():
"""Test that voyageai embedder uses the config passed to get_embedder."""
custom_config = AppConfig(
embeddings=EmbeddingsConfig(
provider="voyageai",
model="voyage-large-2",
vector_dim=1536,
model=EmbeddingModelConfig(
provider="voyageai", name="voyage-large-2", vector_dim=1536
),
),
)

View file

@ -125,7 +125,7 @@ async def test_app_info_with_vector_index(temp_db_path, capsys):
[
SettingsRecord(
id="settings",
settings='{"version": "1.0.0", "embeddings": {"provider": "ollama", "model": "test", "vector_dim": 3}}',
settings='{"version": "1.0.0", "embeddings": {"model": {"provider": "ollama", "name": "test", "vector_dim": 3}}}',
)
]
)

View file

@ -1264,7 +1264,7 @@ wheels = [
[[package]]
name = "haiku-rag"
version = "0.19.5"
version = "0.19.6"
source = { editable = "." }
dependencies = [
{ name = "haiku-rag-slim", extra = ["cohere", "docling", "inspector", "mxbai", "voyageai", "zeroentropy"] },
@ -1312,7 +1312,7 @@ dev = [
[[package]]
name = "haiku-rag-evals"
version = "0.19.5"
version = "0.19.6"
source = { editable = "evaluations" }
dependencies = [
{ name = "datasets" },
@ -1333,7 +1333,7 @@ requires-dist = [
[[package]]
name = "haiku-rag-slim"
version = "0.19.5"
version = "0.19.6"
source = { editable = "haiku_rag_slim" }
dependencies = [
{ name = "docling-core" },