MCPcopy Create free account
hub / github.com/dcharatan/flowmap / LossFlow

Class LossFlow

flowmap/loss/loss_flow.py:26–70  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24
25
26class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected