Support batch embedding

This commit is contained in:
Yiorgis Gozadinos 2025-09-05 10:48:31 +03:00
parent b102b6dde3
commit b99f88f104
No known key found for this signature in database
7 changed files with 58 additions and 24 deletions

View file

@ -9,7 +9,7 @@ class EmbedderBase:
self._model = model self._model = model
self._vector_dim = vector_dim self._vector_dim = vector_dim
async def embed(self, text: str) -> list[float]: async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
raise NotImplementedError( raise NotImplementedError(
"Embedder is an abstract class. Please implement the embed method in a subclass." "Embedder is an abstract class. Please implement the embed method in a subclass."
) )

View file

@ -1,11 +1,17 @@
from ollama import AsyncClient from openai import AsyncOpenAI
from haiku.rag.config import Config from haiku.rag.config import Config
from haiku.rag.embeddings.base import EmbedderBase from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase): class Embedder(EmbedderBase):
async def embed(self, text: str) -> list[float]: async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncClient(host=Config.OLLAMA_BASE_URL) client = AsyncOpenAI(base_url=f"{Config.OLLAMA_BASE_URL}/v1", api_key="dummy")
res = await client.embeddings(model=self._model, prompt=text) response = await client.embeddings.create(
return list(res["embedding"]) 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

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

View file

@ -5,7 +5,7 @@ from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase): class Embedder(EmbedderBase):
async def embed(self, text: str) -> list[float]: async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncOpenAI( client = AsyncOpenAI(
base_url=f"{Config.VLLM_EMBEDDINGS_BASE_URL}/v1", api_key="dummy" base_url=f"{Config.VLLM_EMBEDDINGS_BASE_URL}/v1", api_key="dummy"
) )
@ -13,4 +13,7 @@ class Embedder(EmbedderBase):
model=self._model, model=self._model,
input=text, input=text,
) )
return response.data[0].embedding if isinstance(text, str):
return response.data[0].embedding
else:
return [item.embedding for item in response.data]

View file

@ -4,10 +4,14 @@ try:
from haiku.rag.embeddings.base import EmbedderBase from haiku.rag.embeddings.base import EmbedderBase
class Embedder(EmbedderBase): class Embedder(EmbedderBase):
async def embed(self, text: str) -> list[float]: async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = Client() client = Client()
res = client.embed([text], model=self._model, output_dtype="float") if isinstance(text, str):
return res.embeddings[0] # type: ignore[return-value] res = client.embed([text], model=self._model, output_dtype="float")
return res.embeddings[0] # type: ignore[return-value]
else:
res = client.embed(text, model=self._model, output_dtype="float")
return res.embeddings # type: ignore[return-value]
except ImportError: except ImportError:
pass pass

View file

@ -154,13 +154,7 @@ class ChunkRepository:
"""Create chunks and embeddings for a document from DoclingDocument.""" """Create chunks and embeddings for a document from DoclingDocument."""
chunk_texts = await chunker.chunk(document) chunk_texts = await chunker.chunk(document)
# Generate embeddings in parallel for all chunks embeddings = await self.embedder.embed(chunk_texts)
embeddings_tasks = []
for chunk_text in chunk_texts:
embeddings_tasks.append(self.embedder.embed(chunk_text))
# Wait for all embeddings to complete
embeddings = await asyncio.gather(*embeddings_tasks)
# Prepare all chunk records for batch insertion # Prepare all chunk records for batch insertion
chunk_records = [] chunk_records = []

View file

@ -28,7 +28,13 @@ async def test_ollama_embedder():
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases]
# Test batch embedding
embeddings = await embedder.embed(phrases)
assert isinstance(embeddings, list)
assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings]
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed(test_phrase)
@ -58,7 +64,13 @@ async def test_openai_embedder():
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases]
# Test batch embedding
embeddings = await embedder.embed(phrases)
assert isinstance(embeddings, list)
assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings]
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed(test_phrase)
@ -91,7 +103,13 @@ async def test_voyageai_embedder():
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases]
# Test batch embedding
embeddings = await embedder.embed(phrases)
assert isinstance(embeddings, list)
assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings]
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed(test_phrase)
@ -126,7 +144,13 @@ async def test_vllm_embedder():
"Python is my favorite programming language.", "Python is my favorite programming language.",
"I love to travel and see new places.", "I love to travel and see new places.",
] ]
embeddings = [np.array(await embedder.embed(phrase)) for phrase in phrases]
# Test batch embedding
embeddings = await embedder.embed(phrases)
assert isinstance(embeddings, list)
assert len(embeddings) == 3
assert all(isinstance(emb, list) for emb in embeddings)
embeddings = [np.array(emb) for emb in embeddings]
test_phrase = "I am going for a camping trip." test_phrase = "I am going for a camping trip."
test_embedding = await embedder.embed(test_phrase) test_embedding = await embedder.embed(test_phrase)