(self, batch)
| 174 | return optimizer_d, lr_scheduler_d, disc, adv_criterion |
| 175 | |
| 176 | def _train_batch(self, batch): |
| 177 | inputs = batch[0].to(device) |
| 178 | gts = batch[1].to(device) |
| 179 | class_labels = batch[2].to(device) |
| 180 | scaled_preds, class_preds_lst = self.model(inputs) |
| 181 | if config.out_ref: |
| 182 | (outs_gdt_pred, outs_gdt_label), scaled_preds = scaled_preds |
| 183 | for _idx, (_gdt_pred, _gdt_label) in enumerate(zip(outs_gdt_pred, outs_gdt_label)): |
| 184 | _gdt_pred = nn.functional.interpolate(_gdt_pred, size=_gdt_label.shape[2:], mode='bilinear', align_corners=True).sigmoid() |
| 185 | _gdt_label = _gdt_label.sigmoid() |
| 186 | loss_gdt = self.criterion_gdt(_gdt_pred, _gdt_label) if _idx == 0 else self.criterion_gdt(_gdt_pred, _gdt_label) + loss_gdt |
| 187 | # self.loss_dict['loss_gdt'] = loss_gdt.item() |
| 188 | if None in class_preds_lst: |
| 189 | loss_cls = 0. |
| 190 | else: |
| 191 | loss_cls = self.cls_loss(class_preds_lst, class_labels) * 1.0 |
| 192 | self.loss_dict['loss_cls'] = loss_cls.item() |
| 193 | |
| 194 | # Loss |
| 195 | loss_pix = self.pix_loss(scaled_preds, torch.clamp(gts, 0, 1)) * 1.0 |
| 196 | self.loss_dict['loss_pix'] = loss_pix.item() |
| 197 | # since there may be several losses for sal, the lambdas for them (lambdas_pix) are inside the loss.py |
| 198 | loss = loss_pix + loss_cls |
| 199 | if config.out_ref: |
| 200 | loss = loss + loss_gdt * 1.0 |
| 201 | |
| 202 | if config.lambda_adv_g: |
| 203 | # gen |
| 204 | valid = Variable(torch.cuda.FloatTensor(scaled_preds[-1].shape[0], 1).fill_(1.0), requires_grad=False).to(device) |
| 205 | adv_loss_g = self.adv_criterion(self.disc(scaled_preds[-1] * inputs), valid) * config.lambda_adv_g |
| 206 | loss += adv_loss_g |
| 207 | self.loss_dict['loss_adv'] = adv_loss_g.item() |
| 208 | self.disc_update_for_odd += 1 |
| 209 | self.loss_log.update(loss.item(), inputs.size(0)) |
| 210 | self.optimizer.zero_grad() |
| 211 | loss.backward() |
| 212 | self.optimizer.step() |
| 213 | |
| 214 | if config.lambda_adv_g and self.disc_update_for_odd % 2 == 0: |
| 215 | # disc |
| 216 | fake = Variable(torch.cuda.FloatTensor(scaled_preds[-1].shape[0], 1).fill_(0.0), requires_grad=False).to(device) |
| 217 | self.optimizer_d.zero_grad() |
| 218 | adv_loss_real = self.adv_criterion(self.disc(gts * inputs), valid) |
| 219 | adv_loss_fake = self.adv_criterion(self.disc(scaled_preds[-1].detach() * inputs.detach()), fake) |
| 220 | adv_loss_d = (adv_loss_real + adv_loss_fake) / 2 * config.lambda_adv_d |
| 221 | self.loss_dict['loss_adv_d'] = adv_loss_d.item() |
| 222 | adv_loss_d.backward() |
| 223 | self.optimizer_d.step() |
| 224 | |
| 225 | def train_epoch(self, epoch): |
| 226 | global logger_loss_idx |
no test coverage detected