MCPcopy Create free account
hub / github.com/dayu11/Differentially-Private-Deep-Learning / validate

Function validate

language/bert/bert_code/train.py:218–283  ·  view source on GitHub ↗

Evaluate the model on the validation set(s) and return the losses.

(args, trainer, task, epoch_itr, subsets)

Source from the content-addressed store, hash-verified

216
217
218def 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

Callers 2

mainFunction · 0.85
trainFunction · 0.85

Calls 13

AverageMeterClass · 0.90
get_valid_statsFunction · 0.85
datasetMethod · 0.80
get_modelMethod · 0.80
get_meterMethod · 0.80
next_epoch_itrMethod · 0.45
get_batch_iteratorMethod · 0.45
max_positionsMethod · 0.45
resetMethod · 0.45
valid_stepMethod · 0.45
updateMethod · 0.45
printMethod · 0.45

Tested by

no test coverage detected