(
model: BaseTransformer,
args: argparse.Namespace,
early_stopping_callback=False,
logger=True, # can pass WandbLogger() here
extra_callbacks=[],
checkpoint_callback=None,
logging_callback=None,
**extra_train_kwargs
)
| 274 | |
| 275 | |
| 276 | def generic_train( |
| 277 | model: BaseTransformer, |
| 278 | args: argparse.Namespace, |
| 279 | early_stopping_callback=False, |
| 280 | logger=True, # can pass WandbLogger() here |
| 281 | extra_callbacks=[], |
| 282 | checkpoint_callback=None, |
| 283 | logging_callback=None, |
| 284 | **extra_train_kwargs |
| 285 | ): |
| 286 | # init model |
| 287 | set_seed(args) |
| 288 | odir = Path(model.hparams.output_dir) |
| 289 | odir.mkdir(exist_ok=True) |
| 290 | if checkpoint_callback is None: |
| 291 | checkpoint_callback = pl.callbacks.ModelCheckpoint( |
| 292 | filepath=args.output_dir, prefix="checkpoint", monitor="val_loss", mode="min", save_top_k=1 |
| 293 | ) |
| 294 | if logging_callback is None: |
| 295 | logging_callback = LoggingCallback() |
| 296 | |
| 297 | train_params = {} |
| 298 | |
| 299 | if args.fp16: |
| 300 | train_params["use_amp"] = args.fp16 |
| 301 | train_params["amp_level"] = args.fp16_opt_level |
| 302 | |
| 303 | if args.n_tpu_cores > 0: |
| 304 | global xm |
| 305 | import torch_xla.core.xla_model as xm |
| 306 | |
| 307 | train_params["num_tpu_cores"] = args.n_tpu_cores |
| 308 | train_params["gpus"] = 0 |
| 309 | |
| 310 | if args.gpus > 1: |
| 311 | train_params["distributed_backend"] = "ddp" |
| 312 | |
| 313 | trainer = pl.Trainer( |
| 314 | logger=logger, |
| 315 | accumulate_grad_batches=args.gradient_accumulation_steps, |
| 316 | gpus=args.gpus, |
| 317 | max_epochs=args.num_train_epochs, |
| 318 | early_stop_callback=early_stopping_callback, |
| 319 | gradient_clip_val=args.max_grad_norm, |
| 320 | checkpoint_callback=checkpoint_callback, |
| 321 | callbacks=[logging_callback] + extra_callbacks, |
| 322 | fast_dev_run=args.fast_dev_run, |
| 323 | val_check_interval=args.val_check_interval, |
| 324 | weights_summary=None, |
| 325 | resume_from_checkpoint=args.resume_from_checkpoint, |
| 326 | **train_params, |
| 327 | ) |
| 328 | |
| 329 | if args.do_train: |
| 330 | trainer.fit(model) |
| 331 | trainer.logger.log_hyperparams(args) |
| 332 | trainer.logger.save() |
| 333 | return trainer |
no test coverage detected