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
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support