Delete obsolete embedders, rewrite get_embedder
This commit is contained in:
parent
3141c91023
commit
5a30c197af
6 changed files with 62 additions and 165 deletions
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
@ -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]
|
||||
|
|
@ -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]
|
||||
|
|
@ -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]
|
||||
|
|
@ -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]
|
||||
Loading…
Reference in a new issue