This commit is contained in:
sshleifer
2020-05-20 18:14:27 -04:00
parent 67f9553d4b
commit f281724e92
2 changed files with 10 additions and 3 deletions
+8 -2
View File
@@ -206,6 +206,7 @@ class LoggingCallback(pl.Callback):
if not pl_module.is_logger():
return
metrics = trainer.callback_metrics
trainer.logger.log_metrics({'type_path':metrics})
# Log results
od = Path(pl_module.hparams.output_dir)
output_val_results_file = od / f"{type_path}_results.txt"
@@ -237,7 +238,7 @@ class LoggingCallback(pl.Callback):
def on_test_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule):
return self._do_work(trainer, pl_module, "test")
from pytorch_lightning.loggers import WandbLogger
def add_generic_args(parser, root_dir):
parser.add_argument(
"--output_dir",
@@ -286,7 +287,11 @@ def generic_train(model: BaseTransformer, args: argparse.Namespace, extra_callba
checkpoint_callback = pl.callbacks.ModelCheckpoint(
filepath=args.output_dir, prefix="checkpoint", monitor="val_loss", mode="min", save_top_k=5
)
if args.output_dir.startswith('/var/') or args.fast_dev_run:
logger = True
else:
logger = WandbLogger()
logger.log_hyperparams(args)
train_params = dict(
accumulate_grad_batches=args.gradient_accumulation_steps,
gpus=args.n_gpu,
@@ -297,6 +302,7 @@ def generic_train(model: BaseTransformer, args: argparse.Namespace, extra_callba
callbacks=[LoggingCallback()] + extra_callbacks,
fast_dev_run=args.fast_dev_run,
val_check_interval=args.val_check_interval,
logger=logger,
)
if args.fp16:
+2 -1
View File
@@ -268,12 +268,13 @@ class SummarizationDistiller(SummarizationTrainer):
def freeze_stuff(self):
freeze_part(self.model.model.encoder)
freeze_part(self.model.teacher)
freeze_part(self.model.model.shared)
d = self.model.model.decoder
freeze_part(d.embed_positions)
freeze_part(d.embed_tokens)
self.model.teacher.encoder = None
freeze_part(self.model.teacher)
def _step(self, batch):
# assert is_frozen(self.model.teacher)