Support batch embedding
This commit is contained in:
parent
b102b6dde3
commit
b99f88f104
7 changed files with 58 additions and 24 deletions
|
|
@ -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."
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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]
|
||||||
|
|
|
||||||
|
|
@ -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]
|
||||||
|
|
|
||||||
|
|
@ -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]
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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 = []
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue