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
|
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.
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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):
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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}"
|
|
||||||
)
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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":
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue