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