Views
No views yet
[CLS] токена.Encoder (rubert-tiny2)
└── [CLS] embedding (312-dim)
├── Dropout (p=0.2)
├── Linear(312 → 1) → profanity_logit
├── Linear(312 → 1) → threat_logit
└── Linear(312 → 1) → illegal_logit| Класс | Порог | Precision | Recall | F1 |
|---|---|---|---|---|
| profanity | 0.60 | 0.9454 | 0.9256 | 0.9354 |
| threat | 0.75 | 1.0000 | 0.9286 | 0.9630 |
| illegal | 0.16 | 1.0000 | 1.0000 | 1.0000 |
Пороги подобраны индивидуально для каждого класса по максимуму F1-score на валидационной выборке. Классы сильно несбалансированы — это учтено черезpos_weightвBCEWithLogitsLossпри обучении.
1import json
2import torch
3import torch.nn as nn
4from transformers import AutoModel, AutoTokenizer
5
6# 1. Загрузка токенизатора и конфига
7tokenizer = AutoTokenizer.from_pretrained("AtesiT/ru-multitask-toxicity-encoder")
8
9with open("toxicity_config.json") as f:
10 config = json.load(f)
11
12# 2. Определение архитектуры
13class MultiTaskToxicityEncoder(nn.Module):
14 def __init__(self, model_name, hidden_size, dropout=0.2):
15 super().__init__()
16 self.encoder = AutoModel.from_pretrained(model_name)
17 self.dropout = nn.Dropout(dropout)
18 self.profanity_head = nn.Linear(hidden_size, 1)
19 self.threat_head = nn.Linear(hidden_size, 1)
20 self.illegal_head = nn.Linear(hidden_size, 1)
21
22 def forward(self, input_ids, attention_mask):
23 out = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
24 cls = self.dropout(out.last_hidden_state[:, 0, :])
25 return (
26 self.profanity_head(cls),
27 self.threat_head(cls),
28 self.illegal_head(cls),
29 )
30
31# 3. Загрузка весов
32model = MultiTaskToxicityEncoder(
33 model_name=config["base_model"],
34 hidden_size=config["hidden_size"],
35)
36state_dict = torch.load("model_weights.pt", map_location="cpu")
37model.load_state_dict(state_dict)
38model.eval()
39
40# 4. Инференс
41thresholds = config["thresholds"]
42
43def predict(text):
44 enc = tokenizer(
45 text, return_tensors="pt",
46 padding="max_length", truncation=True,
47 max_length=config["max_length"],
48 )
49 with torch.no_grad():
50 p_logit, t_logit, i_logit = model(
51 enc["input_ids"], enc["attention_mask"]
52 )
53 probs = {
54 "profanity": torch.sigmoid(p_logit).item(),
55 "threat": torch.sigmoid(t_logit).item(),
56 "illegal": torch.sigmoid(i_logit).item(),
57 }
58 labels = {k: int(v >= thresholds[k]) for k, v in probs.items()}
59 return {"probs": probs, "labels": labels}
60
61print(predict("Ты полный идиот, заткнись!"))BCEWithLogitsLoss с pos_weight для каждого классаthreat, illegal) — качество на реальных данных может отличаться.