MCPcopy Create free account
hub / github.com/ant-research/CoDeF / training_step

Method training_step

train.py:286–388  ·  view source on GitHub ↗
(self, batch, batch_idx)

Source from the content-addressed store, hash-verified

284 pin_memory=True)
285
286 def training_step(self, batch, batch_idx):
287 # Fetch training data.
288 rgbs = batch['rgbs']
289 ts_w = batch['ts_w']
290 grid = batch['grid']
291 mk = batch['masks']
292 flows = batch['flows']
293 grid_c = batch['grid_c']
294 ref_batch = batch['reference']
295 self.seq_len = batch['seq_len']
296
297 loss = 0
298 rgbs_flattend = rearrange(rgbs, 'b h w c -> (b h w) c')
299
300 # Forward the model.
301 ret = self.forward(ts_w,
302 grid,
303 self.hparams.encode_w,
304 self.global_step,
305 flows=flows)
306
307 # Mannually set a reference frame.
308 if self.hparams.ref_step < 0: self.hparams.step = 1e10
309 if (self.hparams.ref_idx is not None
310 and self.global_step < self.hparams.ref_step):
311 rgbs_c_flattend = rearrange(ref_batch[0],
312 'b h w c -> (b h w) c')
313 ret_c = self(ts_w, grid, False, self.global_step, flows=flows)
314
315 # Loss computation.
316 for i in range(self.num_models):
317 results = ret.rgbs[i]
318 mk_t = rearrange(mk[i], 'b h w c -> (b h w) c')
319 mk_t = mk_t.sum(dim=-1) > 0.05
320
321 if (self.hparams.ref_idx is not None
322 and self.global_step < self.hparams.ref_step):
323 mk_c_t = rearrange(ref_batch[1][i], 'b h w c -> (b h w) c')
324 mk_c_t = mk_c_t.sum(dim=-1) > 0.05
325
326 # Background regularization.
327 if self.hparams.bg_loss:
328 mk1 = torch.logical_not(mk_t)
329 if self.hparams.self_bg:
330 grid_flattened = rgbs_flattend
331 else:
332 grid_flattened = rearrange(grid, 'b n c -> (b n) c')
333 grid_flattened = torch.cat(
334 [grid_flattened, grid_flattened[:, :1]], -1)
335
336 if self.hparams.bg_loss and self.hparams.mask_dir:
337 loss = loss + self.hparams.bg_loss * self.color_loss(
338 results[mk1], grid_flattened[mk1])
339
340 # MSE color loss.
341 loss = loss + self.color_loss(results[mk_t],
342 rgbs_flattend[mk_t])
343

Callers

nothing calls this directly

Calls 4

forwardMethod · 0.95
compute_gradient_lossFunction · 0.90
psnrFunction · 0.90
get_learning_rateFunction · 0.90

Tested by

no test coverage detected