(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)
| 32 | return items |
| 33 | |
| 34 | def 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) |
no test coverage detected