Differentiate between embedding a query and a document

This commit is contained in:
Yiorgis Gozadinos 2025-12-24 12:43:11 +02:00
parent 5a30c197af
commit 132b8a36bc
No known key found for this signature in database
3 changed files with 37 additions and 35 deletions

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, overload from typing import TYPE_CHECKING
from pydantic_ai.embeddings import Embedder from pydantic_ai.embeddings import Embedder
from pydantic_ai.embeddings.openai import OpenAIEmbeddingModel from pydantic_ai.embeddings.openai import OpenAIEmbeddingModel
@ -12,26 +12,22 @@ if TYPE_CHECKING:
class EmbedderWrapper: class EmbedderWrapper:
"""Wrapper around pydantic-ai Embedder to provide simple embed() interface.""" """Wrapper around pydantic-ai Embedder with explicit query/document methods."""
def __init__(self, embedder: Embedder, vector_dim: int): def __init__(self, embedder: Embedder, vector_dim: int):
self._embedder = embedder self._embedder = embedder
self._vector_dim = vector_dim self._vector_dim = vector_dim
@overload async def embed_query(self, text: str) -> list[float]:
async def embed(self, text: str) -> list[float]: ... """Embed a search query."""
@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) result = await self._embedder.embed_query(text)
return list(result.embeddings[0]) return list(result.embeddings[0])
else:
if not text: async def embed_documents(self, texts: list[str]) -> list[list[float]]:
"""Embed documents/chunks for indexing."""
if not texts:
return [] return []
result = await self._embedder.embed_documents(text) result = await self._embedder.embed_documents(texts)
return [list(e) for e in result.embeddings] return [list(e) for e in result.embeddings]
@ -80,7 +76,7 @@ async def embed_chunks(
embedder = get_embedder(config) embedder = get_embedder(config)
texts = contextualize(chunks) texts = contextualize(chunks)
embeddings = await embedder.embed(texts) embeddings = await embedder.embed_documents(texts)
return [ return [
Chunk( Chunk(

View file

@ -1,26 +1,32 @@
try: try:
from typing import overload
from voyageai.client import Client # type: ignore from voyageai.client import Client # type: ignore
from haiku.rag.embeddings.base import EmbedderBase from haiku.rag.config import AppConfig
class Embedder(EmbedderBase): class Embedder:
@overload """VoyageAI embedder with explicit query/document methods."""
async def embed(self, text: str) -> list[float]: ...
@overload def __init__(self, model: str, vector_dim: int, config: AppConfig):
async def embed(self, text: list[str]) -> list[list[float]]: ... self._model = model
self._vector_dim = vector_dim
self._config = config
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]: async def embed_query(self, text: str) -> list[float]:
"""Embed a search query."""
client = Client() client = Client()
if not text: res = client.embed(
return [] [text], model=self._model, input_type="query", output_dtype="float"
if isinstance(text, str): )
res = client.embed([text], model=self._model, output_dtype="float")
return res.embeddings[0] # type: ignore[return-value] return res.embeddings[0] # type: ignore[return-value]
else:
res = client.embed(text, model=self._model, output_dtype="float") async def embed_documents(self, texts: list[str]) -> list[list[float]]:
"""Embed documents/chunks for indexing."""
if not texts:
return []
client = Client()
res = client.embed(
texts, model=self._model, input_type="document", output_dtype="float"
)
return res.embeddings # type: ignore[return-value] return res.embeddings # type: ignore[return-value]
except ImportError: except ImportError:

View file

@ -245,7 +245,7 @@ class ChunkRepository:
# Prepare search query based on search type # Prepare search query based on search type
if search_type == "vector": if search_type == "vector":
query_embedding = await self.embedder.embed(query) query_embedding = await self.embedder.embed_query(query)
vector_query = cast( vector_query = cast(
"LanceVectorQueryBuilder", "LanceVectorQueryBuilder",
self.store.chunks_table.search( self.store.chunks_table.search(
@ -260,7 +260,7 @@ class ChunkRepository:
results = self.store.chunks_table.search(query, query_type="fts") results = self.store.chunks_table.search(query, query_type="fts")
else: # hybrid (default) else: # hybrid (default)
query_embedding = await self.embedder.embed(query) query_embedding = await self.embedder.embed_query(query)
# Create RRF reranker # Create RRF reranker
reranker = RRFReranker() reranker = RRFReranker()
# Perform native hybrid search with RRF reranking # Perform native hybrid search with RRF reranking