| 248 | |
| 249 | |
| 250 | class DecoderEvaluater: |
| 251 | def __init__(self, args, tokenizer): |
| 252 | self.tokenizer = tokenizer |
| 253 | self.start_token = tokenizer.get_command('sop').Id |
| 254 | self.end_token = tokenizer.get_command('eop').Id |
| 255 | self.mask_token = tokenizer.get_command( |
| 256 | 'sMASK').Id if args.task_mask and args.task != 'cmrc' else tokenizer.get_command('MASK').Id |
| 257 | self.pad_token = tokenizer.get_command('pad').Id |
| 258 | self.processors = LogitsProcessorList() |
| 259 | self.mask_pad_token = args.mask_pad_token |
| 260 | if args.min_tgt_length > 0: |
| 261 | processor = MinLengthLogitsProcessor(args.min_tgt_length, self.end_token) |
| 262 | self.processors.append(processor) |
| 263 | if args.no_repeat_ngram_size > 0: |
| 264 | processor = NoRepeatNGramLogitsProcessor(args.no_repeat_ngram_size) |
| 265 | self.processors.append(processor) |
| 266 | |
| 267 | def evaluate(self, model, dataloader, example_dict, args): |
| 268 | """Calculate correct over total answers and return prediction if the |
| 269 | `output_predictions` is true.""" |
| 270 | model.eval() |
| 271 | local_predictions = {} |
| 272 | print_rank_0("Distributed store created") |
| 273 | with torch.no_grad(): |
| 274 | # For all the batches in the dataset. |
| 275 | for idx, data in enumerate(dataloader): |
| 276 | tokens, attention_mask, position_ids = process_batch(data, args) |
| 277 | batch_size = tokens.size(0) |
| 278 | beam_scorer = BeamSearchScorer( |
| 279 | batch_size=batch_size, |
| 280 | max_length=args.out_seq_length, |
| 281 | num_beams=args.num_beams, |
| 282 | device=tokens.device, |
| 283 | length_penalty=args.length_penalty, |
| 284 | do_early_stopping=False, |
| 285 | ) |
| 286 | beam_scores = torch.zeros((batch_size, args.num_beams), dtype=torch.float, device=tokens.device) |
| 287 | beam_scores[:, 1:] = -1e9 |
| 288 | beam_scores = beam_scores.view((batch_size * args.num_beams,)) |
| 289 | # Run the model forward. |
| 290 | counter = 0 |
| 291 | context_length = tokens.size(1) |
| 292 | while counter < args.tgt_seq_length: |
| 293 | if counter == 0: |
| 294 | next_token_logits, *mems = model(tokens, position_ids, attention_mask, return_memory=True) |
| 295 | seq_length = next_token_logits.size(1) |
| 296 | next_token_logits = next_token_logits[:, -1] |
| 297 | next_token_logits = next_token_logits.unsqueeze(1).repeat(1, args.num_beams, 1).view( |
| 298 | batch_size * args.num_beams, -1) |
| 299 | mems = [mem.unsqueeze(1).repeat(1, args.num_beams, 1, 1).view(batch_size * args.num_beams, |
| 300 | seq_length, -1) for mem in mems] |
| 301 | position_ids = tokens.new_ones(batch_size, args.num_beams, 2, 1) |
| 302 | for i, text in enumerate(tokens.tolist()): |
| 303 | mask_pos = text.index(self.mask_token) |
| 304 | position_ids[i, :, 0] = mask_pos |
| 305 | position_ids = position_ids.reshape(batch_size * args.num_beams, 2, 1) |
| 306 | tokens = tokens.new_zeros(batch_size * args.num_beams, 0) |
| 307 | else: |
no outgoing calls
no test coverage detected