MCPcopy Create free account
hub / github.com/OpenGVLab/UniFormerV2 / train_epoch

Function train_epoch

tools/train_net.py:29–186  ·  view source on GitHub ↗

Perform the video training for one epoch. Args: train_loader (loader): video training loader. model (model): the video model to train. optimizer (optim): the optimizer to perform optimization on the model's parameters. loss_scaler (scaler): scaler

(
    train_loader, model, optimizer, loss_scaler, train_meter, cur_epoch, cfg, writer=None
)

Source from the content-addressed store, hash-verified

27
28
29def train_epoch(
30 train_loader, model, optimizer, loss_scaler, train_meter, cur_epoch, cfg, writer=None
31):
32 """
33 Perform the video training for one epoch.
34 Args:
35 train_loader (loader): video training loader.
36 model (model): the video model to train.
37 optimizer (optim): the optimizer to perform optimization on the model's
38 parameters.
39 loss_scaler (scaler): scaler for loss.
40 train_meter (TrainMeter): training meters to log the training performance.
41 cur_epoch (int): current epoch of training.
42 cfg (CfgNode): configs. Details can be found in
43 slowfast/config/defaults.py
44 writer (TensorboardWriter, optional): TensorboardWriter object
45 to writer Tensorboard log.
46 """
47 # Enable train mode.
48 model.train()
49 train_meter.iter_tic()
50 data_size = len(train_loader)
51
52 if cfg.MIXUP.ENABLE:
53 mixup_fn = MixUp(
54 mixup_alpha=cfg.MIXUP.ALPHA,
55 cutmix_alpha=cfg.MIXUP.CUTMIX_ALPHA,
56 mix_prob=cfg.MIXUP.PROB,
57 switch_prob=cfg.MIXUP.SWITCH_PROB,
58 label_smoothing=cfg.MIXUP.LABEL_SMOOTH_VALUE,
59 num_classes=cfg.MODEL.NUM_CLASSES,
60 )
61
62 for cur_iter, (inputs, labels, _, meta) in enumerate(train_loader):
63 # Transfer the data to the current GPU device.
64 if cfg.NUM_GPUS:
65 if isinstance(inputs, (list,)):
66 for i in range(len(inputs)):
67 inputs[i] = inputs[i].cuda(non_blocking=True)
68 else:
69 inputs = inputs.cuda(non_blocking=True)
70 labels = labels.cuda()
71 for key, val in meta.items():
72 if isinstance(val, (list,)):
73 for i in range(len(val)):
74 val[i] = val[i].cuda(non_blocking=True)
75 else:
76 meta[key] = val.cuda(non_blocking=True)
77
78 # Update the learning rate.
79 lr = optim.get_epoch_lr(cur_epoch + float(cur_iter) / data_size, cfg)
80 optim.set_lr(optimizer, lr)
81
82 train_meter.data_toc()
83 if cfg.MIXUP.ENABLE:
84 samples, labels = mixup_fn(inputs[0], labels)
85 inputs[0] = samples
86

Callers 1

trainFunction · 0.85

Calls 9

MixUpClass · 0.90
add_scalarsMethod · 0.80
iter_ticMethod · 0.45
data_tocMethod · 0.45
update_statsMethod · 0.45
iter_tocMethod · 0.45
log_iter_statsMethod · 0.45
log_epoch_statsMethod · 0.45
resetMethod · 0.45

Tested by

no test coverage detected