(train_loader, model, optimizer, epoch, opt, model_name)
| 105 | return DSC / total_images, IOU / total_images, total_images |
| 106 | |
| 107 | def train(train_loader, model, optimizer, epoch, opt, model_name): |
| 108 | model.train() |
| 109 | global best, test_dice_at_best_val, total_train_time, dict_plot |
| 110 | |
| 111 | epoch_start = time.time() |
| 112 | loss_record = AvgMeter() |
| 113 | size_rates = [0.75, 1, 1.25] |
| 114 | total_step = len(train_loader) |
| 115 | |
| 116 | for i, (images, gts) in enumerate(train_loader, start=1): |
| 117 | for rate in size_rates: |
| 118 | optimizer.zero_grad() |
| 119 | images, gts = Variable(images).cuda(), Variable(gts).float().cuda() |
| 120 | |
| 121 | if rate != 1: |
| 122 | trainsize = int(round(opt.img_size * rate / 32) * 32) |
| 123 | images = F.interpolate(images, size=(trainsize, trainsize), mode='bilinear', align_corners=True) |
| 124 | gts = F.interpolate(gts, size=(trainsize, trainsize), mode='nearest') |
| 125 | |
| 126 | P = model(images) |
| 127 | if not isinstance(P, list): |
| 128 | P = [P] |
| 129 | loss_p1 = structure_loss(P[0], gts) |
| 130 | loss_p2 = structure_loss(P[1], gts) |
| 131 | loss_p3 = structure_loss(P[2], gts) |
| 132 | loss_p4 = structure_loss(P[3], gts) |
| 133 | loss_p1234 = structure_loss(P[0]+P[1]+P[2]+P[3], gts) |
| 134 | |
| 135 | weights = [1, 1, 1, 1, 1] |
| 136 | loss = weights[0]*loss_p1 + weights[1]*loss_p2 + weights[2]*loss_p3 + weights[3]*loss_p4 + weights[4]*loss_p1234 |
| 137 | |
| 138 | loss.backward() |
| 139 | clip_gradient(optimizer, opt.clip) |
| 140 | optimizer.step() |
| 141 | |
| 142 | if rate == 1: |
| 143 | loss_record.update(loss.data, opt.batchsize) |
| 144 | |
| 145 | if i % 100 == 0 or i == total_step: |
| 146 | print(f'{datetime.now()} Epoch [{epoch:03d}/{opt.epoch:03d}], Step [{i:04d}/{total_step:04d}], ' |
| 147 | f'LR: {optimizer.param_groups[0]["lr"]:.6f}, Loss: {loss_record.show():.4f}') |
| 148 | |
| 149 | total_train_time += (time.time() - epoch_start) |
| 150 | |
| 151 | # Save Last |
| 152 | save_path = opt.train_save |
| 153 | os.makedirs(save_path, exist_ok=True) |
| 154 | torch.save(model.state_dict(), os.path.join(save_path, f"{model_name}-last.pth")) |
| 155 | |
| 156 | # Validation and Testing |
| 157 | epoch_results = {} |
| 158 | for ds in ['test', 'val']: |
| 159 | d_dice, d_iou, _ = test(model, opt.test_path, ds, opt) |
| 160 | epoch_results[ds] = d_dice |
| 161 | logging.info(f'Epoch: {epoch}, Dataset: {ds}, Dice: {d_dice:.4f}, IoU: {d_iou:.4f}') |
| 162 | print(f'Epoch: {epoch}, Dataset: {ds}, Dice: {d_dice:.4f}, IoU: {d_iou:.4f}') |
| 163 | dict_plot[ds].append(d_dice) |
| 164 |
no test coverage detected