Multi-Task Toxicity Classifier
Описание
Модель для Multi-Task классификации токсичности русскоязычных текстов. Модель одновременно предсказывает три класса:
- profanity: нецензурная лексика (мат)
- threat: угрозы
- illegal: запросы на незаконные действия
Архитектура
- Базовый энкодер:
cointegrated/rubert-tiny2 - Три независимые классификационные головы
- Размер эмбеддинга: 312
- Dropout: 0.1
Метрики
| Класс | F1-Score | Precision | Recall | Порог |
|---|---|---|---|---|
| Profanity | 0.9470 | 0.9529 | 0.9412 | 0.25 |
| Threat | 0.9220 | 0.9134 | 0.9308 | 0.25 |
| Illegal | 0.8681 | 0.8743 | 0.8620 | 0.25 |
Общие метрики
- Macro F1: 0.9124
- Micro F1: 0.9900
Использование
import torch
from transformers import AutoTokenizer
# Загрузка модели
class MultiTaskToxicityEncoder(torch.nn.Module):
def __init__(self, model_name='cointegrated/rubert-tiny2', dropout_rate=0.1):
super().__init__()
from transformers import AutoModel
self.encoder = AutoModel.from_pretrained(model_name)
self.config = self.encoder.config
self.hidden_size = self.config.hidden_size
self.dropout = torch.nn.Dropout(dropout_rate)
self.head_profanity = torch.nn.Linear(self.hidden_size, 1)
self.head_threat = torch.nn.Linear(self.hidden_size, 1)
self.head_illegal = torch.nn.Linear(self.hidden_size, 1)
def forward(self, input_ids, attention_mask):
outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
cls_embeddings = outputs.last_hidden_state[:, 0, :]
cls_embeddings = self.dropout(cls_embeddings)
return (
self.head_profanity(cls_embeddings).squeeze(-1),
self.head_threat(cls_embeddings).squeeze(-1),
self.head_illegal(cls_embeddings).squeeze(-1)
)
device = 'cuda' if torch.cuda.is_available() else 'cpu'
model = MultiTaskToxicityEncoder()
model.load_state_dict(torch.load('pytorch_model.bin', map_location=device))
model = model.to(device)
model.eval()
tokenizer = AutoTokenizer.from_pretrained('dbrovkin/toxicity-multitask-bert')
# Предсказание
def predict(text):
encoding = tokenizer(text, truncation=True, padding='max_length',
max_length=128, return_tensors='pt')
input_ids = encoding['input_ids'].to(device)
attention_mask = encoding['attention_mask'].to(device)
with torch.no_grad():
p, t, i = model(input_ids, attention_mask)
return torch.sigmoid(p).item(), torch.sigmoid(t).item(), torch.sigmoid(i).item()
# Пример
text = 'Ты просто идиот!'
profanity, threat, illegal = predict(text)
print(f'Мат: {profanity:.3f}, Угрозы: {threat:.3f}, Незаконное: {illegal:.3f}')
Пороги отсечения
Для бинарной классификации используются пороги:
- Profanity: 0.25
- Threat: 0.25
- Illegal: 0.25
Датасет
Модель обучена на сбалансированном датасете русскоязычных комментариев.
Лицензия
MIT
Контакты
Для вопросов и предложений создавайте Issue в репозитории.
- Downloads last month
- 37