| --- |
| language: ru |
| license: mit |
| tags: |
| - toxicity |
| - multi-task |
| - russian |
| - bert |
| - profanity |
| - threats |
| - illegal-acts |
| - text-classification |
| datasets: |
| - toxic-russian-comments-multilabel |
| metrics: |
| - f1 |
| - precision |
| - recall |
| --- |
| |
| # 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 |
|
|
| ## Использование |
|
|
| ```python |
| 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 в репозитории. |