Evaluate the model on the validation set(s) and return the losses.
(args, trainer, task, epoch_itr, subsets)
| 216 | |
| 217 | |
| 218 | def validate(args, trainer, task, epoch_itr, subsets): |
| 219 | """Evaluate the model on the validation set(s) and return the losses.""" |
| 220 | |
| 221 | if args.fixed_validation_seed is not None: |
| 222 | # set fixed seed for every validation |
| 223 | utils.set_torch_seed(args.fixed_validation_seed) |
| 224 | |
| 225 | valid_losses = [] |
| 226 | |
| 227 | |
| 228 | for i, subset in enumerate(subsets): |
| 229 | # Initialize data iterator |
| 230 | itr = task.get_batch_iterator( |
| 231 | dataset=task.dataset(subset), |
| 232 | max_tokens=args.max_tokens_valid, |
| 233 | max_sentences=args.max_sentences_valid, |
| 234 | max_positions=utils.resolve_max_positions( |
| 235 | task.max_positions(), |
| 236 | trainer.get_model().max_positions(), |
| 237 | ), |
| 238 | ignore_invalid_inputs=args.skip_invalid_size_inputs_valid_test, |
| 239 | required_batch_size_multiple=args.required_batch_size_multiple, |
| 240 | seed=args.seed, |
| 241 | num_shards=args.distributed_world_size, |
| 242 | shard_id=args.distributed_rank, |
| 243 | num_workers=args.num_workers, |
| 244 | ).next_epoch_itr(shuffle=False) |
| 245 | progress = progress_bar.build_progress_bar( |
| 246 | args, itr, epoch_itr.epoch, |
| 247 | prefix='valid on \'{}\' subset'.format(subset), |
| 248 | no_progress_bar='simple' |
| 249 | ) |
| 250 | |
| 251 | epoch_loss_list = [] |
| 252 | |
| 253 | |
| 254 | # reset validation loss meters |
| 255 | for k in ['valid_loss', 'valid_nll_loss']: |
| 256 | meter = trainer.get_meter(k) |
| 257 | if meter is not None: |
| 258 | meter.reset() |
| 259 | extra_meters = collections.defaultdict(lambda: AverageMeter()) |
| 260 | |
| 261 | for sample in progress: |
| 262 | log_output = trainer.valid_step(sample) |
| 263 | epoch_loss_list += log_output['losses'] |
| 264 | for k, v in log_output.items(): |
| 265 | if k in ['loss', 'nll_loss', 'ntokens', 'nsentences', 'sample_size']: |
| 266 | continue |
| 267 | extra_meters[k].update(v) |
| 268 | |
| 269 | epoch_loss_list = np.array(epoch_loss_list) |
| 270 | #np.save('loss_stats/%s_valid%d_loss.npy'%(args.sess, i), epoch_loss_list) |
| 271 | |
| 272 | # log validation stats |
| 273 | stats = get_valid_stats(trainer, args, extra_meters) |
| 274 | for k, meter in extra_meters.items(): |
| 275 | stats[k] = meter.avg |
no test coverage detected