feat(embeddings): add native embed_batch to OllamaEmbedding (#5415)
Co-authored-by: Kartik <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -63,3 +63,13 @@ class OllamaEmbedding(EmbeddingBase):
|
||||
if not embeddings:
|
||||
raise ValueError(f"Ollama embed() returned no embeddings for model '{self.config.model}'")
|
||||
return embeddings[0]
|
||||
|
||||
def embed_batch(self, texts, memory_action="add"):
|
||||
"""Embed multiple texts in a single Ollama API call."""
|
||||
if not texts:
|
||||
return []
|
||||
response = self.client.embed(model=self.config.model, input=texts)
|
||||
embeddings = response.get("embeddings") or []
|
||||
if len(embeddings) != len(texts):
|
||||
raise ValueError(f"Ollama embed() returned {len(embeddings)} embeddings for {len(texts)} texts using model '{self.config.model}'")
|
||||
return embeddings
|
||||
|
||||
@@ -60,3 +60,37 @@ def test_embed_empty_response_raises(mock_ollama_client):
|
||||
|
||||
with pytest.raises(ValueError, match="returned no embeddings"):
|
||||
embedder.embed("some text")
|
||||
|
||||
|
||||
def test_embed_batch_single_call(mock_ollama_client):
|
||||
config = BaseEmbedderConfig(model="nomic-embed-text", embedding_dims=512)
|
||||
embedder = OllamaEmbedding(config)
|
||||
|
||||
mock_response = {"embeddings": [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6], [0.7, 0.8, 0.9]]}
|
||||
mock_ollama_client.embed.return_value = mock_response
|
||||
|
||||
texts = ["First text.", "Second text.", "Third text."]
|
||||
embeddings = embedder.embed_batch(texts)
|
||||
|
||||
mock_ollama_client.embed.assert_called_once_with(model="nomic-embed-text", input=texts)
|
||||
assert embeddings == [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6], [0.7, 0.8, 0.9]]
|
||||
|
||||
|
||||
def test_embed_batch_empty_list(mock_ollama_client):
|
||||
config = BaseEmbedderConfig(model="nomic-embed-text", embedding_dims=512)
|
||||
embedder = OllamaEmbedding(config)
|
||||
|
||||
result = embedder.embed_batch([])
|
||||
|
||||
assert result == []
|
||||
mock_ollama_client.embed.assert_not_called()
|
||||
|
||||
|
||||
def test_embed_batch_count_mismatch_raises(mock_ollama_client):
|
||||
config = BaseEmbedderConfig(model="nomic-embed-text", embedding_dims=512)
|
||||
embedder = OllamaEmbedding(config)
|
||||
|
||||
mock_ollama_client.embed.return_value = {"embeddings": [[0.1, 0.2, 0.3]]}
|
||||
|
||||
with pytest.raises(ValueError, match="returned 1 embeddings for 2 texts"):
|
||||
embedder.embed_batch(["first text", "second text"])
|
||||
|
||||
Reference in New Issue
Block a user