Revert embeddings to use flat config
This commit is contained in:
parent
ba4dddeeaf
commit
e83978f88b
10 changed files with 42 additions and 64 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
Loading…
Reference in a new issue