| 24 | |
| 25 | |
| 26 | class LossFlow(Loss[LossFlowCfg]): |
| 27 | def __init__(self, cfg: LossFlowCfg) -> None: |
| 28 | super().__init__(cfg) |
| 29 | self.mapping = get_mapping(cfg.mapping) |
| 30 | |
| 31 | def compute_unweighted_loss( |
| 32 | self, |
| 33 | batch: Batch, |
| 34 | flows: Flows, |
| 35 | tracks: list[Tracks] | None, |
| 36 | model_output: ModelOutput, |
| 37 | global_step: int, |
| 38 | ) -> Float[Tensor, ""]: |
| 39 | _, _, _, h, w = batch.videos.shape |
| 40 | device = batch.videos.device |
| 41 | xy, _ = sample_image_grid((h, w), device) |
| 42 | |
| 43 | loss_sum = 0 |
| 44 | valid_sum = 0 |
| 45 | |
| 46 | # Compute a loss based on forward flow. |
| 47 | xy_flowed_forward = compute_forward_flow( |
| 48 | model_output.surfaces, |
| 49 | model_output.extrinsics, |
| 50 | model_output.intrinsics, |
| 51 | ) |
| 52 | forward_loss = self.mapping.forward( |
| 53 | xy_flowed_forward - xy, flows.forward, (h, w) |
| 54 | ) |
| 55 | loss_sum = loss_sum + (forward_loss * flows.forward_mask).sum() |
| 56 | valid_sum = valid_sum + flows.forward_mask.sum() |
| 57 | |
| 58 | # Compute a loss based on backward flow. |
| 59 | xy_flowed_backward = compute_backward_flow( |
| 60 | model_output.surfaces, |
| 61 | model_output.extrinsics, |
| 62 | model_output.intrinsics, |
| 63 | ) |
| 64 | backward_loss = self.mapping.forward( |
| 65 | xy_flowed_backward - xy, flows.backward, (h, w) |
| 66 | ) |
| 67 | loss_sum = loss_sum + (backward_loss * flows.backward_mask).sum() |
| 68 | valid_sum = valid_sum + flows.backward_mask.sum() |
| 69 | |
| 70 | return loss_sum / (valid_sum or 1) |
nothing calls this directly
no outgoing calls
no test coverage detected