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

Function compute_track_flow

flowmap/model/projection.py:255–298  ·  view source on GitHub ↗
(
    surfaces: Float[Tensor, "batch frame height width xyz=3"],
    extrinsics: Float[Tensor, "batch frame 4 4"],
    intrinsics: Float[Tensor, "batch frame 3 3"],
    tracks: Tracks,
)

Source from the content-addressed store, hash-verified

253
254
255def compute_track_flow(
256 surfaces: Float[Tensor, "batch frame height width xyz=3"],
257 extrinsics: Float[Tensor, "batch frame 4 4"],
258 intrinsics: Float[Tensor, "batch frame 3 3"],
259 tracks: Tracks,
260) -> tuple[
261 Float[Tensor, "batch frame_source frame_target point 2"], # flow
262 Bool[Tensor, "batch frame_source frame_target point"], # visibility
263]:
264 # Sample the surfaces at the track locations.
265 b, f, _, _, _ = surfaces.shape
266 xyz = F.grid_sample(
267 rearrange(surfaces, "b f h w xyz -> (b f) xyz h w"),
268 rearrange(tracks.xy * 2 - 1, "b f p xy -> (b f) () p xy"),
269 mode="bilinear",
270 padding_mode="border",
271 align_corners=False,
272 )
273 xyz = rearrange(xyz, "(b f) xy () p -> b f p xy", b=b, f=f)
274
275 # Add singleton dimensions so that everything broadcasts to the following shape:
276 # (b = batch, fs = source frame, ft = target frame, p = point)
277 xy_source = rearrange(tracks.xy, "b fs p xy -> b fs () p xy")
278 xyz_source = rearrange(xyz, "b fs p xyz -> b fs () p xyz")
279 extrinsics_source = rearrange(extrinsics, "b fs i j -> b fs () () i j")
280 extrinsics_target = rearrange(extrinsics, "b ft i j -> b () ft () i j")
281 intrinsics_target = rearrange(intrinsics, "b ft i j -> b () ft () i j")
282 visibility_source = rearrange(tracks.visibility, "b fs p -> b fs () p")
283 visibility_target = rearrange(tracks.visibility, "b ft p -> b () ft p")
284
285 # Compute flow and visibility.
286 xy_target = reproject_points(
287 xyz_source,
288 extrinsics_target.inverse() @ extrinsics_source,
289 intrinsics_target,
290 )
291 visibility = visibility_source & visibility_target
292
293 # Filter out points that are not in the frame for either the source or target.
294 source_in_frame = (xy_source >= 0).all(dim=-1) & (xy_source < 1).all(dim=-1)
295 target_in_frame = (xy_target >= 0).all(dim=-1) & (xy_target < 1).all(dim=-1)
296 visibility = visibility & source_in_frame & target_in_frame
297
298 return xy_target, visibility

Callers 1

Calls 1

reproject_pointsFunction · 0.85

Tested by

no test coverage detected