MCPcopy Create free account
hub / github.com/RylonW/DocNLC / optimize_parameters

Method optimize_parameters

models/SIEN_model.py:163–204  ·  view source on GitHub ↗
(self, step)

Source from the content-addressed store, hash-verified

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):

Callers

nothing calls this directly

Calls 2

set_params_lr_zeroMethod · 0.95
psnr_npFunction · 0.90

Tested by

no test coverage detected