Views
No views yet
rubert-tiny2. Модель одновременно предсказывает 5 различных параметров текста, связанных с психологическим состоянием, стадиями конфликта и эмоциональными триггерами.Стадии конфликта
Предконфликтная
Инцидент
Эскалация
Пик
Разрешение
Интенсивность эмоции
Низкая
Средняя
Высокая
Группа эмоций
Гнев
Грусть
Страх
Позитив
Социальные
Нейтральные
Детальные эмоции
Злость
Раздражение
Неодобрение
Обида
Горе/Грусть
Разочарование
Раскаяние
Напряжение/Нервозность
Страх
Забота
Любовь
Признательность/Радость
Облегчение
Оптимизм
Желание/Мотивация
Осознание
Одобрение
Нейтральность
Удивление
Эмпатия
Триггеры
Быт
Внимание
Деньги
Дети
Друзья
Личность
Родня
Секс
Ревность
Контроль1import torch
2import torch.nn as nn
3from transformers import AutoModel
4from safetensors.torch import load_file
5from transformers import AutoTokenizer
6from huggingface_hub import hf_hub_download
7
8class MultiTaskRuBERT(nn.Module):
9 def __init__(self, num_stages=5, num_stranges=3, num_groups=6, num_emotions=20, num_triggers=10):
10 super().__init__()
11 self.bert = AutoModel.from_pretrained('cointegrated/rubert-tiny2')
12 self.hidden_size = self.bert.config.hidden_size
13
14 self.stage_head = nn.Linear(self.hidden_size, num_stages)
15 self.strange_head = nn.Linear(self.hidden_size, num_stranges)
16 self.group_head = nn.Linear(self.hidden_size, num_groups)
17 self.emotion_head = nn.Linear(self.hidden_size, num_emotions)
18 self.trigger_head = nn.Linear(self.hidden_size, num_triggers)
19
20 repo_id = "vaveyko/rubert-tiny2-mtl-conflict"
21 weights_path = hf_hub_download(repo_id=repo_id, filename="model.safetensors")
22 state_dict = load_file(weights_path)
23 self.load_state_dict(state_dict)
24 self.eval()
25
26 self.tokenizer = AutoTokenizer.from_pretrained(repo_id)
27
28 self.task_order = ['stage_id', 'strange_id', 'group_id', 'emotion_id', 'trigger_id']
29 self.ids2labels = {'stage_id': ['Предконфликтная', 'Пик', 'Инцидент', 'Разрешение', 'Эскалация'],
30 'strange_id': ['Низкая', 'Высокая', 'Средняя'],
31 'group_id': ['Гнев', 'Страх', 'Социальные', 'Нейтральные', 'Позитив', 'Грусть'],
32 'emotion_id': ['Раздражение', 'Злость', 'Напряжение/Нервозность', 'Раскаяние',
33 'Неодобрение', 'Нейтральность', 'Облегчение', 'Горе/Грусть',
34 'Признательность/Радость', 'Страх', 'Эмпатия', 'Обида',
35 'Осознание', 'Разочарование', 'Удивление', 'Оптимизм',
36 'Желание/Мотивация', 'Любовь', 'Одобрение', 'Забота'],
37 'trigger_id': ['Быт', 'Личность', 'Ревность', 'Внимание', 'Родня', 'Контроль',
38 'Деньги', 'Дети', 'Секс', 'Друзья']
39 }
40 self.labels2ids = {'stage_id': {'Предконфликтная': 0, 'Пик': 1, 'Инцидент': 2, 'Разрешение': 3, 'Эскалация': 4},
41 'strange_id': {'Низкая': 0, 'Высокая': 1, 'Средняя': 2},
42 'group_id': {'Гнев': 0,'Страх': 1,'Социальные': 2,'Нейтральные': 3,'Позитив': 4,'Грусть': 5},
43 'emotion_id': {'Раздражение': 0,'Злость': 1,'Напряжение/Нервозность': 2,'Раскаяние': 3,'Неодобрение': 4,
44 'Нейтральность': 5,'Облегчение': 6,'Горе/Грусть': 7,'Признательность/Радость': 8,
45 'Страх': 9,'Эмпатия': 10,'Обида': 11,'Осознание': 12,'Разочарование': 13,'Удивление': 14,
46 'Оптимизм': 15,'Желание/Мотивация': 16,'Любовь': 17,'Одобрение': 18,'Забота': 19},
47 'trigger_id': {'Быт': 0,'Личность': 1,'Ревность': 2,'Внимание': 3,'Родня': 4,'Контроль': 5,'Деньги': 6,
48 'Дети': 7,'Секс': 8,'Друзья': 9}}
49
50 def forward(self, input_ids, attention_mask, token_type_ids, stage_id=None, strange_id=None,
51 group_id=None, emotion_id=None, trigger_id=None, **kwargs):
52 outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids)
53 cls = outputs.pooler_output
54
55 stage_logits = self.stage_head(cls)
56 strange_logits = self.strange_head(cls)
57 group_logits = self.group_head(cls)
58 emotion_logits = self.emotion_head(cls)
59 trigger_logits = self.trigger_head(cls)
60
61 loss = None
62 if stage_id is not None and strange_id is not None and group_id is not None and emotion_id is not None and trigger_id is not None:
63 loss_fct = nn.CrossEntropyLoss()
64 loss = (loss_fct(stage_logits, stage_id) +
65 loss_fct(strange_logits, strange_id) +
66 loss_fct(group_logits, group_id) +
67 2 * loss_fct(emotion_logits, emotion_id) +
68 1.5 * loss_fct(trigger_logits, trigger_id)) / 5
69
70 return {
71 'loss': loss,
72 'logits': (stage_logits, strange_logits, group_logits, emotion_logits, trigger_logits)
73 }
74
75 def predict_samples(self, inputs: list[str], top_k=1):
76 top_k = top_k if 1 <= top_k <= 3 else 1
77 samples = self.tokenizer(inputs, padding=True, truncation=True)
78 logits = self.forward(**{key: torch.tensor(val) for key, val in samples.items()})["logits"]
79
80 out = {}
81 for i, task in enumerate(self.task_order):
82 val, idxes = logits[i].topk(top_k, dim=-1)
83 out[task] = [[self.ids2labels[task][idx] for idx in row] for row in idxes]
84
85 return out
86
87inp = input()
88
89model = MultiTaskRuBERT()
90print(model.predict_samples([inp]))