FRIDA is a full-scale finetuned general text embedding model inspired by denoising architecture based on T5. The model is based on the encoder part of
FRED-T5 model and continues research of text embedding models (
ruMTEB,
ru-en-RoSBERTa). It has been pre-trained on a Russian-English dataset and fine-tuned for improved performance on the target task.
For more model details please refer to our
article (RU).
The model can be used as is with prefixes. It is recommended to use CLS pooling. The choice of prefix and pooling depends on the task.
To better tailor the model to your needs, you can fine-tune it with relevant high-quality Russian and English datasets.
Below are examples of texts encoding using the Transformers and SentenceTransformers libraries.
1import torch
2import torch.nn.functional as F
3from transformers import AutoTokenizer, T5EncoderModel
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 "paraphrase: В Ярославской области разрешили работу бань, но без посетителей",
17 "categorize_entailment: Женщину доставили в больницу, за ее жизнь сейчас борются врачи.",
18 "search_query: Сколько программистов нужно, чтобы вкрутить лампочку?",
19 #
20 "paraphrase: Ярославским баням разрешили работать без посетителей",
21 "categorize_entailment: Женщину спасают врачи.",
22 "search_document: Чтобы вкрутить лампочку, требуется три программиста: один напишет программу извлечения лампочки, другой — вкручивания лампочки, а третий проведет тестирование."
23]
24
25tokenizer = AutoTokenizer.from_pretrained("ai-forever/FRIDA")
26model = T5EncoderModel.from_pretrained("ai-forever/FRIDA")
27
28tokenized_inputs = tokenizer(inputs, max_length=512, padding=True, truncation=True, return_tensors="pt")
29
30with torch.no_grad():
31 outputs = model(**tokenized_inputs)
32
33embeddings = pool(
34 outputs.last_hidden_state,
35 tokenized_inputs["attention_mask"],
36 pooling_method="cls" # or try "mean"
37)
38
39embeddings = F.normalize(embeddings, p=2, dim=1)
40sim_scores = embeddings[:3] @ embeddings[3:].T
41print(sim_scores.diag().tolist())
42# [0.9360030293464661, 0.8591322302818298, 0.728583037853241]
1from sentence_transformers import SentenceTransformer
2
3inputs = [
4 #
5 "paraphrase: В Ярославской области разрешили работу бань, но без посетителей",
6 "categorize_entailment: Женщину доставили в больницу, за ее жизнь сейчас борются врачи.",
7 "search_query: Сколько программистов нужно, чтобы вкрутить лампочку?",
8 #
9 "paraphrase: Ярославским баням разрешили работать без посетителей",
10 "categorize_entailment: Женщину спасают врачи.",
11 "search_document: Чтобы вкрутить лампочку, требуется три программиста: один напишет программу извлечения лампочки, другой — вкручивания лампочки, а третий проведет тестирование."
12]
13
14# loads model with CLS pooling
15model = SentenceTransformer("ai-forever/FRIDA")
16
17# embeddings are normalized by default
18embeddings = model.encode(inputs, convert_to_tensor=True)
19
20sim_scores = embeddings[:3] @ embeddings[3:].T
21print(sim_scores.diag().tolist())
22# [0.9360026717185974, 0.8591331243515015, 0.7285830974578857]
1from sentence_transformers import SentenceTransformer
2
3# loads model with CLS pooling
4model = SentenceTransformer("ai-forever/FRIDA")
5
6paraphrase = model.encode(["В Ярославской области разрешили работу бань, но без посетителей", "Ярославским баням разрешили работать без посетителей"], prompt_name="paraphrase")
7print(paraphrase[0] @ paraphrase[1].T) # 0.9360032
8
9categorize_entailment = model.encode(["Женщину доставили в больницу, за ее жизнь сейчас борются врачи.", "Женщину спасают врачи."], prompt_name="categorize_entailment")
10print(categorize_entailment[0] @ categorize_entailment[1].T) # 0.8591322
11
12query_embedding = model.encode("Сколько программистов нужно, чтобы вкрутить лампочку?", prompt_name="search_query")
13document_embedding = model.encode("Чтобы вкрутить лампочку, требуется три программиста: один напишет программу извлечения лампочки, другой — вкручивания лампочки, а третий проведет тестирование.", prompt_name="search_document")
14print(query_embedding @ document_embedding.T) # 0.7285831
The model is designed to process texts in Russian, the quality in English is unknown. Maximum input text length is limited to 512 tokens.