(self, step)
| 161 | |
| 162 | |
| 163 | def optimize_parameters(self, step): |
| 164 | if self.opt['train']['fix_some_part'] and step < self.opt['train']['fix_some_part']: |
| 165 | self.set_params_lr_zero() |
| 166 | |
| 167 | self.netG.zero_grad() ################################################# new add |
| 168 | self.optimizer_G.zero_grad() |
| 169 | |
| 170 | # LR_right = self.varright_L |
| 171 | out,instf,fusef = self.netG(self.var_L) |
| 172 | |
| 173 | if self.opt['train']['distill']: |
| 174 | var_fake = self.real_H |
| 175 | # fakeout,gtinst, = self.netG(var_fake) |
| 176 | with torch.no_grad(): |
| 177 | self.netG_Pre.eval() |
| 178 | _,gtinstf,gtfusef = self.netG_Pre(var_fake.detach()) |
| 179 | |
| 180 | |
| 181 | gt = self.real_H |
| 182 | #print('gt range:', max(gt), min(gt)) |
| 183 | l_total = self.mse(out, gt) |
| 184 | # 0809 fail low psnr |
| 185 | # l_total = self.mse(out, gt) + 5*self.bce(out, gt) |
| 186 | |
| 187 | # + 1.2*self.cri_pix(out, out_pre) |
| 188 | if self.opt['train']['distill']: |
| 189 | l_total += self.opt['train']['distill_coff']*(self.cri_pix(instf,gtinstf.detach())+self.cri_pix(fusef,gtfusef.detach())) |
| 190 | |
| 191 | if self.opt['train']['ewc']: |
| 192 | for i, w in enumerate(self.netG.parameters()): |
| 193 | l_total += self.opt['train']['ewc_coff']/2 * torch.sum(torch.mul(self.Importance_Pre[i], torch.abs(w - self.Star_vals_Pre[i])))\ |
| 194 | + self.opt['train']['ewc_coff']/4 * torch.square(torch.sum(torch.mul(self.Importance_Pre[i], torch.abs(w - self.Star_vals_Pre[i])))) |
| 195 | |
| 196 | |
| 197 | l_total.backward() |
| 198 | self.optimizer_G.step() |
| 199 | self.fake_H = out |
| 200 | psnr = psnr_np(self.fake_H.detach(), self.real_H.detach()) |
| 201 | |
| 202 | # set log |
| 203 | self.log_dict['psnr'] = psnr.item() |
| 204 | self.log_dict['l_total'] = l_total.item() |
| 205 | |
| 206 | ################# test function and it helpers |
| 207 | def feed_val_data(self, data, need_GT=True): |
nothing calls this directly
no test coverage detected