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