Merge pull request #48 from ggozad/feat/batch_embedding

Support batch embeddings
This commit is contained in:
Yiorgis Gozadinos 2025-09-05 10:52:18 +03:00 committed by GitHub
commit 3b38d9b327
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 58 additions and 24 deletions

View file

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

@ -4,10 +4,13 @@ from haiku.rag.embeddings.base import 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()
response = await client.embeddings.create(
model=self._model,
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):
async def embed(self, text: str) -> list[float]:
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
client = AsyncOpenAI(
base_url=f"{Config.VLLM_EMBEDDINGS_BASE_URL}/v1", api_key="dummy"
)
@ -13,4 +13,7 @@ class Embedder(EmbedderBase):
model=self._model,
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
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()
res = client.embed([text], model=self._model, output_dtype="float")
return res.embeddings[0] # type: ignore[return-value]
if isinstance(text, str):
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:
pass

View file

@ -154,13 +154,7 @@ class ChunkRepository:
"""Create chunks and embeddings for a document from DoclingDocument."""
chunk_texts = await chunker.chunk(document)
# Generate embeddings in parallel for all chunks
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)
embeddings = await self.embedder.embed(chunk_texts)
# Prepare all chunk records for batch insertion
chunk_records = []

View file

@ -28,7 +28,13 @@ async def test_ollama_embedder():
"Python is my favorite programming language.",
"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_embedding = await embedder.embed(test_phrase)
@ -58,7 +64,13 @@ async def test_openai_embedder():
"Python is my favorite programming language.",
"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_embedding = await embedder.embed(test_phrase)
@ -91,7 +103,13 @@ async def test_voyageai_embedder():
"Python is my favorite programming language.",
"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_embedding = await embedder.embed(test_phrase)
@ -126,7 +144,13 @@ async def test_vllm_embedder():
"Python is my favorite programming language.",
"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_embedding = await embedder.embed(test_phrase)