Skip to content

Commit 210fd32

Browse files
author
Nick Vigilante
committed
Add support for Ollama embedding
Fixes DOC-821
1 parent 9fefadb commit 210fd32

2 files changed

Lines changed: 51 additions & 0 deletions

File tree

apis/python/src/tiledb/vector_search/embeddings/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from .image_resnetv2_embedding import ImageResNetV2Embedding
33
from .langchain_embedding import LangChainEmbedding
44
from .object_embedding import ObjectEmbedding
5+
from .ollama_embedding import OllamaEmbedding
56
from .random_embedding import RandomEmbedding
67
from .sentence_transformers_embedding import SentenceTransformersEmbedding
78
from .soma_geneptw_embedding import SomaGenePTwEmbedding
@@ -18,4 +19,5 @@
1819
"LangChainEmbedding",
1920
"SomaScGPTEmbedding",
2021
"SomaSCVIEmbedding",
22+
"OllamaEmbedding",
2123
]
Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,49 @@
1+
from typing import Dict, Optional, OrderedDict, Sequence, Union
2+
3+
import numpy as np
4+
5+
# from tiledb.vector_search.embeddings import ObjectEmbedding
6+
7+
8+
class OllamaEmbedding:
9+
"""
10+
Embedding functions from Ollama.
11+
12+
This attempts to import the embedding_class from the ollama module.
13+
"""
14+
15+
def __init__(
16+
self,
17+
dimensions: int,
18+
embedding_class: str = "embed", # really it's the method
19+
embedding_kwargs: Optional[Dict] = None,
20+
):
21+
self.dim_num = dimensions
22+
self.embedding_class = embedding_class
23+
self.embedding_kwargs = embedding_kwargs
24+
25+
def init_kwargs(self) -> Dict:
26+
return {
27+
"dimensions": self.dim_num,
28+
"embedding_class": self.embedding_class,
29+
"embedding_kwargs": self.embedding_kwargs,
30+
}
31+
32+
def dimensions(self) -> int:
33+
return self.dim_num
34+
35+
def vector_type(self) -> np.dtype:
36+
return np.float32
37+
38+
def load(self) -> None:
39+
import importlib
40+
41+
try:
42+
embeddings_module = importlib.import_module("ollama")
43+
embedding_method_ = getattr(embeddings_module, self.embedding_class)
44+
self.embedding = embedding_method_(**self.embedding_kwargs)
45+
except ImportError as e:
46+
print(e)
47+
48+
def embed(self, objects: Union[str, Sequence[str]]) -> np.ndarray:
49+
return np.array(self.embedding(input=objects).embeddings, dtype=np.float32)

0 commit comments

Comments
 (0)