Use EmbeddingModelConfig similar to ModelConfig for embeddings
This commit is contained in:
parent
f07590bd5b
commit
3525fae625
19 changed files with 165 additions and 89 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
65
haiku_rag_slim/haiku/rag/store/upgrades/v0_19_6.py
Normal file
65
haiku_rag_slim/haiku/rag/store/upgrades/v0_19_6.py
Normal 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",
|
||||
)
|
||||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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}}}',
|
||||
)
|
||||
]
|
||||
)
|
||||
|
|
|
|||
6
uv.lock
6
uv.lock
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Reference in a new issue