NightPrince commited on
Commit
2fed74d
·
verified ·
1 Parent(s): e66564e

Add catt/ed_pl.py

Browse files
Files changed (1) hide show
  1. catt/ed_pl.py +164 -0
catt/ed_pl.py ADDED
@@ -0,0 +1,164 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import torch
3
+ import torch.nn as nn
4
+ import pytorch_lightning as pl
5
+ from torch.nn import functional as F
6
+ from torch.utils.data import DataLoader
7
+ from ed import Transformer
8
+ from tqdm import tqdm
9
+ import math
10
+ import torch
11
+ import torch.nn as nn
12
+
13
+ from torch.nn.utils.rnn import pad_sequence
14
+ # sequences is a list of tensors of shape TxH where T is the seqlen and H is the feats dim
15
+ def pad_seq(sequences, batch_first=True, padding_value=0.0, prepadding=True):
16
+ lens = [i.shape[0]for i in sequences]
17
+ padded_sequences = pad_sequence(sequences, batch_first=True, padding_value=padding_value) # NxTxH
18
+ if prepadding:
19
+ for i in range(len(lens)):
20
+ padded_sequences[i] = padded_sequences[i].roll(-lens[i])
21
+ if not batch_first:
22
+ padded_sequences = padded_sequences.transpose(0, 1) # TxNxH
23
+ return padded_sequences
24
+
25
+
26
+
27
+ def get_batches(X, batch_size=16):
28
+ num_batches = math.ceil(len(X) / batch_size)
29
+ for i in range(num_batches):
30
+ x = X[i*batch_size : (i+1)*batch_size]
31
+ yield x
32
+
33
+
34
+ class TashkeelModel(pl.LightningModule):
35
+ def __init__(self, tokenizer, max_seq_len, d_model=512, n_layers=3, n_heads=16, drop_prob=0.1, learnable_pos_emb=True):
36
+
37
+ super(TashkeelModel, self).__init__()
38
+
39
+ ffn_hidden = 4 * d_model
40
+ src_pad_idx = tokenizer.letters_map['<PAD>']
41
+ trg_pad_idx = tokenizer.tashkeel_map['<PAD>']
42
+ enc_voc_size = len(tokenizer.letters_map) # 37 + 3
43
+ dec_voc_size = len(tokenizer.tashkeel_map) # 15 + 3
44
+ self.transformer = Transformer(src_pad_idx=src_pad_idx,
45
+ trg_pad_idx=trg_pad_idx,
46
+ d_model=d_model,
47
+ enc_voc_size=enc_voc_size,
48
+ dec_voc_size=dec_voc_size,
49
+ max_len=max_seq_len,
50
+ ffn_hidden=ffn_hidden,
51
+ n_head=n_heads,
52
+ n_layers=n_layers,
53
+ drop_prob=drop_prob,
54
+ learnable_pos_emb=learnable_pos_emb
55
+ )
56
+
57
+ self.criterion = nn.CrossEntropyLoss(ignore_index=tokenizer.tashkeel_map['<PAD>'])
58
+ self.tokenizer = tokenizer
59
+
60
+
61
+ def forward(self, x, y=None):
62
+ y_pred = self.transformer(x, y)
63
+ return y_pred
64
+
65
+
66
+ def training_step(self, batch, batch_idx):
67
+ input_ids, target_ids = batch
68
+ input_ids = input_ids[:, :-1]
69
+ y_in = target_ids[:, :-1]
70
+ y_out = target_ids[:, 1:]
71
+ y_pred = self(input_ids, y_in)
72
+ loss = self.criterion(y_pred.transpose(1, 2), y_out)
73
+
74
+ self.log('train_loss', loss, prog_bar=True)
75
+ sch = self.lr_schedulers()
76
+ sch.step()
77
+ self.log('lr', sch.get_last_lr()[0], prog_bar=True)
78
+ return loss
79
+
80
+
81
+ def validation_step(self, batch, batch_idx):
82
+ input_ids, target_ids = batch
83
+ input_ids = input_ids[:, :-1]
84
+ y_in = target_ids[:, :-1]
85
+ y_out = target_ids[:, 1:]
86
+ y_pred = self(input_ids, y_in)
87
+ loss = self.criterion(y_pred.transpose(1, 2), y_out)
88
+
89
+ pred_text_with_tashkeels = self.tokenizer.decode(input_ids, y_pred.argmax(2).squeeze())
90
+ true_text_with_tashkeels = self.tokenizer.decode(input_ids, y_out)
91
+ total_val_der_distance = 0
92
+ total_val_der_ref_length = 0
93
+ for i in range(len(true_text_with_tashkeels)):
94
+ pred_text_with_tashkeel = pred_text_with_tashkeels[i]
95
+ true_text_with_tashkeel = true_text_with_tashkeels[i]
96
+ val_der = self.tokenizer.compute_der(true_text_with_tashkeel, pred_text_with_tashkeel)
97
+ total_val_der_distance += val_der['distance']
98
+ total_val_der_ref_length += val_der['ref_length']
99
+
100
+ total_der_error = total_val_der_distance / total_val_der_ref_length
101
+ self.log('val_loss', loss)
102
+ self.log('val_der', torch.FloatTensor([total_der_error]))
103
+ self.log('val_der_distance', torch.FloatTensor([total_val_der_distance]))
104
+ self.log('val_der_ref_length', torch.FloatTensor([total_val_der_ref_length]))
105
+
106
+
107
+ def test_step(self, batch, batch_idx):
108
+ input_ids, target_ids = batch
109
+ y_pred = self(input_ids, None)
110
+ loss = self.criterion(y_pred.transpose(1, 2), target_ids)
111
+ self.log('test_loss', loss)
112
+
113
+
114
+ def configure_optimizers(self):
115
+ optimizer = torch.optim.AdamW(self.parameters(), lr=3e-4)
116
+ #max_iters = 10000
117
+ #lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max_iters, eta_min=3e-6)
118
+ gamma = 1 / 1.000001
119
+ #gamma = 1 / 1.0001
120
+ lr_scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma)
121
+ opts = {"optimizer": optimizer, "lr_scheduler": lr_scheduler}
122
+ return opts
123
+
124
+
125
+ @torch.no_grad()
126
+ def do_tashkeel_batch(self, texts, batch_size=16, verbose=True):
127
+ self.eval()
128
+ device = next(self.parameters()).device
129
+ text_with_tashkeel = []
130
+ data_iter = get_batches(texts, batch_size)
131
+ if verbose:
132
+ num_batches = math.ceil(len(texts) / batch_size)
133
+ data_iter = tqdm(data_iter, total=num_batches)
134
+ for texts_mini in data_iter:
135
+ input_ids_list = []
136
+ for text in texts_mini:
137
+ input_ids, _ = self.tokenizer.encode(text, test_match=False)
138
+ input_ids_list.append(input_ids)
139
+ batch_input_ids = pad_seq(input_ids_list, batch_first=True, padding_value=self.tokenizer.letters_map['<PAD>'], prepadding=False)
140
+ target_ids = torch.LongTensor([[self.tokenizer.tashkeel_map['<BOS>']]] * len(texts_mini)).to(device)
141
+ src = batch_input_ids.to(device)
142
+
143
+ src_mask = self.transformer.make_pad_mask(src, src, self.transformer.src_pad_idx, self.transformer.src_pad_idx).to(device)
144
+ enc_src = self.transformer.encoder(src, src_mask)
145
+
146
+ for i in range(src.shape[1] - 1):
147
+ trg = target_ids
148
+ src_trg_mask = self.transformer.make_pad_mask(trg, src, self.transformer.trg_pad_idx, self.transformer.src_pad_idx).to(device)
149
+ trg_mask = self.transformer.make_pad_mask(trg, trg, self.transformer.trg_pad_idx, self.transformer.trg_pad_idx).to(device) * \
150
+ self.transformer.make_no_peak_mask(trg, trg).to(device)
151
+
152
+ preds = self.transformer.decoder(trg, enc_src, trg_mask, src_trg_mask)
153
+ # IMPORTANT NOTE: the following code snippet is to FORCE the prediction of the input space char to output no_tashkeel tag '<NT>'
154
+ target_ids = torch.cat([target_ids, preds[:, -1].argmax(1).unsqueeze(1)], axis=1)
155
+ target_ids[self.tokenizer.letters_map[' '] == src[:, :target_ids.shape[1]]] = self.tokenizer.tashkeel_map[self.tokenizer.no_tashkeel_tag]
156
+ # target_ids = torch.cat([target_ids, preds[:, -1].argmax(1).unsqueeze(1)], axis=1)
157
+ text_with_tashkeel_mini = self.tokenizer.decode(src, target_ids)
158
+ text_with_tashkeel += text_with_tashkeel_mini
159
+ return text_with_tashkeel
160
+
161
+
162
+ @torch.no_grad()
163
+ def do_tashkeel(self, text):
164
+ return self.do_tashkeel_batch([text])[0]