Модель для одновременного обнаружения трёх типов токсичности в русскоязычных текстах:
1from transformers import AutoTokenizer, AutoModel
2import torch
3import torch.nn as nn
4import json
5import requests
6
7# Загрузка модели и токенизатора
8
9tokenizer = AutoTokenizer.from_pretrained("qquarkq/multitask-toxicity-classifier")
10
11# Загрузка конфигурации с порогами
12url = "https://huggingface.co/qquarkq/multitask-toxicity-classifier/resolve/main/config.json"
13response = requests.get(url)
14config = response.json()
15#Или
16#with open("config.json", "r") as f:
17# config = json.load(f)
18thresholds = config["thresholds"]
19
20class MultiTaskToxicityEncoder(nn.Module):
21 def __init__(self, model_name, dropout=0.2, freeze_encoder=False):
22 super().__init__()
23
24 self.encoder = AutoModel.from_pretrained(model_name)
25 if freeze_encoder:
26 for param in self.encoder.parameters():
27 param.requires_grad = False
28 self.hidden_size = self.encoder.config.hidden_size
29 self.dropout = nn.Dropout(dropout)
30
31 #Головы
32 self.profanity_head = nn.Linear(self.hidden_size, 1) # Нецензурная лексика
33 self.threat_head = nn.Linear(self.hidden_size, 1) # Угрозы
34 self.illegal_head = nn.Linear(self.hidden_size, 1) # Незаконный контент
35
36
37
38 def forward(self, input_ids, attention_mask):
39 outputs = self.encoder(
40 input_ids=input_ids,
41 attention_mask=attention_mask,
42 return_dict=True
43 )
44 cls_embedding = outputs.last_hidden_state[:, 0, :]
45 cls_embedding = self.dropout(cls_embedding)
46
47 profanity_logits = self.profanity_head(cls_embedding)
48 threat_logits = self.threat_head(cls_embedding)
49 illegal_logits = self.illegal_head(cls_embedding)
50
51 return profanity_logits, threat_logits, illegal_logits
52
53model = MultiTaskToxicityEncoder(model_name=config["model_name"])
54
55def predict_toxicity(text):
56 encoding = tokenizer(
57 text,
58 truncation=True,
59 padding='max_length',
60 max_length=128,
61 return_tensors='pt'
62 )
63
64 with torch.no_grad():
65 profanity_logits, threat_logits, illegal_logits = model(
66 encoding['input_ids'],
67 encoding['attention_mask']
68 )
69
70 profanity_prob = torch.sigmoid(profanity_logits).item()
71 threat_prob = torch.sigmoid(threat_logits).item()
72 illegal_prob = torch.sigmoid(illegal_logits).item()
73
74 profanity_pred = int(profanity_prob >= thresholds['profanity'])
75 threat_pred = int(threat_prob >= thresholds['threat'])
76 illegal_pred = int(illegal_prob >= thresholds['illegal'])
77 return {
78 'profanity': {
79 'probability': profanity_prob,
80 'prediction': profanity_pred,
81 'confidence': profanity_prob * 100
82 },
83 'threat': {
84 'probability': threat_prob,
85 'prediction': threat_pred,
86 'confidence': threat_prob * 100
87 },
88 'illegal': {
89 'probability': illegal_prob,
90 'prediction': illegal_pred,
91 'confidence': illegal_prob * 100
92 }
93 }
94
95# Пример использования
96text = "Ты полный идиот!"
97result = predict_toxicity(text)
98print(result)