| 132 | |
| 133 | |
| 134 | class Trainer: |
| 135 | def __init__( |
| 136 | self, data_loaders, model_opt_lrsch, |
| 137 | ): |
| 138 | self.model, self.optimizer, self.lr_scheduler = model_opt_lrsch |
| 139 | self.train_loader, self.test_loaders = data_loaders |
| 140 | if config.out_ref: |
| 141 | self.criterion_gdt = nn.BCELoss() |
| 142 | |
| 143 | # Setting Losses |
| 144 | self.pix_loss = PixLoss() |
| 145 | self.cls_loss = ClsLoss() |
| 146 | |
| 147 | # Others |
| 148 | self.loss_log = AverageMeter() |
| 149 | if config.lambda_adv_g: |
| 150 | self.optimizer_d, self.lr_scheduler_d, self.disc, self.adv_criterion = self._load_adv_components() |
| 151 | self.disc_update_for_odd = 0 |
| 152 | |
| 153 | def _load_adv_components(self): |
| 154 | # AIL |
| 155 | from loss import Discriminator |
| 156 | disc = Discriminator(channels=3, img_size=config.size) |
| 157 | if to_be_distributed: |
| 158 | disc = disc.to(device) |
| 159 | disc = DDP(disc, device_ids=[device], broadcast_buffers=False) |
| 160 | else: |
| 161 | disc = disc.to(device) |
| 162 | if config.compile: |
| 163 | disc = torch.compile(disc, mode=['default', 'reduce-overhead', 'max-autotune'][0]) |
| 164 | adv_criterion = nn.BCELoss() |
| 165 | if config.optimizer == 'AdamW': |
| 166 | optimizer_d = optim.AdamW(params=disc.parameters(), lr=config.lr, weight_decay=1e-2) |
| 167 | elif config.optimizer == 'Adam': |
| 168 | optimizer_d = optim.Adam(params=disc.parameters(), lr=config.lr, weight_decay=0) |
| 169 | lr_scheduler_d = torch.optim.lr_scheduler.MultiStepLR( |
| 170 | optimizer_d, |
| 171 | milestones=[lde if lde > 0 else args.epochs + lde + 1 for lde in config.lr_decay_epochs], |
| 172 | gamma=config.lr_decay_rate |
| 173 | ) |
| 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 |