MCPcopy Create free account
hub / github.com/SLDGroup/EMCAD / train

Function train

train_polyp.py:107–171  ·  view source on GitHub ↗
(train_loader, model, optimizer, epoch, opt, model_name)

Source from the content-addressed store, hash-verified

105 return DSC / total_images, IOU / total_images, total_images
106
107def 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

Callers 1

train_polyp.pyFile · 0.85

Calls 7

updateMethod · 0.95
showMethod · 0.95
AvgMeterClass · 0.90
clip_gradientFunction · 0.90
structure_lossFunction · 0.85
stepMethod · 0.80
testFunction · 0.70

Tested by

no test coverage detected