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 { return {
"dataset": dataset_key, "dataset": dataset_key,
"test_cases": test_cases, "test_cases": test_cases,
"embedder_provider": config.embeddings.provider, "embedder_provider": config.embeddings.model.provider,
"embedder_model": config.embeddings.model, "embedder_model": config.embeddings.model.name,
"embedder_dim": config.embeddings.vector_dim, "embedder_dim": config.embeddings.model.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

View file

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

View file

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

View file

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

View file

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

View file

@ -25,6 +25,20 @@ class ModelConfig(BaseModel):
max_tokens: int | None = None 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): class StorageConfig(BaseModel):
data_dir: Path = Field(default_factory=get_default_data_dir) data_dir: Path = Field(default_factory=get_default_data_dir)
vacuum_retention_seconds: int = 86400 vacuum_retention_seconds: int = 86400
@ -44,9 +58,7 @@ class LanceDBConfig(BaseModel):
class EmbeddingsConfig(BaseModel): class EmbeddingsConfig(BaseModel):
provider: str = "ollama" model: EmbeddingModelConfig = Field(default_factory=EmbeddingModelConfig)
model: str = "qwen3-embedding:4b"
vector_dim: int = 2560
class RerankingConfig(BaseModel): class RerankingConfig(BaseModel):

View file

@ -13,13 +13,12 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
Returns: Returns:
An embedder instance configured according to the config. An embedder instance configured according to the config.
""" """
embedding_model = config.embeddings.model
if config.embeddings.provider == "ollama": if embedding_model.provider == "ollama":
return OllamaEmbedder( return OllamaEmbedder(embedding_model.name, embedding_model.vector_dim, config)
config.embeddings.model, config.embeddings.vector_dim, config
)
if config.embeddings.provider == "voyageai": if embedding_model.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,28 +28,24 @@ 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, 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 from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
return OpenAIEmbedder( return OpenAIEmbedder(embedding_model.name, embedding_model.vector_dim, config)
config.embeddings.model, config.embeddings.vector_dim, config
)
if config.embeddings.provider == "vllm": if embedding_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(embedding_model.name, embedding_model.vector_dim, config)
config.embeddings.model, config.embeddings.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 from haiku.rag.embeddings.lm_studio import Embedder as LMStudioEmbedder
return 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: class EmbedderBase:
_model: str = Config.embeddings.model _model: str = Config.embeddings.model.name
_vector_dim: int = Config.embeddings.vector_dim _vector_dim: int = Config.embeddings.model.vector_dim
_config: AppConfig = Config _config: AppConfig = Config
def __init__(self, model: str, vector_dim: int, 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") current_config = self.store._config.model_dump(mode="json")
# Check if embedding provider or model has changed # 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", {}) stored_embeddings = stored_settings.get("embeddings", {})
current_embeddings = current_config.get("embeddings", {}) current_embeddings = current_config.get("embeddings", {})
# Try nested structure first, fall back to flat for old databases stored_model_obj = stored_embeddings.get("model", {})
stored_provider = stored_embeddings.get("provider") or stored_settings.get( current_model_obj = current_embeddings.get("model", {})
"EMBEDDINGS_PROVIDER"
)
current_provider = current_embeddings.get("provider")
stored_model = stored_embeddings.get("model") or stored_settings.get( stored_provider = stored_model_obj.get("provider")
"EMBEDDINGS_MODEL" current_provider = current_model_obj.get("provider")
)
current_model = current_embeddings.get("model")
stored_vector_dim = stored_embeddings.get("vector_dim") or stored_settings.get( stored_model = stored_model_obj.get("name")
"EMBEDDINGS_VECTOR_DIM" current_model = current_model_obj.get("name")
)
current_vector_dim = current_embeddings.get("vector_dim") stored_vector_dim = stored_model_obj.get("vector_dim")
current_vector_dim = current_model_obj.get("vector_dim")
# Check for incompatible changes # Check for incompatible changes
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_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_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_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_order)
upgrades.append(upgrade_0_9_3_fts) upgrades.append(upgrade_0_9_3_fts)
upgrades.append(upgrade_0_10_1_add_title) 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 # Collect Ollama models from config
required_models: set[str] = set() required_models: set[str] = set()
if Config.embeddings.provider == "ollama": if Config.embeddings.model.provider == "ollama":
required_models.add(Config.embeddings.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.name) required_models.add(Config.qa.model.name)
if Config.research.model.provider == "ollama": if Config.research.model.provider == "ollama":

View file

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

View file

@ -2,7 +2,7 @@
name = "haiku.rag" name = "haiku.rag"
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling" 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" }] authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
license = { text = "MIT" } license = { text = "MIT" }
readme = { file = "README.md", content-type = "text/markdown" } readme = { file = "README.md", content-type = "text/markdown" }
@ -22,7 +22,7 @@ classifiers = [
] ]
dependencies = [ 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] [project.scripts]

View file

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

View file

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

View file

@ -2,6 +2,7 @@ import pytest
from haiku.rag.config import ( from haiku.rag.config import (
AppConfig, AppConfig,
EmbeddingModelConfig,
EmbeddingsConfig, EmbeddingsConfig,
LMStudioConfig, LMStudioConfig,
OllamaConfig, OllamaConfig,
@ -15,9 +16,9 @@ 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(
provider="ollama", model=EmbeddingModelConfig(
model="custom-model", provider="ollama", name="custom-model", vector_dim=512
vector_dim=512, ),
), ),
providers=ProvidersConfig( providers=ProvidersConfig(
ollama=OllamaConfig(base_url="http://custom-ollama:8080"), 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.""" """Test that vllm embedder uses the config passed to get_embedder."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
provider="vllm", model=EmbeddingModelConfig(
model="custom-vllm-model", provider="vllm", name="custom-vllm-model", vector_dim=768
vector_dim=768, ),
), ),
providers=ProvidersConfig( providers=ProvidersConfig(
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9001"), 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(): 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."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
provider="openai", model=EmbeddingModelConfig(
model="text-embedding-3-large", provider="openai", name="text-embedding-3-large", vector_dim=3072
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.""" """Test that lm_studio embedder uses the config passed to get_embedder."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
provider="lm_studio", model=EmbeddingModelConfig(
model="custom-lm-studio-model", provider="lm_studio", name="custom-lm-studio-model", vector_dim=1024
vector_dim=1024, ),
), ),
providers=ProvidersConfig( providers=ProvidersConfig(
lm_studio=LMStudioConfig(base_url="http://custom-lmstudio:5678"), 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.""" """Test that voyageai embedder uses the config passed to get_embedder."""
custom_config = AppConfig( custom_config = AppConfig(
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
provider="voyageai", model=EmbeddingModelConfig(
model="voyage-large-2", provider="voyageai", name="voyage-large-2", vector_dim=1536
vector_dim=1536, ),
), ),
) )

View file

@ -125,7 +125,7 @@ async def test_app_info_with_vector_index(temp_db_path, capsys):
[ [
SettingsRecord( SettingsRecord(
id="settings", 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]] [[package]]
name = "haiku-rag" name = "haiku-rag"
version = "0.19.5" version = "0.19.6"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "haiku-rag-slim", extra = ["cohere", "docling", "inspector", "mxbai", "voyageai", "zeroentropy"] }, { name = "haiku-rag-slim", extra = ["cohere", "docling", "inspector", "mxbai", "voyageai", "zeroentropy"] },
@ -1312,7 +1312,7 @@ dev = [
[[package]] [[package]]
name = "haiku-rag-evals" name = "haiku-rag-evals"
version = "0.19.5" version = "0.19.6"
source = { editable = "evaluations" } source = { editable = "evaluations" }
dependencies = [ dependencies = [
{ name = "datasets" }, { name = "datasets" },
@ -1333,7 +1333,7 @@ requires-dist = [
[[package]] [[package]]
name = "haiku-rag-slim" name = "haiku-rag-slim"
version = "0.19.5" version = "0.19.6"
source = { editable = "haiku_rag_slim" } source = { editable = "haiku_rag_slim" }
dependencies = [ dependencies = [
{ name = "docling-core" }, { name = "docling-core" },