Multi-Task Toxicity Classifier
Описание
Модель для одновременного обнаружения трёх типов токсичности в русскоязычных текстах:
- Profanity (ненормативная лексика)
- Threat (угрозы)
- Illegal (запросы на нарушение закона)
Модель основана на энкодере cointegrated/rubert-tiny2 с тремя независимыми классификационными головами.
Метрики на валидационной выборке
| Класс | Threshold | Precision | Recall | F1-Score |
|---|---|---|---|---|
| Profanity | 0.51 | 0.9048 | 0.9429 | 0.9235 |
| Threat | 0.35 | 0.7462 | 0.8097 | 0.7766 |
| Illegal | 0.25 | 0.6038 | 0.6598 | 0.6305 |
Macro F1-Score: 0.7769
Использование
from transformers import AutoTokenizer, AutoModel
import torch
import torch.nn as nn
import json
import requests
# Загрузка модели и токенизатора
tokenizer = AutoTokenizer.from_pretrained("qquarkq/multitask-toxicity-classifier")
# Загрузка конфигурации с порогами
url = "https://huggingface.co/qquarkq/multitask-toxicity-classifier/resolve/main/config.json"
response = requests.get(url)
config = response.json()
#Или
#with open("config.json", "r") as f:
# config = json.load(f)
thresholds = config["thresholds"]
class MultiTaskToxicityEncoder(nn.Module):
def __init__(self, model_name, dropout=0.2, freeze_encoder=False):
super().__init__()
self.encoder = AutoModel.from_pretrained(model_name)
if freeze_encoder:
for param in self.encoder.parameters():
param.requires_grad = False
self.hidden_size = self.encoder.config.hidden_size
self.dropout = nn.Dropout(dropout)
#Головы
self.profanity_head = nn.Linear(self.hidden_size, 1) # Нецензурная лексика
self.threat_head = nn.Linear(self.hidden_size, 1) # Угрозы
self.illegal_head = nn.Linear(self.hidden_size, 1) # Незаконный контент
def forward(self, input_ids, attention_mask):
outputs = self.encoder(
input_ids=input_ids,
attention_mask=attention_mask,
return_dict=True
)
cls_embedding = outputs.last_hidden_state[:, 0, :]
cls_embedding = self.dropout(cls_embedding)
profanity_logits = self.profanity_head(cls_embedding)
threat_logits = self.threat_head(cls_embedding)
illegal_logits = self.illegal_head(cls_embedding)
return profanity_logits, threat_logits, illegal_logits
model = MultiTaskToxicityEncoder(model_name=config["model_name"])
def predict_toxicity(text):
encoding = tokenizer(
text,
truncation=True,
padding='max_length',
max_length=128,
return_tensors='pt'
)
with torch.no_grad():
profanity_logits, threat_logits, illegal_logits = model(
encoding['input_ids'],
encoding['attention_mask']
)
profanity_prob = torch.sigmoid(profanity_logits).item()
threat_prob = torch.sigmoid(threat_logits).item()
illegal_prob = torch.sigmoid(illegal_logits).item()
profanity_pred = int(profanity_prob >= thresholds['profanity'])
threat_pred = int(threat_prob >= thresholds['threat'])
illegal_pred = int(illegal_prob >= thresholds['illegal'])
return {
'profanity': {
'probability': profanity_prob,
'prediction': profanity_pred,
'confidence': profanity_prob * 100
},
'threat': {
'probability': threat_prob,
'prediction': threat_pred,
'confidence': threat_prob * 100
},
'illegal': {
'probability': illegal_prob,
'prediction': illegal_pred,
'confidence': illegal_prob * 100
}
}
# Пример использования
text = "Ты полный идиот!"
result = predict_toxicity(text)
print(result)
Датасет Модель обучена на датасете qquarkq/russian-toxic-multilabel-comments
- Downloads last month
- 165