MCPcopy Create free account
hub / github.com/huggingface/transformers / generic_train

Function generic_train

examples/lightning_base.py:276–333  ·  view source on GitHub ↗
(
    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
)

Source from the content-addressed store, hash-verified

274
275
276def 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

Callers 4

evaluate_checkpointFunction · 0.90
mainFunction · 0.90
run_pl_ner.pyFile · 0.90
run_pl_glue.pyFile · 0.90

Calls 3

LoggingCallbackClass · 0.85
set_seedFunction · 0.70
saveMethod · 0.45

Tested by

no test coverage detected