diff --git a/src/haiku/rag/embeddings/base.py b/src/haiku/rag/embeddings/base.py index 369a53a6..0edf568f 100644 --- a/src/haiku/rag/embeddings/base.py +++ b/src/haiku/rag/embeddings/base.py @@ -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." ) diff --git a/src/haiku/rag/embeddings/ollama.py b/src/haiku/rag/embeddings/ollama.py index 600afe65..2dbd8ea4 100644 --- a/src/haiku/rag/embeddings/ollama.py +++ b/src/haiku/rag/embeddings/ollama.py @@ -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] diff --git a/src/haiku/rag/embeddings/openai.py b/src/haiku/rag/embeddings/openai.py index 485c97fe..14d9129a 100644 --- a/src/haiku/rag/embeddings/openai.py +++ b/src/haiku/rag/embeddings/openai.py @@ -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] diff --git a/src/haiku/rag/embeddings/vllm.py b/src/haiku/rag/embeddings/vllm.py index 0f9a1aee..cae33398 100644 --- a/src/haiku/rag/embeddings/vllm.py +++ b/src/haiku/rag/embeddings/vllm.py @@ -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] diff --git a/src/haiku/rag/embeddings/voyageai.py b/src/haiku/rag/embeddings/voyageai.py index ac7aa1b6..4c0e7e89 100644 --- a/src/haiku/rag/embeddings/voyageai.py +++ b/src/haiku/rag/embeddings/voyageai.py @@ -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 diff --git a/src/haiku/rag/store/repositories/chunk.py b/src/haiku/rag/store/repositories/chunk.py index 870c697d..391a5809 100644 --- a/src/haiku/rag/store/repositories/chunk.py +++ b/src/haiku/rag/store/repositories/chunk.py @@ -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 = [] diff --git a/tests/test_embedder.py b/tests/test_embedder.py index 984ff889..fbde114a 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -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)