Views
No views yet
"search_query: " and "search_document: " prefixes are for answer or relevant paragraph retrieval"classification: " prefix is for symmetric paraphrasing related tasks (STS, NLI, Bitext Mining)"clustering: " prefix is for any tasks that rely on thematic features (topic classification, title-body retrieval)1import torch
2import torch.nn.functional as F
3from transformers import AutoTokenizer, AutoModel
4
5
6def pool(hidden_state, mask, pooling_method="cls"):
7 if pooling_method == "mean":
8 s = torch.sum(hidden_state * mask.unsqueeze(-1).float(), dim=1)
9 d = mask.sum(axis=1, keepdim=True).float()
10 return s / d
11 elif pooling_method == "cls":
12 return hidden_state[:, 0]
13
14inputs = [
15 #
16 "classification: Он нам и <unk> не нужон ваш Интернет!",
17 "clustering: В Ярославской области разрешили работу бань, но без посетителей",
18 "search_query: Сколько программистов нужно, чтобы вкрутить лампочку?",
19
20 #
21 "classification: What a time to be alive!",
22 "clustering: Ярославским баням разрешили работать без посетителей",
23 "search_document: Чтобы вкрутить лампочку, требуется три программиста: один напишет программу извлечения лампочки, другой — вкручивания лампочки, а третий проведет тестирование.",
24]
25
26tokenizer = AutoTokenizer.from_pretrained("ai-forever/ru-en-RoSBERTa")
27model = AutoModel.from_pretrained("ai-forever/ru-en-RoSBERTa")
28
29tokenized_inputs = tokenizer(inputs, max_length=512, padding=True, truncation=True, return_tensors="pt")
30
31with torch.no_grad():
32 outputs = model(**tokenized_inputs)
33
34embeddings = pool(
35 outputs.last_hidden_state,
36 tokenized_inputs["attention_mask"],
37 pooling_method="cls" # or try "mean"
38)
39
40embeddings = F.normalize(embeddings, p=2, dim=1)
41
42sim_scores = embeddings[:3] @ embeddings[3:].T
43print(sim_scores.diag().tolist())
44# [0.4796873927116394, 0.9409002065658569, 0.7761015892028809]1from sentence_transformers import SentenceTransformer
2
3
4inputs = [
5 #
6 "classification: Он нам и <unk> не нужон ваш Интернет!",
7 "clustering: В Ярославской области разрешили работу бань, но без посетителей",
8 "search_query: Сколько программистов нужно, чтобы вкрутить лампочку?",
9
10 #
11 "classification: What a time to be alive!",
12 "clustering: Ярославским баням разрешили работать без посетителей",
13 "search_document: Чтобы вкрутить лампочку, требуется три программиста: один напишет программу извлечения лампочки, другой — вкручивания лампочки, а третий проведет тестирование.",
14]
15
16# loads model with CLS pooling
17model = SentenceTransformer("ai-forever/ru-en-RoSBERTa")
18
19# embeddings are normalized by default
20embeddings = model.encode(inputs, convert_to_tensor=True)
21
22sim_scores = embeddings[:3] @ embeddings[3:].T
23print(sim_scores.diag().tolist())
24# [0.47968706488609314, 0.940900444984436, 0.7761018872261047]1from sentence_transformers import SentenceTransformer
2
3
4# loads model with CLS pooling
5model = SentenceTransformer("ai-forever/ru-en-RoSBERTa")
6
7classification = model.encode(["Он нам и <unk> не нужон ваш Интернет!", "What a time to be alive!"], prompt_name="classification")
8print(classification[0] @ classification[1].T) # 0.47968706488609314
9
10clustering = model.encode(["В Ярославской области разрешили работу бань, но без посетителей", "Ярославским баням разрешили работать без посетителей"], prompt_name="clustering")
11print(clustering[0] @ clustering[1].T) # 0.940900444984436
12
13query_embedding = model.encode("Сколько программистов нужно, чтобы вкрутить лампочку?", prompt_name="search_query")
14document_embedding = model.encode("Чтобы вкрутить лампочку, требуется три программиста: один напишет программу извлечения лампочки, другой — вкручивания лампочки, а третий проведет тестирование.", prompt_name="search_document")
15print(query_embedding @ document_embedding.T) # 0.7761018872261047@misc{snegirev2024russianfocusedembeddersexplorationrumteb,
title={The Russian-focused embedders' exploration: ruMTEB benchmark and Russian embedding model design},
author={Artem Snegirev and Maria Tikhonova and Anna Maksimova and Alena Fenogenova and Alexander Abramov},
year={2024},
eprint={2408.12503},
archivePrefix={arXiv},
primaryClass={cs.CL},
url={https://arxiv.org/abs/2408.12503},
}