1from huggingface_hub import hf_hub_download, PyTorchModelHubMixin
2import torch
3import torch.nn as nn
4from torchcrf import CRF
5
6device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
7
8REPO_ID = "hoangduy0610/uit-cs112-sentiment-span-extraction-vietnamese"
9FILENAME = "bilstm_crf_state.pth"
10CONFIG = {
11 'vocab_size': 18182,
12 'tag_to_idx': {
13 'I-CAMERA#NEUTRAL': 0,
14 'B-SER&ACC#NEGATIVE': 1,
15 'B-SER&ACC#NEUTRAL': 2,
16 'B-CAMERA#NEUTRAL': 3,
17 'B-STORAGE#NEUTRAL': 4,
18 'I-STORAGE#NEUTRAL': 5,
19 'I-BATTERY#NEUTRAL': 6,
20 'B-STORAGE#NEGATIVE': 7,
21 'I-STORAGE#NEGATIVE': 8,
22 'B-GENERAL#POSITIVE': 9,
23 'I-CAMERA#NEGATIVE': 10,
24 'B-GENERAL#NEUTRAL': 11,
25 'I-DESIGN#NEUTRAL': 12,
26 'I-PERFORMANCE#POSITIVE': 13,
27 'B-SCREEN#NEUTRAL': 14,
28 'B-FEATURES#POSITIVE': 15,
29 'I-DESIGN#NEGATIVE': 16,
30 'I-BATTERY#NEGATIVE': 17,
31 'B-FEATURES#NEGATIVE': 18,
32 'I-DESIGN#POSITIVE': 19,
33 'I-FEATURES#NEGATIVE': 20,
34 'B-CAMERA#NEGATIVE': 21,
35 'I-PRICE#NEUTRAL': 22,
36 'I-FEATURES#NEUTRAL': 23,
37 'B-PRICE#POSITIVE': 24,
38 'I-PERFORMANCE#NEUTRAL': 25,
39 'I-FEATURES#POSITIVE': 26,
40 'I-PERFORMANCE#NEGATIVE': 27,
41 'B-SER&ACC#POSITIVE': 28,
42 'B-PRICE#NEGATIVE': 29,
43 'I-SCREEN#NEUTRAL': 30,
44 'B-DESIGN#NEUTRAL': 31,
45 'B-BATTERY#POSITIVE': 32,
46 'B-STORAGE#POSITIVE': 33,
47 'I-GENERAL#POSITIVE': 34,
48 'B-CAMERA#POSITIVE': 35,
49 'B-PERFORMANCE#NEGATIVE': 36,
50 'B-PERFORMANCE#NEUTRAL': 37,
51 'B-GENERAL#NEGATIVE': 38,
52 'I-CAMERA#POSITIVE': 39,
53 'I-BATTERY#POSITIVE': 40,
54 'I-GENERAL#NEGATIVE': 41,
55 'B-BATTERY#NEUTRAL': 42,
56 'I-SER&ACC#POSITIVE': 43,
57 'I-SER&ACC#NEUTRAL': 44,
58 'I-SER&ACC#NEGATIVE': 45,
59 'I-PRICE#NEGATIVE': 46,
60 'B-FEATURES#NEUTRAL': 47,
61 'B-SCREEN#POSITIVE': 48,
62 'B-BATTERY#NEGATIVE': 49,
63 'I-SCREEN#NEGATIVE': 50,
64 'B-SCREEN#NEGATIVE': 51,
65 'O': 52,
66 'I-GENERAL#NEUTRAL': 53,
67 'I-SCREEN#POSITIVE': 54,
68 'B-PERFORMANCE#POSITIVE': 55,
69 'B-PRICE#NEUTRAL': 56,
70 'I-STORAGE#POSITIVE': 57,
71 'B-DESIGN#NEGATIVE': 58,
72 'I-PRICE#POSITIVE': 59,
73 'B-DESIGN#POSITIVE': 60
74 },
75 'embedding_dim': 100,
76 'hidden_dim': 256,
77 'lstm_layers': 1
78}
79
80# BiLSTM-CRF model
81class BiLSTM_CRF(
82 nn.Module,
83 PyTorchModelHubMixin
84 ):
85 def __init__(self, config: dict):
86 super().__init__()
87 self.embedding = nn.Embedding(config["vocab_size"], config["embedding_dim"])
88 self.lstm = nn.LSTM(config["embedding_dim"], config["hidden_dim"] // 2,
89 num_layers=config["lstm_layers"], bidirectional=True, batch_first=True)
90 self.ln = nn.Linear(config["hidden_dim"], len(config["tag_to_idx"]))
91 self.crf = CRF(len(config["tag_to_idx"]))
92
93 def forward(self, sentence):
94 embeds = self.embedding(sentence)
95 lstm_out, _ = self.lstm(embeds)
96 emissions = self.ln(lstm_out)
97 return emissions
98
99 def summary(self):
100 print(self)
101 print('\n\nModel Summary:')
102 print('=================================================================')
103 print('Layer (type) Output Shape Param # ')
104 print('=================================================================')
105 total_params = 0
106 for name, param in self.named_parameters():
107 print(f'{name:<30} {str(param.shape):<30} {param.numel():<10}')
108 total_params += param.numel()
109 print('=================================================================')
110 print(f'Total params: {total_params}')
111
112loaded_model = BiLSTM_CRF(config=CONFIG).to(device)
113loaded_model.load_state_dict(torch.load(hf_hub_download(repo_id=REPO_ID, filename=FILENAME)))
114
115loaded_model.eval()