From cf1495d8966675a05ae38c3340591b411c46e819 Mon Sep 17 00:00:00 2001 From: Anton Kiselev Date: Sat, 18 May 2019 15:19:44 +0300 Subject: [PATCH] Changes in configuration. --- data.py | 13 ++++++------- modules.py | 2 +- training.py | 9 +++++++-- 3 files changed, 14 insertions(+), 10 deletions(-) diff --git a/data.py b/data.py index c589507..e975681 100644 --- a/data.py +++ b/data.py @@ -36,9 +36,9 @@ class TokenizedDataFrameDataset(Dataset): def preprocess_text(self, text: str) -> BertInput: tokens = self.tokenizer.tokenize(text) tokens = ["[CLS]"] + tokens + ["[SEP]"] - input_ids = self.tokenizer.convert_tokens_to_ids(tokens) + input_ids = self.tokenizer.convert_tokens_to_ids(tokens)[:self.max_seq_len] - segment_ids = [0] * len(tokens) + segment_ids = [0] * len(input_ids) input_mask = [1] * len(input_ids) padding = [0] * (self.max_seq_len - len(input_ids)) @@ -46,16 +46,15 @@ class TokenizedDataFrameDataset(Dataset): input_mask += padding segment_ids += padding - assert len(input_ids) == self.max_seq_len - assert len(input_mask) == self.max_seq_len - assert len(segment_ids) == self.max_seq_len - + assert len(input_ids) == self.max_seq_len, f'{len(input_ids)} != {self.max_seq_len}' + assert len(input_mask) == self.max_seq_len, f'{len(input_mask)} != {self.max_seq_len}' + assert len(segment_ids) == self.max_seq_len, f'{len(segment_ids)} != {self.max_seq_len}' return BertInput(*[torch.LongTensor(_) for _ in [input_ids, input_mask, segment_ids]]) def preprocess_label(self, label: int): result = torch.zeros(len(self.y_labels)) result[label] = 1 - return result + return torch.LongTensor(result) def __getitem__(self, index) -> dict: sample = self.df.iloc[index] diff --git a/modules.py b/modules.py index 3e8d377..473a9ad 100644 --- a/modules.py +++ b/modules.py @@ -92,7 +92,7 @@ class ClassificationModel(nn.Module): if labels is not None: loss_function = nn.CrossEntropyLoss() - loss = loss_function(logits.view(-1, self.n_labels), labels.view(-1)) + loss = loss_function(logits, labels) return loss, logits else: return logits diff --git a/training.py b/training.py index 59d50a4..ca19b10 100644 --- a/training.py +++ b/training.py @@ -64,7 +64,7 @@ if __name__ == '__main__': model.eval() model.to(device) - optimizer = Adam(model.learnable_parameters(), lr=0.001, amsgrad=True) + optimizer = Adam(model.parameters(), lr=0.001, amsgrad=True) print('Model have initialized') for i in range(args.num_epochs): @@ -73,7 +73,12 @@ if __name__ == '__main__': optimizer.zero_grad() input_ids, input_mask, segment_ids = batch['x'] + input_ids = input_ids.to(device) + input_mask = input_mask.to(device) + segment_ids = segment_ids.to(device) + y = batch['y'] + y = y.to(device) loss, _ = model.forward(input_ids, input_mask, segment_ids, labels=y) loss.backward() @@ -83,7 +88,7 @@ if __name__ == '__main__': model.eval() labels = [] predictions = [] - for batch in test_loader: + for batch in tqdm(test_loader): optimizer.zero_grad() input_ids, input_mask, segment_ids = batch['x']