(model: torch.nn.Module,
data_loader, val_loader, optimizer: torch.optim.Optimizer,
epoch: int, start_iter, loss_scaler,
log_writer=None,
args=None)
| 12 | |
| 13 | |
| 14 | def train_one_epoch(model: torch.nn.Module, |
| 15 | data_loader, val_loader, optimizer: torch.optim.Optimizer, |
| 16 | epoch: int, start_iter, loss_scaler, |
| 17 | log_writer=None, |
| 18 | args=None): |
| 19 | model.train(True) |
| 20 | metric_logger = misc.MetricLogger(delimiter=" ") |
| 21 | metric_logger.add_meter('lr', misc.SmoothedValue(window_size=1, fmt='{value:.6f}')) |
| 22 | header = 'Epoch: [{}]'.format(epoch) |
| 23 | print_freq = 10 |
| 24 | |
| 25 | accum_iter = args.accum_iter |
| 26 | |
| 27 | model.zero_grad(set_to_none=True) |
| 28 | |
| 29 | dataset_state = {} |
| 30 | |
| 31 | if log_writer is not None: |
| 32 | print('log_dir: {}'.format(log_writer.log_dir)) |
| 33 | for data_iter_step, (examples, labels, item_states) in enumerate( |
| 34 | metric_logger.log_every(data_loader, print_freq, header, start_iter), start=start_iter |
| 35 | ): |
| 36 | |
| 37 | if data_iter_step % accum_iter == 0: |
| 38 | lr_sched.adjust_learning_rate(optimizer, data_iter_step, args) |
| 39 | |
| 40 | autocast_ctx = { |
| 41 | "bf16": torch.cuda.amp.autocast(dtype=torch.bfloat16), |
| 42 | "fp16": torch.cuda.amp.autocast(dtype=torch.float16), |
| 43 | "tf32": contextlib.nullcontext(), |
| 44 | }[args.precision] |
| 45 | with autocast_ctx: |
| 46 | c_loss, additional_loss_dict = model(examples, labels) |
| 47 | loss = c_loss |
| 48 | for (add_loss, weight) in additional_loss_dict.values(): |
| 49 | loss = loss + add_loss * weight |
| 50 | loss_value = loss.item() |
| 51 | c_loss_value = c_loss.item() |
| 52 | if not math.isfinite(loss_value): |
| 53 | print("Loss is {}, stopping training".format(loss_value)) |
| 54 | sys.exit(1) |
| 55 | |
| 56 | loss /= accum_iter |
| 57 | |
| 58 | update_grad = (data_iter_step + 1) % accum_iter == 0 |
| 59 | grad_norm = loss_scaler( |
| 60 | loss, optimizer, model, |
| 61 | parameters=model.parameters(), |
| 62 | update_grad=update_grad, |
| 63 | clip_grad=None if args.clip_grad <= 0 else args.clip_grad, |
| 64 | ) |
| 65 | |
| 66 | if update_grad: |
| 67 | assert grad_norm is not None |
| 68 | if torch.any(torch.isinf(grad_norm)): |
| 69 | print("grad norm is inf") |
| 70 | else: |
| 71 | metric_logger.update(grad_norm=grad_norm) |
no test coverage detected