MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / train_one_epoch

Function train_one_epoch

engine.py:34–179  ·  view source on GitHub ↗
(model: torch.nn.Module,
                    criterion: torch.nn.Module,
                    data_loader: Iterable,
                    optimizer: torch.optim.Optimizer,
                    device: torch.device,
                    epoch: int,
                    max_norm: float = 0,
                    wo_class_error=False,
                    lr_scheduler=None,
                    args=None,
                    logger=None,
                    ema_m=None,
                    tf_writer=None)

Source from the content-addressed store, hash-verified

32 return items
33
34def train_one_epoch(model: torch.nn.Module,
35 criterion: torch.nn.Module,
36 data_loader: Iterable,
37 optimizer: torch.optim.Optimizer,
38 device: torch.device,
39 epoch: int,
40 max_norm: float = 0,
41 wo_class_error=False,
42 lr_scheduler=None,
43 args=None,
44 logger=None,
45 ema_m=None,
46 tf_writer=None):
47 scaler = torch.cuda.amp.GradScaler(enabled=args.amp)
48
49 try:
50 need_tgt_for_training = args.use_dn
51 except:
52 need_tgt_for_training = False
53
54 model.train()
55 criterion.train()
56 # criterion_smpl.to(device)
57 metric_logger = utils.MetricLogger(delimiter=' ')
58 metric_logger.add_meter(
59 'lr', utils.SmoothedValue(window_size=1, fmt='{value:.6f}'))
60 if not wo_class_error:
61 metric_logger.add_meter(
62 'class_error', utils.SmoothedValue(window_size=1,
63 fmt='{value:.2f}'))
64 header = 'Epoch: [{}]'.format(epoch)
65 print_freq = 10
66
67 _cnt = 0
68
69 for step_i, data_batch in enumerate(metric_logger.log_every(data_loader,
70 print_freq,
71 header,
72 logger=logger)):
73 with torch.cuda.amp.autocast(enabled=args.amp):
74 if need_tgt_for_training:
75 outputs, targets, data_batch_nc = model(data_batch)
76 else:
77 outputs, targets, data_batch_nc = model(data_batch)
78
79 ['hand_kp3d_4', 'face_kp3d_4', 'hand_kp2d_4',]
80 loss_dict = criterion(outputs, targets, data_batch=data_batch_nc)
81 weight_dict = criterion.weight_dict
82
83 for k,v in weight_dict.items():
84 for n in ['hand_kp3d_4', 'face_kp3d_4', 'hand_kp2d_4']:
85 if n in k:
86 weight_dict[k] = weight_dict[k]/10
87
88 losses = sum(loss_dict[k] * weight_dict[k]
89 for k in loss_dict.keys() if k in weight_dict)
90
91 loss_dict_reduced = utils.reduce_dict(loss_dict)

Callers 1

mainFunction · 0.90

Calls 15

add_meterMethod · 0.95
log_everyMethod · 0.95
updateMethod · 0.95
round_floatFunction · 0.85
backwardMethod · 0.80
parametersMethod · 0.80
writeMethod · 0.80
printFunction · 0.50
trainMethod · 0.45
itemsMethod · 0.45
keysMethod · 0.45

Tested by

no test coverage detected