MCPcopy Create free account
hub / github.com/ZhengPeng7/BiRefNet / _train_batch

Method _train_batch

train.py:176–223  ·  view source on GitHub ↗
(self, batch)

Source from the content-addressed store, hash-verified

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

Callers 1

train_epochMethod · 0.95

Calls 2

updateMethod · 0.80
stepMethod · 0.45

Tested by

no test coverage detected