mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-31 18:48:25 +00:00
eb9784d7f2
Other embedding integrations such as Cohere and OpenAI already send requests in batches. We should do that for Ollama too to improve throughput. The Ollama [`.embed` API](https://github.com/ollama/ollama-python/blob/63ca74762284100b2f0ad207bc00fa3d32720fbd/ollama/_client.py#L359-L378) was added in version 0.3.0 (almost a year ago) so I updated the version requirement in pyproject. <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit - **Bug Fixes** - Improved compatibility with newer versions of the "ollama" package by requiring version 0.3.0 or higher. - Enhanced embedding generation to process batches of texts more efficiently and reliably. - **Refactor** - Improved type consistency and clarity for embedding-related methods. <!-- end of auto-generated comment: release notes by coderabbit.ai -->
64 lines
1.9 KiB
Python
64 lines
1.9 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
|
|
|
from functools import cached_property
|
|
from typing import TYPE_CHECKING, List, Optional, Sequence, Union
|
|
|
|
import numpy as np
|
|
|
|
from ..util import attempt_import_or_raise
|
|
from .base import TextEmbeddingFunction
|
|
from .registry import register
|
|
|
|
if TYPE_CHECKING:
|
|
import ollama
|
|
|
|
|
|
@register("ollama")
|
|
class OllamaEmbeddings(TextEmbeddingFunction):
|
|
"""
|
|
An embedding function that uses Ollama
|
|
|
|
https://github.com/ollama/ollama/blob/main/docs/api.md#generate-embeddings
|
|
https://ollama.com/blog/embedding-models
|
|
"""
|
|
|
|
name: str = "nomic-embed-text"
|
|
host: str = "http://localhost:11434"
|
|
options: Optional[dict] = None # type = ollama.Options
|
|
keep_alive: Optional[Union[float, str]] = None
|
|
ollama_client_kwargs: Optional[dict] = {}
|
|
|
|
def ndims(self) -> int:
|
|
return len(self.generate_embeddings(["foo"])[0])
|
|
|
|
def _compute_embedding(self, text: Sequence[str]) -> Sequence[Sequence[float]]:
|
|
response = self._ollama_client.embed(
|
|
model=self.name,
|
|
input=text,
|
|
options=self.options,
|
|
keep_alive=self.keep_alive,
|
|
)
|
|
return response.embeddings
|
|
|
|
def generate_embeddings(
|
|
self, texts: Union[List[str], np.ndarray]
|
|
) -> list[Union[np.array, None]]:
|
|
"""
|
|
Get the embeddings for the given texts
|
|
|
|
Parameters
|
|
----------
|
|
texts: list[str] or np.ndarray (of str)
|
|
The texts to embed
|
|
"""
|
|
# TODO retry, rate limit, token limit
|
|
embeddings = self._compute_embedding(texts)
|
|
return list(embeddings)
|
|
|
|
@cached_property
|
|
def _ollama_client(self) -> "ollama.Client":
|
|
ollama = attempt_import_or_raise("ollama")
|
|
# ToDo explore ollama.AsyncClient
|
|
return ollama.Client(host=self.host, **self.ollama_client_kwargs)
|