Evaluation.
(data_loader, model, eval_metric)
| 107 | |
| 108 | |
| 109 | def evaluate(data_loader, model, eval_metric): |
| 110 | """Evaluation.""" |
| 111 | args = get_args() |
| 112 | |
| 113 | # Turn on evaluation mode which disables dropout. |
| 114 | model.eval() |
| 115 | |
| 116 | total_output = 0.0 |
| 117 | with torch.no_grad(): |
| 118 | # For all the batches in the dataset. |
| 119 | for iteration, batch in enumerate(data_loader): |
| 120 | if iteration % args.log_interval == 0: |
| 121 | print_rank_0('> working on iteration: {}'.format(iteration)) |
| 122 | # Forward evaluation. |
| 123 | output = forward_step(batch, model, eval_metric) |
| 124 | |
| 125 | # Reduce across processes. |
| 126 | torch.distributed.all_reduce(output, |
| 127 | group=mpu.get_data_parallel_group()) |
| 128 | |
| 129 | total_output += output |
| 130 | |
| 131 | return total_output |
| 132 | |
| 133 | |
| 134 | def evaluate_and_print_results(task, data_loader, model, eval_metric): |
no test coverage detected