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

Class LossTracking

flowmap/loss/loss_tracking.py:23–61  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

21
22
23class LossTracking(Loss[LossTrackingCfg]):
24 def __init__(self, cfg: LossTrackingCfg) -> None:
25 super().__init__(cfg)
26 self.mapping = get_mapping(cfg.mapping)
27
28 def compute_unweighted_loss(
29 self,
30 batch: Batch,
31 flows: Flows,
32 tracks: list[Tracks] | None,
33 model_output: ModelOutput,
34 global_step: int,
35 ) -> Float[Tensor, ""]:
36 # Tracks must be available for the tracking loss.
37 assert tracks is not None
38
39 _, _, _, h, w = batch.videos.shape
40
41 loss_sum = 0
42 valid_sum = 0
43
44 for segment_tracks in tracks:
45 _, f, _, _ = segment_tracks.xy.shape
46 s = segment_tracks.start_frame
47
48 xy_target, visibility = compute_track_flow(
49 model_output.surfaces[:, s : s + f],
50 model_output.extrinsics[:, s : s + f],
51 model_output.intrinsics[:, s : s + f],
52 segment_tracks,
53 )
54 xy_target_gt = rearrange(segment_tracks.xy, "b ft p xy -> b () ft p xy")
55
56 loss = self.mapping.forward(xy_target, xy_target_gt, (h, w)) * visibility
57
58 loss_sum = loss_sum + loss.sum()
59 valid_sum = valid_sum + visibility.sum()
60
61 return loss_sum / (valid_sum or 1)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected