Spaces:
Sleeping
Sleeping
Add catt/ed_pl.py
Browse files- 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]
|