Delete obsolete embedders, rewrite get_embedder

This commit is contained in:
Yiorgis Gozadinos 2025-12-24 12:37:17 +02:00
parent 3141c91023
commit 5a30c197af
No known key found for this signature in database
6 changed files with 62 additions and 165 deletions

View file

@ -1,13 +1,40 @@
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, overload
from pydantic_ai.embeddings import Embedder
from pydantic_ai.embeddings.openai import OpenAIEmbeddingModel
from pydantic_ai.providers.ollama import OllamaProvider
from pydantic_ai.providers.openai import OpenAIProvider
from haiku.rag.config import AppConfig, Config
from haiku.rag.embeddings.base import EmbedderBase
from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
if TYPE_CHECKING:
from haiku.rag.store.models.chunk import Chunk
class EmbedderWrapper:
"""Wrapper around pydantic-ai Embedder to provide simple embed() interface."""
def __init__(self, embedder: Embedder, vector_dim: int):
self._embedder = embedder
self._vector_dim = vector_dim
@overload
async def embed(self, text: str) -> list[float]: ...
@overload
async def embed(self, text: list[str]) -> list[list[float]]: ...
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
if isinstance(text, str):
result = await self._embedder.embed_query(text)
return list(result.embeddings[0])
else:
if not text:
return []
result = await self._embedder.embed_documents(text)
return [list(e) for e in result.embeddings]
def contextualize(chunks: list["Chunk"]) -> list[str]:
"""Prepare chunk content for embedding by adding context.
@ -71,9 +98,8 @@ async def embed_chunks(
]
def get_embedder(config: AppConfig = Config) -> EmbedderBase:
"""
Factory function to get the appropriate embedder based on the configuration.
def get_embedder(config: AppConfig = Config) -> EmbedderWrapper:
"""Factory function to get the appropriate embedder based on the configuration.
Args:
config: Configuration to use. Defaults to global Config.
@ -82,11 +108,29 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
An embedder instance configured according to the config.
"""
embedding_model = config.embeddings.model
provider = embedding_model.provider
model_name = embedding_model.name
vector_dim = embedding_model.vector_dim
if embedding_model.provider == "ollama":
return OllamaEmbedder(embedding_model.name, embedding_model.vector_dim, config)
if provider == "ollama":
# Use model-level base_url if set, otherwise fall back to providers config
base_url = embedding_model.base_url or f"{config.providers.ollama.base_url}/v1"
model = OpenAIEmbeddingModel(
model_name,
provider=OllamaProvider(base_url=base_url),
)
return EmbedderWrapper(Embedder(model), vector_dim)
if embedding_model.provider == "voyageai":
if provider == "openai":
if embedding_model.base_url:
model = OpenAIEmbeddingModel(
model_name,
provider=OpenAIProvider(base_url=embedding_model.base_url),
)
return EmbedderWrapper(Embedder(model), vector_dim)
return EmbedderWrapper(Embedder(f"openai:{model_name}"), vector_dim)
if provider == "voyageai":
try:
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
except ImportError:
@ -95,25 +139,14 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
"Please install haiku.rag with the 'voyageai' extra: "
"uv pip install haiku.rag[voyageai]"
)
return VoyageAIEmbedder(
embedding_model.name, embedding_model.vector_dim, config
return VoyageAIEmbedder(model_name, vector_dim, config) # type: ignore[return-value]
if provider == "cohere":
return EmbedderWrapper(Embedder(f"cohere:{model_name}"), vector_dim)
if provider == "sentence-transformers":
return EmbedderWrapper(
Embedder(f"sentence-transformers:{model_name}"), vector_dim
)
if embedding_model.provider == "openai":
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
return OpenAIEmbedder(embedding_model.name, embedding_model.vector_dim, config)
if embedding_model.provider == "vllm":
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
return VllmEmbedder(embedding_model.name, embedding_model.vector_dim, config)
if embedding_model.provider == "lm_studio":
from haiku.rag.embeddings.lm_studio import Embedder as LMStudioEmbedder
return LMStudioEmbedder(
embedding_model.name, embedding_model.vector_dim, config
)
raise ValueError(f"Unsupported embedding provider: {embedding_model.provider}")
raise ValueError(f"Unsupported embedding provider: {provider}")

View file

@ -1,25 +0,0 @@
from typing import overload
from haiku.rag.config import AppConfig, Config
class EmbedderBase:
_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):
self._model = model
self._vector_dim = vector_dim
self._config = config
@overload
async def embed(self, text: str) -> list[float]: ...
@overload
async def embed(self, text: list[str]) -> list[list[float]]: ...
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
raise NotImplementedError(
"Embedder is an abstract class. Please implement the embed method in a subclass."
)

View file

@ -1,28 +0,0 @@
from typing import overload
from openai import AsyncOpenAI
from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase): # pragma: no cover
@overload
async def embed(self, text: str) -> list[float]: ...
@overload
async def embed(self, text: list[str]) -> list[list[float]]: ...
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncOpenAI(
base_url=f"{self._config.providers.lm_studio.base_url}/v1", api_key="dummy"
)
if not text:
return []
response = await client.embeddings.create(
model=self._model,
input=text,
)
if isinstance(text, str):
return response.data[0].embedding
else:
return [item.embedding for item in response.data]

View file

@ -1,28 +0,0 @@
from typing import overload
from openai import AsyncOpenAI
from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase):
@overload
async def embed(self, text: str) -> list[float]: ...
@overload
async def embed(self, text: list[str]) -> list[list[float]]: ...
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncOpenAI(
base_url=f"{self._config.providers.ollama.base_url}/v1", api_key="dummy"
)
if not text:
return []
response = await client.embeddings.create(
model=self._model,
input=text,
)
if isinstance(text, str):
return response.data[0].embedding
else:
return [item.embedding for item in response.data]

View file

@ -1,26 +0,0 @@
from typing import overload
from openai import AsyncOpenAI
from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase):
@overload
async def embed(self, text: str) -> list[float]: ...
@overload
async def embed(self, text: list[str]) -> list[list[float]]: ...
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncOpenAI()
if not text:
return []
response = await client.embeddings.create(
model=self._model,
input=text,
)
if isinstance(text, str):
return response.data[0].embedding
else:
return [item.embedding for item in response.data]

View file

@ -1,29 +0,0 @@
from typing import overload
from openai import AsyncOpenAI
from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase): # pragma: no cover
@overload
async def embed(self, text: str) -> list[float]: ...
@overload
async def embed(self, text: list[str]) -> list[list[float]]: ...
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncOpenAI(
base_url=f"{self._config.providers.vllm.embeddings_base_url}/v1",
api_key="dummy",
)
if not text:
return []
response = await client.embeddings.create(
model=self._model,
input=text,
)
if isinstance(text, str):
return response.data[0].embedding
else:
return [item.embedding for item in response.data]