yiningmao/metaphor-detection-baseline
2
1import numpy as np2 3import torch4import torch.nn as nn5import string6from sklearn.model_selection import StratifiedKFold7from torch.utils.data import DataLoader, RandomSampler, SequentialSampler, TensorDataset8from run_classifier_dataset_utils import (9 convert_examples_to_two_features,10 convert_examples_to_features,11 convert_two_examples_to_features,12)13 14 15def load_train_data(args, logger, processor, task_name, label_list, tokenizer, output_mode, k=None):16 # Prepare data loader17 if task_name == "vua":18 train_examples = processor.get_train_examples(args.data_dir)19 elif task_name == "trofi":20 train_examples = processor.get_train_examples(args.data_dir, k)21 else:22 raise ("task_name should be 'vua' or 'trofi'!")23 import pdb; pdb.set_trace()24 print(args.model_type, args.max_data_num)25 # make features file26 if args.model_type == "BERT_BASE":27 train_features = convert_two_examples_to_features(28 train_examples, label_list, args.max_seq_length, tokenizer, output_mode29 )30 if args.model_type in ["BERT_SEQ", "MELBERT_SPV"]:31 train_features = convert_examples_to_features(32 train_examples, label_list, args.max_seq_length, tokenizer, output_mode, args33 )34 if args.model_type in ["MELBERT_MIP", "MELBERT"]:35 train_features = convert_examples_to_two_features(36 train_examples, label_list, args.max_seq_length, tokenizer, output_mode, args37 )38 39 # make features into tensor40 all_input_ids = torch.tensor([f.input_ids for f in train_features], dtype=torch.long)41 all_input_mask = torch.tensor([f.input_mask for f in train_features], dtype=torch.long)42 all_segment_ids = torch.tensor([f.segment_ids for f in train_features], dtype=torch.long)43 all_label_ids = torch.tensor([f.label_id for f in train_features], dtype=torch.long)44 45 # add additional features for MELBERT_MIP and MELBERT46 if args.model_type in ["MELBERT_MIP", "MELBERT"]:47 all_input_ids_2 = torch.tensor([f.input_ids_2 for f in train_features], dtype=torch.long)48 all_input_mask_2 = torch.tensor([f.input_mask_2 for f in train_features], dtype=torch.long)49 all_segment_ids_2 = torch.tensor(50 [f.segment_ids_2 for f in train_features], dtype=torch.long51 )52 train_data = TensorDataset(53 all_input_ids,54 all_input_mask,55 all_segment_ids,56 all_label_ids,57 all_input_ids_2,58 all_input_mask_2,59 all_segment_ids_2,60 )61 else:62 train_data = TensorDataset(all_input_ids, all_input_mask, all_segment_ids, all_label_ids)63 train_sampler = RandomSampler(train_data)64 train_dataloader = DataLoader(65 train_data, sampler=train_sampler, batch_size=args.train_batch_size66 )67 68 return train_dataloader69 70 71def load_train_data_kf(72 args, logger, processor, task_name, label_list, tokenizer, output_mode, k=None73):74 # Prepare data loader75 if task_name == "vua":76 train_examples = processor.get_train_examples(args.data_dir)77 elif task_name == "trofi":78 train_examples = processor.get_train_examples(args.data_dir, k)79 else:80 raise ("task_name should be 'vua' or 'trofi'!")81 82 # make features file83 if args.model_type == "BERT_BASE":84 train_features = convert_two_examples_to_features(85 train_examples, label_list, args.max_seq_length, tokenizer, output_mode86 )87 if args.model_type in ["BERT_SEQ", "MELBERT_SPV"]:88 train_features = convert_examples_to_features(89 train_examples, label_list, args.max_seq_length, tokenizer, output_mode, args90 )91 if args.model_type in ["MELBERT_MIP", "MELBERT"]:92 train_features = convert_examples_to_two_features(93 train_examples, label_list, args.max_seq_length, tokenizer, output_mode, args94 )95 96 # make features into tensor97 all_input_ids = torch.tensor([f.input_ids for f in train_features], dtype=torch.long)98 all_input_mask = torch.tensor([f.input_mask for f in train_features], dtype=torch.long)99 all_segment_ids = torch.tensor([f.segment_ids for f in train_features], dtype=torch.long)100 all_label_ids = torch.tensor([f.label_id for f in train_features], dtype=torch.long)101 102 # add additional features for MELBERT_MIP and MELBERT103 if args.model_type in ["MELBERT_MIP", "MELBERT"]:104 all_input_ids_2 = torch.tensor([f.input_ids_2 for f in train_features], dtype=torch.long)105 all_input_mask_2 = torch.tensor([f.input_mask_2 for f in train_features], dtype=torch.long)106 all_segment_ids_2 = torch.tensor(107 [f.segment_ids_2 for f in train_features], dtype=torch.long108 )109 train_data = TensorDataset(110 all_input_ids,111 all_input_mask,112 all_segment_ids,113 all_label_ids,114 all_input_ids_2,115 all_input_mask_2,116 all_segment_ids_2,117 )118 else:119 train_data = TensorDataset(all_input_ids, all_input_mask, all_segment_ids, all_label_ids)120 gkf = StratifiedKFold(n_splits=args.num_bagging).split(X=all_input_ids, y=all_label_ids.numpy())121 return train_data, gkf122 123 124def load_test_data(args, logger, processor, task_name, label_list, tokenizer, output_mode, k=None):125 if task_name == "vua":126 eval_examples = processor.get_test_examples(args.data_dir)127 elif task_name == "trofi":128 eval_examples = processor.get_test_examples(args.data_dir, k)129 else:130 raise ("task_name should be 'vua' or 'trofi'!")131 import pdb; pdb.set_trace()132 eval_examples = eval_examples[14185:14216]133 if args.model_type == "BERT_BASE":134 eval_features = convert_two_examples_to_features(135 eval_examples, label_list, args.max_seq_length, tokenizer, output_mode136 )137 if args.model_type in ["BERT_SEQ", "MELBERT_SPV"]:138 eval_features = convert_examples_to_features(139 eval_examples, label_list, args.max_seq_length, tokenizer, output_mode, args140 )141 if args.model_type in ["MELBERT_MIP", "MELBERT"]:142 eval_features = convert_examples_to_two_features(143 eval_examples, label_list, args.max_seq_length, tokenizer, output_mode, args144 )145 import pdb; pdb.set_trace()146 logger.info("***** Running evaluation *****")147 if args.model_type in ["MELBERT_MIP", "MELBERT"]:148 all_input_ids = torch.tensor([f.input_ids for f in eval_features], dtype=torch.long)149 all_input_mask = torch.tensor([f.input_mask for f in eval_features], dtype=torch.long)150 all_segment_ids = torch.tensor([f.segment_ids for f in eval_features], dtype=torch.long)151 all_guids = [f.guid for f in eval_features]152 all_idx = torch.tensor([i for i in range(len(eval_features))], dtype=torch.long)153 all_label_ids = torch.tensor([f.label_id for f in eval_features], dtype=torch.long)154 all_input_ids_2 = torch.tensor([f.input_ids_2 for f in eval_features], dtype=torch.long)155 all_input_mask_2 = torch.tensor([f.input_mask_2 for f in eval_features], dtype=torch.long)156 all_segment_ids_2 = torch.tensor([f.segment_ids_2 for f in eval_features], dtype=torch.long)157 eval_data = TensorDataset(158 all_input_ids,159 all_input_mask,160 all_segment_ids,161 all_label_ids,162 all_idx,163 all_input_ids_2,164 all_input_mask_2,165 all_segment_ids_2,166 )167 else:168 all_input_ids = torch.tensor([f.input_ids for f in eval_features], dtype=torch.long)169 all_input_mask = torch.tensor([f.input_mask for f in eval_features], dtype=torch.long)170 all_segment_ids = torch.tensor([f.segment_ids for f in eval_features], dtype=torch.long)171 all_guids = [f.guid for f in eval_features]172 all_idx = torch.tensor([i for i in range(len(eval_features))], dtype=torch.long)173 all_label_ids = torch.tensor([f.label_id for f in eval_features], dtype=torch.long)174 eval_data = TensorDataset(175 all_input_ids, all_input_mask, all_segment_ids, all_label_ids, all_idx176 )177 178 # Run prediction for full data179 eval_sampler = SequentialSampler(eval_data)180 eval_dataloader = DataLoader(eval_data, sampler=eval_sampler, batch_size=args.eval_batch_size)181 182 return all_guids, eval_dataloader183 184from run_classifier_dataset_utils import InputExample185def load_sentence_data(args, sentence, label_list, tokenizer, output_mode, ):186 #tokens = tokenizer.tokenize(sentence)187 #print('tokens:', tokens)188 examples = []189 example_idxs = []190 for index, token in enumerate(sentence.split()):191 if token not in string.punctuation:192 examples.append(193 InputExample(194 guid='', text_a=sentence, text_b=str(index), label='0', POS='', FGPOS=''195 )196 )197 print('[', index, token, ']', end=', ')198 example_idxs.append(index)199 eval_features = convert_examples_to_two_features(200 examples, label_list, args.max_seq_length, tokenizer, output_mode, args201 )202 203 if args.model_type in ["MELBERT_MIP", "MELBERT"]:204 all_input_ids = torch.tensor([f.input_ids for f in eval_features], dtype=torch.long)205 all_input_mask = torch.tensor([f.input_mask for f in eval_features], dtype=torch.long)206 all_segment_ids = torch.tensor([f.segment_ids for f in eval_features], dtype=torch.long)207 all_guids = [f.guid for f in eval_features]208 all_idx = torch.tensor(example_idxs, dtype=torch.long)209 all_label_ids = torch.tensor([f.label_id for f in eval_features], dtype=torch.long)210 all_input_ids_2 = torch.tensor([f.input_ids_2 for f in eval_features], dtype=torch.long)211 all_input_mask_2 = torch.tensor([f.input_mask_2 for f in eval_features], dtype=torch.long)212 all_segment_ids_2 = torch.tensor([f.segment_ids_2 for f in eval_features], dtype=torch.long)213 eval_data = (214 all_input_ids,215 all_input_mask,216 all_segment_ids,217 all_label_ids,218 all_idx,219 all_input_ids_2,220 all_input_mask_2,221 all_segment_ids_2,222 )223 else:224 all_input_ids = torch.tensor([f.input_ids for f in eval_features], dtype=torch.long)225 all_input_mask = torch.tensor([f.input_mask for f in eval_features], dtype=torch.long)226 all_segment_ids = torch.tensor([f.segment_ids for f in eval_features], dtype=torch.long)227 all_guids = [f.guid for f in eval_features]228 all_idx = torch.tensor(example_idxs, dtype=torch.long)229 all_label_ids = torch.tensor([f.label_id for f in eval_features], dtype=torch.long)230 eval_data = (231 all_input_ids, all_input_mask, all_segment_ids, all_label_ids, all_idx232 )233 return eval_data