(self, args, tokenizer)
| 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 |
nothing calls this directly
no test coverage detected