SafeTechDev commited on
Commit
ee3db37
·
verified ·
1 Parent(s): a5215f6

Update README.md

Browse files
Files changed (1) hide show
  1. README.md +147 -0
README.md CHANGED
@@ -1,3 +1,150 @@
1
  ---
 
 
2
  license: mit
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ language:
3
+ - ru
4
  license: mit
5
+ base_model: DeepPavlov/rubert-base-cased
6
+ tags:
7
+ - text-classification
8
+ - spam-detection
9
+ - russian
10
+ - telegram-moderation
11
+ - bert
12
+ pipeline_tag: text-classification
13
+ datasets:
14
+ - alt-gnome/telegram-spam
15
+ metrics:
16
+ - f1
17
+ - precision
18
+ - recall
19
  ---
20
+
21
+ # AIMS-RU-Binary-Spam-Classifier
22
+
23
+ Бинарный классификатор спама для русскоязычных сообщений, часть проекта **AIMS (AI Moderation System)**.
24
+
25
+ Модель дообучена поверх [`DeepPavlov/rubert-base-cased`](https://huggingface.co/DeepPavlov/rubert-base-cased) и определяет, является ли сообщение **SPAM** или **SAFE**.
26
+ Лёгкая и быстрая модель.
27
+
28
+ ## Метрики
29
+
30
+ Оценка проводилась на внешнем датасете [`alt-gnome/telegram-spam`](https://huggingface.co/datasets/alt-gnome/telegram-spam), не пересекающемся с обучающей выборкой:
31
+
32
+ | Метрика | Значение |
33
+ |-----------|----------|
34
+ | F1 | 0.94 |
35
+ | Precision | 0.97 |
36
+ | Recall | 0.91 |
37
+
38
+ Высокий precision (0.97) означает низкую долю ложных срабатываний — модель редко помечает безопасные сообщения как спам, что важно для автоматической модерации без риска задеть живых пользователей.
39
+
40
+ ## Архитектура
41
+
42
+ - Энкодер: `DeepPavlov/rubert-base-cased`
43
+ - Голова классификации: `Dropout(0.2) → Linear(hidden_size, 1)`
44
+ - Пулинг: эмбеддинг токена `[CLS]`
45
+ - Функция потерь при обучении: `BCEWithLogitsLoss` с весом положительного класса (`pos_weight`) для компенсации дисбаланса SAFE/SPAM
46
+ - Выход: один логит → `sigmoid` → вероятность класса SPAM
47
+
48
+ ## РЕКОМЕНДАЦИЯ Нормализация текста
49
+
50
+ Перед подачей в модель рекомендуется нормализовать текст, это использовалось при обучении.
51
+
52
+ - убирает "разрежённое" написание (`п.р.и.в.е.т.` → `привет`)
53
+ - заменяет гомоглифы (похожие символы из других алфавитов — `ᴀ`, `α`, латиница вместо кириллицы и т.д.) на кириллические эквиваленты
54
+ - удаляет ссылки и команды ботов (`/cmd@bot`)
55
+ - схлопывает повторяющиеся буквы (`приввееееет` → `привет`)
56
+ - приводит текст к нижнему регистру
57
+
58
+
59
+
60
+ ## Быстрый запуск
61
+
62
+ ```bash
63
+ pip install torch transformers huggingface_hub
64
+ ```
65
+
66
+ ```python
67
+ import re
68
+ import json
69
+ import torch
70
+ import torch.nn as nn
71
+ from transformers import AutoTokenizer, AutoModel
72
+ from huggingface_hub import hf_hub_download
73
+
74
+ REPO = "SafeTechDev/AIMS-RU-Binary-Spam-Classifier"
75
+ device = "cuda" if torch.cuda.is_available() else "cpu"
76
+
77
+ # ── Архитектура ──────────────────────────────────────────────────────────
78
+ class BinaryModel(nn.Module):
79
+ def __init__(self, model_name):
80
+ super().__init__()
81
+ self.bert = AutoModel.from_pretrained(model_name, low_cpu_mem_usage=True)
82
+ hidden = self.bert.config.hidden_size
83
+ self.binary = nn.Sequential(
84
+ nn.Dropout(0.2),
85
+ nn.Linear(hidden, 1)
86
+ )
87
+
88
+ def forward(self, input_ids, attention_mask):
89
+ outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask)
90
+ pooled = outputs.last_hidden_state[:, 0]
91
+ return self.binary(pooled).squeeze(-1)
92
+
93
+ # ── Загрузка ─────────────────────────────────────────────────────────────
94
+ config_path = hf_hub_download(REPO, "config.json")
95
+ weights_path = hf_hub_download(REPO, "pytorch_model.bin")
96
+
97
+ with open(config_path, encoding="utf-8") as f:
98
+ cfg = json.load(f)
99
+
100
+ MAX_LENGTH = cfg.get("max_length", 40)
101
+
102
+ tokenizer = AutoTokenizer.from_pretrained(REPO)
103
+
104
+ model = BinaryModel(cfg["model"]).to(device)
105
+ model.load_state_dict(torch.load(weights_path, map_location=device, weights_only=True))
106
+ model.eval()
107
+
108
+ # ── Инференс ─────────────────────────────────────────────────────────────
109
+ def classify(text: str) -> dict:
110
+ enc = tokenizer(
111
+ text, truncation=True, padding="max_length",
112
+ max_length=MAX_LENGTH, return_tensors="pt"
113
+ )
114
+ with torch.no_grad():
115
+ logits = model(enc["input_ids"].to(device), enc["attention_mask"].to(device))
116
+
117
+ prob_spam = float(torch.sigmoid(logits).squeeze().cpu().item())
118
+ label = "SPAM" if prob_spam >= 0.7 else "SAFE"
119
+
120
+ return {"label": label, "prob_spam": prob_spam}
121
+
122
+
123
+ print(classify("Дам денег, работу, пишите в лс"))
124
+ # {'label': 'SPAM', 'prob_spam': 0.98...}
125
+
126
+ print(classify("Привет, как дела?"))
127
+ # {'label': 'SAFE', 'prob_spam': 0.02...}
128
+ ```
129
+
130
+ ### Через pipeline (упрощённо, без кастомной нормализации)
131
+
132
+ Модель хранит веса в формате `pytorch_model.bin` с кастомной головой, поэтому напрямую через `transformers.pipeline("text-classification", ...)` она не запустится — используйте код выше.
133
+
134
+ ## Файлы в репозитории
135
+
136
+ | Файл | Назначение |
137
+ |----------------------|-----------------------------------------------|
138
+ | `pytorch_model.bin` | Веса модели (`state_dict` для `BinaryModel`) |
139
+ | `config.json` | `model` (базовый энкодер), `max_length` |
140
+ | `tokenizer.json`, `vocab.txt`, ... | Файлы токенизатора `rubert-base-cased` |
141
+
142
+ ## Ограничения
143
+
144
+ - Модель обучена и валидирована на **русскоязычном** тексте; на других языках качество не гарантируется.
145
+ - Порог 0.7 используется по умолчанию для разделения SAFE/SPAM — в проде рекомендуется подбирать порог под свои данные (например, по максимальному F1), т.к. соотношение precision/recall меняется в зависимости от домена.
146
+ - Модель не учитывает контекст переписки — классифицируется только отдельное сообщение.
147
+
148
+ ## Автор
149
+
150
+ [SafeTechDev](https://huggingface.co/SafeTechDev)