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

Function evaluate

pretrain_glm.py:387–438  ·  view source on GitHub ↗

Evaluation.

(data_iterator, model, args, timers, forward_step_func, verbose=False)

Source from the content-addressed store, hash-verified

385
386
387def evaluate(data_iterator, model, args, timers, forward_step_func, verbose=False):
388 """Evaluation."""
389 # Turn on evaluation mode which disables dropout.
390 model.eval()
391
392 total_lm_loss, total_gpt_loss, total_bert_loss, total_sent_loss, total_multi_loss = 0, 0, 0, 0, 0
393 gpt_iters, bert_iters, sent_iters, multi_iters = 0, 0, 0, 0
394 mems = []
395 with torch.no_grad():
396 iteration = 0
397 while iteration < args.eval_iters:
398 iteration += 1
399 if verbose and iteration % args.log_interval == 0:
400 print_rank_0('Evaluating iter {}/{}'.format(iteration, args.eval_iters))
401 # Forward evaluation.
402 lm_loss, mems, mode = forward_step_func(data_iterator, model, args, timers, mems=mems)
403
404 ''&#x27;when contiguous memory optimizations are enabled, the buffers
405 allocated by the optimizations are deallocated during backward pass
406 in the absence of backward pass the buffers should be reset after each
407 forward pass''&#x27;
408 if args.deepspeed and args.deepspeed_activation_checkpointing:
409 deepspeed.checkpointing.reset()
410
411 lm_loss = lm_loss.data.detach().float().item()
412 total_lm_loss += lm_loss
413 if mode == 'gpt':
414 total_gpt_loss += lm_loss
415 gpt_iters += 1
416 elif mode == 'bert':
417 total_bert_loss += lm_loss
418 bert_iters += 1
419 elif mode == 'sentence':
420 total_sent_loss += lm_loss
421 sent_iters += 1
422 elif mode == 'multi-task':
423 total_multi_loss += lm_loss
424 multi_iters += 1
425 # Move model back to the train mode.
426 model.train()
427 # Reduce across processes.
428 loss_data = torch.cuda.FloatTensor(
429 [total_lm_loss, total_gpt_loss, total_bert_loss, total_sent_loss, total_multi_loss, gpt_iters, bert_iters,
430 sent_iters, multi_iters])
431 torch.distributed.all_reduce(loss_data, group=mpu.get_data_parallel_group())
432 loss_data = loss_data.tolist()
433 total_lm_loss = loss_data[0] / args.eval_iters / (args.world_size / args.model_parallel_size)
434 total_gpt_loss = loss_data[1] / loss_data[5] if loss_data[5] > 0 else 0
435 total_bert_loss = loss_data[2] / loss_data[6] if loss_data[6] > 0 else 0
436 total_sent_loss = loss_data[3] / loss_data[7] if loss_data[7] > 0 else 0
437 total_multi_loss = loss_data[4] / loss_data[8] if loss_data[8] > 0 else 0
438 return total_lm_loss, total_gpt_loss, total_bert_loss, total_sent_loss, total_multi_loss
439
440
441def evaluate_and_print_results(prefix, data_iterator, model,

Callers 1

Calls 2

print_rank_0Function · 0.90
resetMethod · 0.45

Tested by

no test coverage detected