Views
No views yet
DeepPavlov/rubert-base-cased и определяет, является ли сообщение SPAM или SAFE.
Лёгкая и быстрая модель.alt-gnome/telegram-spam, не пересекающемся с обучающей выборкой:| Метрика | Значение |
|---|---|
| F1 | 0.95 |
| Precision | 0.98 |
| Recall | 0.92 |
DeepPavlov/rubert-base-casedDropout(0.2) → Linear(hidden_size, 1)[CLS]BCEWithLogitsLoss с весом положительного класса (pos_weight) для компенсации дисбаланса SAFE/SPAMsigmoid → вероятность класса SPAMп.р.и.в.е.т. → привет)ᴀ, α, латиница вместо кириллицы и т.д.) на кириллические эквиваленты/cmd@bot)приввееееет → привет)pip install torch transformers huggingface_hub1import re
2import json
3import torch
4import torch.nn as nn
5from transformers import AutoTokenizer, AutoModel
6from huggingface_hub import hf_hub_download
7
8REPO = "SafeTechDev/Russian-Spam-classifier"
9device = "cuda" if torch.cuda.is_available() else "cpu"
10
11# ── Архитектура ──────────────────────────────────────────────────────────
12class BinaryModel(nn.Module):
13 def __init__(self, model_name):
14 super().__init__()
15 self.bert = AutoModel.from_pretrained(model_name, low_cpu_mem_usage=True)
16 hidden = self.bert.config.hidden_size
17 self.binary = nn.Sequential(
18 nn.Dropout(0.2),
19 nn.Linear(hidden, 1)
20 )
21
22 def forward(self, input_ids, attention_mask):
23 outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask)
24 pooled = outputs.last_hidden_state[:, 0]
25 return self.binary(pooled).squeeze(-1)
26
27# ── Загрузка ─────────────────────────────────────────────────────────────
28config_path = hf_hub_download(REPO, "config.json")
29weights_path = hf_hub_download(REPO, "pytorch_model.bin")
30
31with open(config_path, encoding="utf-8") as f:
32 cfg = json.load(f)
33
34MAX_LENGTH = cfg.get("max_length", 40)
35
36tokenizer = AutoTokenizer.from_pretrained(REPO)
37
38model = BinaryModel(cfg["model"]).to(device)
39model.load_state_dict(torch.load(weights_path, map_location=device, weights_only=True))
40model.eval()
41
42# ── Инференс ─────────────────────────────────────────────────────────────
43def classify(text: str) -> dict:
44 enc = tokenizer(
45 text, truncation=True, padding="max_length",
46 max_length=MAX_LENGTH, return_tensors="pt"
47 )
48 with torch.no_grad():
49 logits = model(enc["input_ids"].to(device), enc["attention_mask"].to(device))
50
51 prob_spam = float(torch.sigmoid(logits).squeeze().cpu().item())
52 label = "SPAM" if prob_spam >= 0.7 else "SAFE"
53
54 return {"label": label, "prob_spam": prob_spam}
55
56
57print(classify("Дам денег, работу, пишите в лс"))
58# {'label': 'SPAM', 'prob_spam': 0.98...}
59
60print(classify("Привет, как дела?"))
61# {'label': 'SAFE', 'prob_spam': 0.02...}pytorch_model.bin с кастомной головой, поэтому напрямую через transformers.pipeline("text-classification", ...) она не запустится — используйте код выше.