Loading train.py +13 −17 Changes for train.py: 13 added lines, 17 removed lines. Original line number Diff line number Diff line Loading @@ -88,19 +88,20 @@ def main(args): first_val_loss = None train_meter = StopwatchMeter() train_meter.start() valid_subsets = args.valid_subset.split(',') while lr > args.min_lr and epoch <= max_epoch and trainer.get_num_updates() < max_update: # train for one epoch train(args, trainer, next_ds, epoch, dataset) if epoch % args.validate_interval == 0: first_val_loss = val_loss(args, trainer, dataset, epoch) valid_losses = validate(args, trainer, dataset, valid_subsets, epoch) # only use first validation loss to update the learning rate lr = trainer.lr_step(epoch, first_val_loss) lr = trainer.lr_step(epoch, valid_losses[0]) # save checkpoint if epoch % args.save_interval == 0: save_checkpoint(args, trainer, epoch, end_of_epoch=True, val_loss=first_val_loss) save_checkpoint(args, trainer, epoch, end_of_epoch=True, val_loss=valid_losses[0]) epoch += 1 next_ds = next(train_dataloader) Loading Loading @@ -135,6 +136,7 @@ def train(args, trainer, itr, epoch, dataset): update_freq = args.update_freq[-1] extra_meters = collections.defaultdict(lambda: AverageMeter()) first_valid = args.valid_subset.split(',')[0] max_update = args.max_update or math.inf num_batches = len(itr) progress = progress_bar.build_progress_bar(args, itr, epoch, no_progress_bar='simple') Loading Loading @@ -164,8 +166,8 @@ def train(args, trainer, itr, epoch, dataset): num_updates = trainer.get_num_updates() if args.save_interval_updates > 0 and num_updates % args.save_interval_updates == 0: first_val_loss = val_loss(args, trainer, dataset, epoch, num_updates) save_checkpoint(args, trainer, epoch, end_of_epoch=False, val_loss=first_val_loss) valid_losses = validate(args, trainer, dataset, [first_valid], epoch) save_checkpoint(args, trainer, epoch, end_of_epoch=False, val_loss=valid_losses[0]) if num_updates >= max_update: break Loading Loading @@ -201,9 +203,10 @@ def get_training_stats(trainer): return stats def validate(args, trainer, dataset, subset, epoch, num_updates): """Evaluate the model on the validation set and return the average loss.""" def validate(args, trainer, dataset, subsets, epoch): """Evaluate the model on the validation set(s) and return the losses.""" valid_losses = [] for subset in subsets: # Initialize dataloader max_positions_valid = ( trainer.get_model().max_encoder_positions(), Loading Loading @@ -246,7 +249,8 @@ def validate(args, trainer, dataset, subset, epoch, num_updates): stats[k] = meter.avg progress.print(stats) return stats['valid_loss'] valid_losses.append(stats['valid_loss']) return valid_losses def get_valid_stats(trainer): Loading @@ -271,14 +275,6 @@ def get_perplexity(loss): return float('inf') def val_loss(args, trainer, dataset, epoch, num_updates=None): # evaluate on validate set subsets = args.valid_subset.split(',') # we want to validate all subsets so the results get printed out, but return only the first losses = [validate(args, trainer, dataset, subset, epoch, num_updates) for subset in subsets] return losses[0] if len(losses) > 0 else None def save_checkpoint(args, trainer, epoch, end_of_epoch, val_loss): if args.no_save or args.distributed_rank > 0: return Loading Loading
train.py +13 −17 Changes for train.py: 13 added lines, 17 removed lines. Original line number Diff line number Diff line Loading @@ -88,19 +88,20 @@ def main(args): first_val_loss = None train_meter = StopwatchMeter() train_meter.start() valid_subsets = args.valid_subset.split(',') while lr > args.min_lr and epoch <= max_epoch and trainer.get_num_updates() < max_update: # train for one epoch train(args, trainer, next_ds, epoch, dataset) if epoch % args.validate_interval == 0: first_val_loss = val_loss(args, trainer, dataset, epoch) valid_losses = validate(args, trainer, dataset, valid_subsets, epoch) # only use first validation loss to update the learning rate lr = trainer.lr_step(epoch, first_val_loss) lr = trainer.lr_step(epoch, valid_losses[0]) # save checkpoint if epoch % args.save_interval == 0: save_checkpoint(args, trainer, epoch, end_of_epoch=True, val_loss=first_val_loss) save_checkpoint(args, trainer, epoch, end_of_epoch=True, val_loss=valid_losses[0]) epoch += 1 next_ds = next(train_dataloader) Loading Loading @@ -135,6 +136,7 @@ def train(args, trainer, itr, epoch, dataset): update_freq = args.update_freq[-1] extra_meters = collections.defaultdict(lambda: AverageMeter()) first_valid = args.valid_subset.split(',')[0] max_update = args.max_update or math.inf num_batches = len(itr) progress = progress_bar.build_progress_bar(args, itr, epoch, no_progress_bar='simple') Loading Loading @@ -164,8 +166,8 @@ def train(args, trainer, itr, epoch, dataset): num_updates = trainer.get_num_updates() if args.save_interval_updates > 0 and num_updates % args.save_interval_updates == 0: first_val_loss = val_loss(args, trainer, dataset, epoch, num_updates) save_checkpoint(args, trainer, epoch, end_of_epoch=False, val_loss=first_val_loss) valid_losses = validate(args, trainer, dataset, [first_valid], epoch) save_checkpoint(args, trainer, epoch, end_of_epoch=False, val_loss=valid_losses[0]) if num_updates >= max_update: break Loading Loading @@ -201,9 +203,10 @@ def get_training_stats(trainer): return stats def validate(args, trainer, dataset, subset, epoch, num_updates): """Evaluate the model on the validation set and return the average loss.""" def validate(args, trainer, dataset, subsets, epoch): """Evaluate the model on the validation set(s) and return the losses.""" valid_losses = [] for subset in subsets: # Initialize dataloader max_positions_valid = ( trainer.get_model().max_encoder_positions(), Loading Loading @@ -246,7 +249,8 @@ def validate(args, trainer, dataset, subset, epoch, num_updates): stats[k] = meter.avg progress.print(stats) return stats['valid_loss'] valid_losses.append(stats['valid_loss']) return valid_losses def get_valid_stats(trainer): Loading @@ -271,14 +275,6 @@ def get_perplexity(loss): return float('inf') def val_loss(args, trainer, dataset, epoch, num_updates=None): # evaluate on validate set subsets = args.valid_subset.split(',') # we want to validate all subsets so the results get printed out, but return only the first losses = [validate(args, trainer, dataset, subset, epoch, num_updates) for subset in subsets] return losses[0] if len(losses) > 0 else None def save_checkpoint(args, trainer, epoch, end_of_epoch, val_loss): if args.no_save or args.distributed_rank > 0: return Loading