Views
No views yet
cointegrated/rubert-tiny2 и содержит три независимые классификационные «головы» для одновременного предсказания трёх классов:illegalprofanity и threat| Класс | Порог | Precision | Recall | F1-score |
|---|---|---|---|---|
| profanity | 0.50 | 0.821 | 0.908 | 0.862 |
| threat | 0.15 | 0.261 | 0.571 | 0.358 |
| illegal | 0.40 | 1.000 | 1.000 | 1.000 |
Примечание: высокий F1 для классаillegalобусловлен небольшим количеством положительных примеров в валидации; на более крупных выборках результаты могут отличаться.
1from transformers import AutoTokenizer, AutoModel
2import torch
3
4MODEL_NAME = "Arrtemwolf/rubert-tiny2-toxicity-multitask"
5tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
6encoder = AutoModel.from_pretrained(MODEL_NAME)
7
8# Загрузка обученных голов (классификаторов)
9# Внимание: головы сохранены отдельно, их нужно загрузить и прикрепить к модели
10# Ниже приведён пример класса-обёртки, который можно использовать после загрузки весов голов.
11Для удобства рекомендуется использовать класс MultiTaskToxicityEncoder, который объединяет энкодер и три головы. Веса голов сохранены в файле multitask_heads.pt в репозитории. Пример загрузки:
12
13python
14class MultiTaskToxicityEncoder(torch.nn.Module):
15 def __init__(self, encoder):
16 super().__init__()
17 self.encoder = encoder
18 hidden_size = encoder.config.hidden_size
19 self.head_profanity = torch.nn.Linear(hidden_size, 1)
20 self.head_threat = torch.nn.Linear(hidden_size, 1)
21 self.head_illegal = torch.nn.Linear(hidden_size, 1)
22 self.dropout = torch.nn.Dropout(0.3)
23
24 def forward(self, input_ids, attention_mask):
25 outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
26 cls_embedding = outputs.last_hidden_state[:, 0, :]
27 cls_embedding = self.dropout(cls_embedding)
28 return (self.head_profanity(cls_embedding),
29 self.head_threat(cls_embedding),
30 self.head_illegal(cls_embedding))
31
32# Загружаем энкодер
33encoder = AutoModel.from_pretrained(MODEL_NAME)
34model = MultiTaskToxicityEncoder(encoder)
35
36# Загружаем веса голов
37state_dict = torch.load("multitask_heads.pt", map_location="cpu")
38model.load_state_dict(state_dict, strict=False) # strict=False, т.к. веса только для голов
39
40model.eval()
41Предсказание для одного текста
42python
43def predict(text, model, tokenizer, device="cpu"):
44 encoded = tokenizer(text, padding=True, truncation=True, max_length=256, return_tensors="pt")
45 input_ids = encoded["input_ids"].to(device)
46 attention_mask = encoded["attention_mask"].to(device)
47
48 with torch.no_grad():
49 logit_p, logit_t, logit_i = model(input_ids, attention_mask)
50 prob_p = torch.sigmoid(logit_p).item()
51 prob_t = torch.sigmoid(logit_t).item()
52 prob_i = torch.sigmoid(logit_i).item()
53
54 # Пороги (оптимальные, полученные на валидации)
55 thresholds = {"profanity": 0.50, "threat": 0.15, "illegal": 0.40}
56 return {
57 "profanity": {"prob": prob_p, "label": prob_p >= thresholds["profanity"]},
58 "threat": {"prob": prob_t, "label": prob_t >= thresholds["threat"]},
59 "illegal": {"prob": prob_i, "label": prob_i >= thresholds["illegal"]},
60 }
61
62# Пример
63text = "Ты мне угрожаешь? Я вызову полицию!"
64result = predict(text, model, tokenizer)
65print(result)