-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathembmodels.py
More file actions
57 lines (41 loc) · 2.14 KB
/
Copy pathembmodels.py
File metadata and controls
57 lines (41 loc) · 2.14 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
from sentence_transformers import SentenceTransformer
from FlagEmbedding import FlagModel
from text2vec import SentenceModel
import torch
def binary_encoding(input_vector):
input_tensor = torch.tensor(input_vector)
threshold = torch.median(input_tensor)
binary_vector = (input_tensor > threshold).type(torch.FloatTensor)
return binary_vector
def get_MokaAI_embedding(text, model_name='moka-ai/m3e-base'):
model = SentenceTransformer(model_name)
embedding = model.encode([text])
return binary_encoding(embedding[0])
def get_BAAI_embedding(text, model_name='BAAI/bge-small-en-v1.5'):
model = FlagModel(model_name)
embedding = model.encode([text])
return binary_encoding(embedding[0])
def get_text2vec_embedding(text, model_name):
model = SentenceModel(model_name)
embedding = model.encode([text])
return binary_encoding(embedding[0])
def get_sentence_embedding(embedding_model, text):
moka_ai_models = ['moka-ai/m3e-base', 'moka-ai/m3e-small', 'moka-ai/m3e-large']
baai_models = ['BAAI/bge-small-zh', 'BAAI/bge-base-zh', 'BAAI/bge-large-zh',
'BAAI/bge-small-zh-v1.5', 'BAAI/bge-base-zh-v1.5', 'BAAI/bge-large-zh-v1.5',
'BAAI/bge-large-zh-noinstruct', 'BAAI/bge-reranker-large', 'BAAI/bge-reranker-base',
'BAAI/bge-small-en-v1.5']
text2vec_models = ['shibing624/text2vec-base-chinese-sentence', 'shibing624/text2vec-base-chinese-paraphrase',
'shibing624/text2vec-base-multilingual', 'shibing624/text2vec-base-chinese',
'shibing624/text2vec-bge-large-chinese', 'GanymedeNil/text2vec-large-chinese']
if embedding_model in moka_ai_models:
return get_MokaAI_embedding(text, embedding_model)
elif embedding_model in baai_models:
return get_BAAI_embedding(text, embedding_model)
elif embedding_model in text2vec_models:
return get_text2vec_embedding(text, embedding_model)
else:
raise ValueError("Invalid embedding model specified.")
if __name__ == "__main__":
sentence1 = "Your sentence here."
print(get_sentence_embedding('BAAI/bge-small-en-v1.5', sentence1))