Views
No views yet
| Model | HR@5 [%] | Improvement [%] | Embedding Size |
|---|---|---|---|
| Google Text-Embedding-004 | 84 | 5 | 768 |
| Cohere Embed-English-v3.0 | 85 | 4 | 1024 |
| OpenAI Text-Embedding-3-Large | 86 | 2 | 3072 |
| MistralAI Mistral-Embed | 87 | 1 | 1024 |
| VoyageAI Voyage-Finance-2 | 88 | 0 | 1024 |
| Ours | 88 | - | 768 |
1from sentence_transformers import SentenceTransformer
2
3model = SentenceTransformer("aminhaeri/RiskEmbed")
4
5queries = ['what is snowflake?', 'Where can I get the best tacos?']
6documents = ['The Data Cloud!', 'Mexico City of Course!']
7
8query_embeddings = model.encode(queries, prompt_name="query")
9document_embeddings = model.encode(documents)
10
11scores = query_embeddings @ document_embeddings.T
12for query, query_scores in zip(queries, scores):
13 doc_score_pairs = list(zip(documents, query_scores))
14 doc_score_pairs = sorted(doc_score_pairs, key=lambda x: x[1], reverse=True)
15 # Output passages & scores
16 print("Query:", query)
17 for document, score in doc_score_pairs:
18 print(score, document)1import torch
2from transformers import AutoModel, AutoTokenizer
3
4tokenizer = AutoTokenizer.from_pretrained('aminhaeri/RiskEmbed')
5model = AutoModel.from_pretrained('aminhaeri/RiskEmbed', add_pooling_layer=False)
6model.eval()
7
8query_prefix = 'Represent this sentence for searching relevant passages: '
9queries = ['what is snowflake?', 'Where can I get the best tacos?']
10queries_with_prefix = ["{}{}".format(query_prefix, i) for i in queries]
11query_tokens = tokenizer(queries_with_prefix, padding=True, truncation=True, return_tensors='pt', max_length=512)
12
13documents = ['The Data Cloud!', 'Mexico City of Course!']
14document_tokens = tokenizer(documents, padding=True, truncation=True, return_tensors='pt', max_length=512)
15
16# Compute token embeddings
17with torch.no_grad():
18 query_embeddings = model(**query_tokens)[0][:, 0]
19 document_embeddings = model(**document_tokens)[0][:, 0]
20
21# normalize embeddings
22query_embeddings = torch.nn.functional.normalize(query_embeddings, p=2, dim=1)
23document_embeddings = torch.nn.functional.normalize(document_embeddings, p=2, dim=1)
24
25scores = torch.mm(query_embeddings, document_embeddings.transpose(0, 1))
26for query, query_scores in zip(queries, scores):
27 doc_score_pairs = list(zip(documents, query_scores))
28 doc_score_pairs = sorted(doc_score_pairs, key=lambda x: x[1], reverse=True)
29 #Output passages & scores
30 print("Query:", query)
31 for document, score in doc_score_pairs:
32 print(score, document)