(
self, data_loaders, model_opt_lrsch,
)
| 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 |
nothing calls this directly
no test coverage detected