-
Notifications
You must be signed in to change notification settings - Fork 4
Expand file tree
/
Copy pathbackends.py
More file actions
130 lines (92 loc) · 3.58 KB
/
Copy pathbackends.py
File metadata and controls
130 lines (92 loc) · 3.58 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
import os
import logging
from sentence_transformers import SentenceTransformer
from openai import OpenAI
import tiktoken
from ckan.plugins import toolkit
log = logging.getLogger(__name__)
class BaseEmbeddingsBackend:
def get_dataset_values(self, dataset_dict):
if dataset_dict.get("notes"):
return dataset_dict["title"] + " " + dataset_dict["notes"]
else:
return dataset_dict["title"]
#return dataset_dict["title"]
def get_embedding_for_dataset(self, dataset_dict):
return self.create_embedding(self.get_dataset_values(dataset_dict))
def get_embedding_for_string(self, value):
return self.create_embedding(value)
def create_embedding(self, values):
raise NotImplemented
class SentenceTransformerBackend(BaseEmbeddingsBackend):
model = None
def __init__(self):
# TODO: config model
self.model = SentenceTransformer("all-MiniLM-L6-v2")
# self.model = SentenceTransformer("distiluse-base-multilingual-cased-v1")
def _check_input_length(self, values):
num_input = len(
self.model[0]
.tokenizer(values, return_attention_mask=False, return_token_type_ids=False)
.input_ids
)
max_input = self.model.max_seq_length
if num_input > max_input:
log.debug(
f"Too many input values, input will be truncated ({num_input} vs {max_input})"
)
def create_embedding(self, values):
self._check_input_length(values)
return self.model.encode(values)
class OpenAIBackend(BaseEmbeddingsBackend):
client = None
def __init__(self):
# TODO: config declaration
api_key = toolkit.config.get(
"ckanext.embeddings.openai.api_key", os.environ.get("OPENAI_API_KEY")
)
self.client = OpenAI(api_key=api_key)
def _check_input_length(self, values):
# TODO: configure
encoding = tiktoken.get_encoding("cl100k_base")
max_input = 8191
num_tokens = len(encoding.encode(values))
log.debug(f"[OpenAI Embeddings API] Input size: {num_tokens}")
if num_tokens > max_input:
log.debug(
f"Too many input values, input will be truncated ({num_tokens} vs {max_input})"
)
def create_embedding(self, values):
self._check_input_length(values)
# TODO: config model
response = self.client.embeddings.create(
input=values, model="text-embedding-ada-002"
)
embeddings = [v.embedding for v in response.data]
return embeddings[0]
embeddings_backends = {}
def _load_embeddings_backends():
from importlib.metadata import entry_points
try:
eps = entry_points(group="ckanext.embeddings.backends")
except:
# python 3.9/3.8
eps = (ep for ep in entry_points()['ckanext.embeddings.backends'])
for ep in eps:
embeddings_backends[ep.name] = ep.load()
log.debug(f"Registering Embeddings Backend: {ep.name}")
_embeddings_backend = None
def get_embeddings_backend():
# TODO: config declaration
global _embeddings_backend
backend = toolkit.config.get("ckanext.embeddings.backend", "sentence_transformers")
log.debug(f"Using Embeddings Backend: {backend}")
import time
start = time.time()
try:
_load_embeddings_backends()
if _embeddings_backend is None:
_embeddings_backend = embeddings_backends[backend]()
return _embeddings_backend
finally:
log.debug("loading embeddings took: %.3f sec", time.time()-start)