feat(chunking): add embedding model registry and ollama config
Why: - Operators need Cloud and Local Embedding Models with stable ids, dimensions, and defaults. Changes: - Add Embedding Model Registry; Ollama client; env defaults for model, Ollama host, and Neighbor Expansion knobs. Impact: - New installs default to text-embedding-3-large; OLLAMA_BASE_URL required for Local provider. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
10
.env.example
10
.env.example
@@ -5,12 +5,18 @@ OPENAI_API_KEY=sk-...
|
|||||||
QDRANT_URL=http://localhost:6333
|
QDRANT_URL=http://localhost:6333
|
||||||
QDRANT_API_KEY=
|
QDRANT_API_KEY=
|
||||||
|
|
||||||
# Fixed models (not configurable per strategy)
|
# Default cloud Embedding Model (registry id / OpenAI model name)
|
||||||
EMBEDDING_MODEL=text-embedding-3-small
|
EMBEDDING_MODEL=text-embedding-3-large
|
||||||
LLM_MODEL=gpt-4o-mini
|
LLM_MODEL=gpt-4o-mini
|
||||||
|
|
||||||
|
# Local Ollama embeddings (Active model switched in Admin)
|
||||||
|
OLLAMA_BASE_URL=http://192.168.10.10:11435
|
||||||
|
|
||||||
# Retrieval defaults
|
# Retrieval defaults
|
||||||
TOP_K=5
|
TOP_K=5
|
||||||
|
# Neighbor Expansion for fixed_size only (ADR-0023); 0 = off
|
||||||
|
NEIGHBOR_PREV=0
|
||||||
|
NEIGHBOR_NEXT=0
|
||||||
|
|
||||||
# LLM generation parameters
|
# LLM generation parameters
|
||||||
TEMPERATURE=0.0
|
TEMPERATURE=0.0
|
||||||
|
|||||||
153
src/chunking/embedding_models.py
Normal file
153
src/chunking/embedding_models.py
Normal file
@@ -0,0 +1,153 @@
|
|||||||
|
"""Embedding Model Registry — static catalog of Cloud and Local models."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class Provider(str, Enum):
|
||||||
|
CLOUD = "cloud"
|
||||||
|
LOCAL = "local"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class EmbeddingModelSpec:
|
||||||
|
"""One entry in the Embedding Model Registry."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
provider: Provider
|
||||||
|
model_name: str
|
||||||
|
dimension: int
|
||||||
|
display_name: str
|
||||||
|
task_prefixes: bool = False
|
||||||
|
# Default Semantic Boundary Detection threshold (Admin may override in SQLite)
|
||||||
|
default_semantic_threshold: float = 0.3
|
||||||
|
|
||||||
|
|
||||||
|
# Stable id for legacy unscoped OpenAI corpora / historical Experiment backfill (1536-d).
|
||||||
|
LEGACY_CLOUD_MODEL_ID = "text-embedding-3-small"
|
||||||
|
# Default Corpus/Boundary when Admin settings are unset (new installs & fallbacks).
|
||||||
|
DEFAULT_CLOUD_MODEL_ID = "text-embedding-3-large"
|
||||||
|
|
||||||
|
_SEMANTIC_THRESHOLD_SETTING_PREFIX = "semantic_threshold:"
|
||||||
|
|
||||||
|
|
||||||
|
def semantic_threshold_setting_key(model_id: str) -> str:
|
||||||
|
return f"{_SEMANTIC_THRESHOLD_SETTING_PREFIX}{model_id}"
|
||||||
|
|
||||||
|
|
||||||
|
def _build_registry() -> dict[str, EmbeddingModelSpec]:
|
||||||
|
# Prefer explicit registry entries over EMBEDDING_MODEL aliasing a different
|
||||||
|
# OpenAI model under a single id (wrong dimension / provenance).
|
||||||
|
cloud_small = EmbeddingModelSpec(
|
||||||
|
id=LEGACY_CLOUD_MODEL_ID,
|
||||||
|
provider=Provider.CLOUD,
|
||||||
|
model_name="text-embedding-3-small",
|
||||||
|
dimension=1536,
|
||||||
|
display_name="OpenAI text-embedding-3-small",
|
||||||
|
task_prefixes=False,
|
||||||
|
default_semantic_threshold=0.3,
|
||||||
|
)
|
||||||
|
cloud_large = EmbeddingModelSpec(
|
||||||
|
id=DEFAULT_CLOUD_MODEL_ID,
|
||||||
|
provider=Provider.CLOUD,
|
||||||
|
model_name="text-embedding-3-large",
|
||||||
|
dimension=3072,
|
||||||
|
display_name="OpenAI text-embedding-3-large",
|
||||||
|
task_prefixes=False,
|
||||||
|
default_semantic_threshold=0.3,
|
||||||
|
)
|
||||||
|
local = EmbeddingModelSpec(
|
||||||
|
id="nomic-embed-text-v2-moe",
|
||||||
|
provider=Provider.LOCAL,
|
||||||
|
model_name="nomic-embed-text-v2-moe:latest",
|
||||||
|
dimension=768,
|
||||||
|
display_name="Ollama nomic-embed-text-v2-moe",
|
||||||
|
task_prefixes=True,
|
||||||
|
default_semantic_threshold=0.6,
|
||||||
|
)
|
||||||
|
return {
|
||||||
|
cloud_small.id: cloud_small,
|
||||||
|
cloud_large.id: cloud_large,
|
||||||
|
local.id: local,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def get_registry() -> dict[str, EmbeddingModelSpec]:
|
||||||
|
"""Return the Embedding Model Registry (built from config)."""
|
||||||
|
return _build_registry()
|
||||||
|
|
||||||
|
|
||||||
|
def get_model(model_id: str) -> EmbeddingModelSpec:
|
||||||
|
"""Look up a registry entry by stable id.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
KeyError: If model_id is not registered.
|
||||||
|
"""
|
||||||
|
registry = get_registry()
|
||||||
|
if model_id not in registry:
|
||||||
|
known = ", ".join(sorted(registry))
|
||||||
|
raise KeyError(f"Unknown Embedding Model '{model_id}'. Registered: {known}")
|
||||||
|
return registry[model_id]
|
||||||
|
|
||||||
|
|
||||||
|
def list_models() -> list[EmbeddingModelSpec]:
|
||||||
|
"""All registered Embedding Models in stable order (default cloud first)."""
|
||||||
|
registry = get_registry()
|
||||||
|
preferred = [
|
||||||
|
DEFAULT_CLOUD_MODEL_ID,
|
||||||
|
LEGACY_CLOUD_MODEL_ID,
|
||||||
|
"nomic-embed-text-v2-moe",
|
||||||
|
]
|
||||||
|
order = preferred + [mid for mid in registry if mid not in preferred]
|
||||||
|
return [registry[mid] for mid in order if mid in registry]
|
||||||
|
|
||||||
|
|
||||||
|
def default_model_id() -> str:
|
||||||
|
"""Default Corpus/Boundary Embedding Model id when Admin settings are unset."""
|
||||||
|
return DEFAULT_CLOUD_MODEL_ID
|
||||||
|
|
||||||
|
|
||||||
|
def get_semantic_threshold(model_id: str) -> float:
|
||||||
|
"""Resolve Semantic Boundary Detection threshold for a model.
|
||||||
|
|
||||||
|
Order: Admin SQLite override → registry default → global SEMANTIC_THRESHOLD.
|
||||||
|
"""
|
||||||
|
from src.storage import sqlite as db
|
||||||
|
|
||||||
|
model = get_model(model_id)
|
||||||
|
stored = db.get_app_setting(semantic_threshold_setting_key(model_id))
|
||||||
|
if stored is not None and stored != "":
|
||||||
|
return float(stored)
|
||||||
|
return model.default_semantic_threshold
|
||||||
|
|
||||||
|
|
||||||
|
def set_semantic_threshold(model_id: str, threshold: float) -> float:
|
||||||
|
"""Persist Admin override for a model's semantic_threshold (0 < t <= 1)."""
|
||||||
|
from src.storage import sqlite as db
|
||||||
|
|
||||||
|
get_model(model_id) # validate registry id
|
||||||
|
if not (0.0 < threshold <= 1.0):
|
||||||
|
raise ValueError("semantic_threshold must be in (0, 1]")
|
||||||
|
db.set_app_setting(semantic_threshold_setting_key(model_id), str(threshold))
|
||||||
|
return threshold
|
||||||
|
|
||||||
|
|
||||||
|
def apply_task_prefixes(
|
||||||
|
texts: list[str],
|
||||||
|
*,
|
||||||
|
model: EmbeddingModelSpec,
|
||||||
|
purpose: str,
|
||||||
|
) -> list[str]:
|
||||||
|
"""Apply Nomic-style task prefixes when the model requires them.
|
||||||
|
|
||||||
|
purpose: \"document\" → search_document; \"query\" → search_query.
|
||||||
|
"""
|
||||||
|
if not model.task_prefixes:
|
||||||
|
return texts
|
||||||
|
if purpose == "query":
|
||||||
|
prefix = "search_query: "
|
||||||
|
else:
|
||||||
|
prefix = "search_document: "
|
||||||
|
return [prefix + t if not t.startswith(prefix) else t for t in texts]
|
||||||
@@ -10,15 +10,21 @@ class Settings(BaseSettings):
|
|||||||
|
|
||||||
# OpenAI
|
# OpenAI
|
||||||
openai_api_key: str
|
openai_api_key: str
|
||||||
embedding_model: str = "text-embedding-3-small"
|
embedding_model: str = "text-embedding-3-large"
|
||||||
llm_model: str = "gpt-4o-mini"
|
llm_model: str = "gpt-4o-mini"
|
||||||
|
|
||||||
|
# Local embeddings (Ollama) — Admin switches models; host stays in config
|
||||||
|
ollama_base_url: str = "http://192.168.10.10:11435"
|
||||||
|
|
||||||
# Qdrant
|
# Qdrant
|
||||||
qdrant_url: str = "http://localhost:6333"
|
qdrant_url: str = "http://localhost:6333"
|
||||||
qdrant_api_key: str | None = None
|
qdrant_api_key: str | None = None
|
||||||
|
|
||||||
# Retrieval
|
# Retrieval
|
||||||
top_k: int = 5
|
top_k: int = 5
|
||||||
|
# Neighbor Expansion for fixed_size (ADR-0023); 0/0 = off
|
||||||
|
neighbor_prev: int = 0
|
||||||
|
neighbor_next: int = 0
|
||||||
|
|
||||||
# LLM generation
|
# LLM generation
|
||||||
temperature: float = 0.0
|
temperature: float = 0.0
|
||||||
|
|||||||
@@ -14,6 +14,13 @@ def get_openai_client() -> OpenAI:
|
|||||||
return OpenAI(api_key=settings.openai_api_key)
|
return OpenAI(api_key=settings.openai_api_key)
|
||||||
|
|
||||||
|
|
||||||
|
@lru_cache()
|
||||||
|
def get_ollama_client() -> OpenAI:
|
||||||
|
"""Return a cached OpenAI-compatible client pointed at Ollama."""
|
||||||
|
base = settings.ollama_base_url.rstrip("/")
|
||||||
|
return OpenAI(base_url=f"{base}/v1", api_key="ollama")
|
||||||
|
|
||||||
|
|
||||||
@lru_cache()
|
@lru_cache()
|
||||||
def get_qdrant_client() -> QdrantClient:
|
def get_qdrant_client() -> QdrantClient:
|
||||||
"""Return a cached Qdrant client singleton."""
|
"""Return a cached Qdrant client singleton."""
|
||||||
|
|||||||
Reference in New Issue
Block a user