MCPcopy Create free account
hub / github.com/THUDM/GLM / DecoderEvaluater

Class DecoderEvaluater

tasks/seq2seq/evaluate.py:250–390  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

248
249
250class 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:

Callers 1

metrics_func_providerFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected