File size: 5,029 Bytes
37586a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
946b455
37586a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8b0a86e
37586a2
 
 
 
 
946b455
 
 
37586a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
from model import DocBERT
from models.lstm_model import DocumentBiLSTM
from dataset import DataLoader, DocumentDataset
from utils.word_segmentation_vi import word_segmentation_vi
import numpy as np
from transformers import AutoTokenizer
import torch.nn.functional as F
import torch

args = {
    "bert_model": "vinai/phobert-base-v2", # Base BERT model name
    "model_path": "./vietnamese_hate_speech_detection_phobert/vinai_phobert-base-v2_finetuned.pth", # Change this if you have a fine-tuned model somewhere else
    "lstm_model_path": "./vietnamese_hate_speech_detection_phobert/distilled_lstm_model.pth", # Change this if you have a fine-tuned model somewhere else
    "max_seq_length": 250,
    "num_classes": 4, # As the fine tuned model has 4 classes per category
    "num_categories": 5, # As the fine tuned model has 5 categories
}

class_names = ["NORMAL", "CLEAN", "OFFENSIVE", "HATE"]

def load_model_bert():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    model = DocBERT(bert_model_name=args["bert_model"], num_classes=args["num_classes"], num_categories=args["num_categories"])
    model.load_state_dict(torch.load(args["model_path"], map_location=device))
    model = model.to(device)
    return model, device

def load_model_lstm():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    tokenizer = AutoTokenizer.from_pretrained(args["bert_model"])
    vocab_size = tokenizer.vocab_size
    model = DocumentBiLSTM(vocab_size=vocab_size,
                           embedding_dim=300,
                           hidden_dim=256,
                           n_layers=2,
                           output_dim=args["num_classes"] * args["num_categories"])
    model.load_state_dict(torch.load(args["lstm_model_path"], map_location=device)["model_state_dict"])
    model = model.to(device)
    return model, device

def inference(model, device, comments: str | list, threshold: float = 0.5):
    if isinstance(comments, str):
        comments = [comments]
    elif not isinstance(comments, list):
        raise ValueError("comment must be a string or a list of strings")
    
    comments = np.array([word_segmentation_vi(comment) for comment in comments])
    data = DocumentDataset(texts=comments, labels=None, tokenizer_name=args["bert_model"], max_length=args["max_seq_length"])
    inference_loader = DataLoader(data, batch_size=comments.shape[0], shuffle=False)

    batch = next(iter(inference_loader))
    input_ids = batch['input_ids']
    attention_mask = batch['attention_mask']
    token_type_ids = batch['token_type_ids']

    input_ids = input_ids.to(device)
    attention_mask = attention_mask.to(device)
    token_type_ids = token_type_ids.to(device)

    with torch.no_grad():
        outputs = model(input_ids, attention_mask=attention_mask)
        if args["num_categories"] > 1:
            batch_size, total_classes = outputs.shape
            if total_classes % args["num_categories"] != 0:
                raise ValueError("Error: Number of total classes in the batch must of divisible by the number of categories.")

            classes_per_group = total_classes // args["num_categories"]
            # Group every classes_per_group values along dim=1
            reshaped = outputs.view(outputs.size(0), -1, classes_per_group)  # shape: (batch, self., classes_per_group)
            probs = F.softmax(reshaped, dim=1)

            # Keep only the probs that are above the threshold (to prevent false positive), else set it to 0 (NORMAL, in this case unconclusive)
            probs = torch.where(probs > threshold, probs, 0.0)
            # Argmax over each group of classes_per_group
            predictions = probs.argmax(dim=-1)
        else:
            predictions = torch.argmax(outputs, dim=-1)

    preds_array = predictions.cpu().numpy()
    result = []
    for i in range(preds_array.shape[0]):
        result.append(
        {
            "Bình luận": comments[i],
            "Cá nhân": class_names[ preds_array[i, 0] ],
            "Nhóm/tổ chức": class_names[ preds_array[i, 1] ],
            "Tôn giáo/tín ngưỡng": class_names[ preds_array[i, 2] ],
            "Chủng tộc/sắc tộc": class_names[ preds_array[i, 3] ],
            "Chính trị": class_names[ preds_array[i, 4] ],
        })
    return result

if __name__ == "__main__":
    
    model, device = load_model_bert()
    comments = [
        "Để avata bít ngay là ngu hơn chó",
        "Hàn Quốc chửi dân Đông Lào và đây là hậu quả",
        "Nguyễn Thuận =)) tư tưởng rừng rú gì vậy",
        "@công danh nguyen thể chế chính trị khác hẳn tư tưởng xã hội nhé. Con cờ hó china liên quan cmn gì?"
    ]
    predictions = inference(model, device, comments)
    print("BERT Predictions:")
    print(predictions)

    lstm_model, device = load_model_lstm()
    lstm_predictions = inference(lstm_model, device, comments)
    print("LSTM Predictions:")
    print(lstm_predictions)