for epoch in range(1, cfg.epochs + 1): model.train() for batch in train_loader: img, lab = batch["image"].to(device), batch["label"].to(device) optimizer.zero_grad(set_to_none=True) with torch.amp.autocast("cuda", enabled=cfg.use_amp): loss = loss_fn(model(img), lab) scaler.scale(loss).backward() scaler.step(optimizer); scaler.update() model.eval(); metric.reset() with torch.no_grad(): for batch in val_loader: img, lab = batch["image"].to(device), batch["label"].to(device) pred = post_pred(model(img)) metric(y_pred=pred, y=lab) val_dice = metric.aggregate().item() if val_dice > best_dice: best_dice = val_dice torch.save(model.state_dict(), cfg.ckpt_path)