(
self,
batch: Batch,
flow_shape: tuple[int, int],
)
| 80 | return rearrange((1 - deltas) ** 8, "(b f) h w -> b f h w", b=b, f=f - 1) |
| 81 | |
| 82 | def compute_bidirectional_flow( |
| 83 | self, |
| 84 | batch: Batch, |
| 85 | flow_shape: tuple[int, int], |
| 86 | ) -> Flows: |
| 87 | forward = self.forward(batch.videos) |
| 88 | forward_mask = self.compute_consistency_mask(batch.videos, forward) |
| 89 | forward = self.rescale_flow(forward, flow_shape) |
| 90 | forward_mask = self.rescale_mask(forward_mask, flow_shape) |
| 91 | |
| 92 | backward_videos = batch.videos.flip(dims=(1,)) |
| 93 | |
| 94 | backward = self.forward(backward_videos) |
| 95 | backward_mask = self.compute_consistency_mask(backward_videos, backward) |
| 96 | backward = self.rescale_flow(backward, flow_shape) |
| 97 | backward_mask = self.rescale_mask(backward_mask, flow_shape) |
| 98 | |
| 99 | backward = backward.flip(dims=(1,)) |
| 100 | backward_mask = backward_mask.flip(dims=(1,)) |
| 101 | |
| 102 | return Flows(forward, backward, forward_mask, backward_mask) |
no test coverage detected