diff --git a/demo_colmap.py b/demo_colmap.py index 836af172..620be6a1 100644 --- a/demo_colmap.py +++ b/demo_colmap.py @@ -278,11 +278,14 @@ def rename_colmap_recons_and_rescale_camera( pycamera.height = real_image_size[1] if shift_point2d_to_original_res: - # Also shift the point2D to original resolution - top_left = original_coords[pyimageid - 1, :2] + x1, y1, x2, y2, w, h = original_coords[pyimageid - 1] + if x1 == 0: + shift = np.array([x1, y1]) * (w / x2) + else: + shift = np.array([x1, y1]) * (h / y2) for point2D in pyimage.points2D: - point2D.xy = (point2D.xy - top_left) * resize_ratio + point2D.xy = point2D.xy * resize_ratio - shift if shared_camera: # If shared_camera, all images share the same camera diff --git a/vggt/dependency/np_to_pycolmap.py b/vggt/dependency/np_to_pycolmap.py index 61ea5786..ce82354b 100644 --- a/vggt/dependency/np_to_pycolmap.py +++ b/vggt/dependency/np_to_pycolmap.py @@ -55,8 +55,8 @@ def batch_np_matrix_to_pycolmap( if max_reproj_error is not None: projected_points_2d, projected_points_cam = project_3D_points_np(points3d, extrinsics, intrinsics) - projected_diff = np.linalg.norm(projected_points_2d - tracks, axis=-1) projected_points_2d[projected_points_cam[:, -1] <= 0] = 1e6 + projected_diff = np.linalg.norm(projected_points_2d - tracks, axis=-1) reproj_mask = projected_diff < max_reproj_error if masks is not None and reproj_mask is not None: diff --git a/vggt/dependency/track_predict.py b/vggt/dependency/track_predict.py index c15c23fe..2dc276ec 100644 --- a/vggt/dependency/track_predict.py +++ b/vggt/dependency/track_predict.py @@ -4,13 +4,16 @@ # This source code is licensed under the license found in the # LICENSE file in the root directory of this source tree. -import torch +import math import numpy as np from .vggsfm_utils import * +from .projection import project_3D_points_np def predict_tracks( images, + extrinsics, + intrinsics, conf=None, points_3d=None, masks=None, @@ -20,6 +23,9 @@ def predict_tracks( max_points_num=163840, fine_tracking=True, complete_non_vis=True, + track_vis_thresh=0.1, + reproj_error_thresh=8, + min_inlier_per_frame=64 ): """ Predict tracks for the given images and masks. @@ -33,6 +39,8 @@ def predict_tracks( Args: images: Tensor of shape [S, 3, H, W] containing the input images. + extrinsics: Array of shape [S, 3, 4] containing the extrinsic parameters for each frame. + intrinsics: Array of shape [S, 3, 3] containing the intrinsic parameters for each frame. conf: Tensor of shape [S, 1, H, W] containing the confidence scores. Default is None. points_3d: Tensor containing 3D points. Default is None. masks: Optional tensor of shape [S, 1, H, W] containing masks. Default is None. @@ -42,6 +50,9 @@ def predict_tracks( max_points_num: Maximum number of points to process at once. Default is 163840. fine_tracking: Whether to use fine tracking. Default is True. complete_non_vis: Whether to augment non-visible frames. Default is True. + track_vis_thresh: Visibility threshold for track filtering + reproj_error_thresh: Reprojection error threshold for track filtering + min_inlier_per_frame: Minimum number of inliers per frame Returns: pred_tracks: Numpy array containing the predicted tracks. @@ -108,6 +119,8 @@ def predict_tracks( pred_points_3d, pred_colors, images, + extrinsics, + intrinsics, conf, points_3d, fmaps_for_tracker, @@ -115,9 +128,10 @@ def predict_tracks( tracker, max_points_num, fine_tracking, - min_vis=500, - non_vis_thresh=0.1, - device=device, + track_vis_thresh=track_vis_thresh, + reproj_error_thresh=reproj_error_thresh, + min_inlier_per_frame=min_inlier_per_frame, + device=device ) pred_tracks = np.concatenate(pred_tracks, axis=1) @@ -209,16 +223,17 @@ def _forward_on_query( fmaps_feed = fmaps_feed[None] # add batch dimension all_points_num = images_feed.shape[1] * query_points.shape[1] + gpu_count = torch.cuda.device_count() + gpu_points_num = math.ceil(all_points_num / gpu_count) - # Don't need to be scared, this is just chunking to make GPU happy - if all_points_num > max_points_num: - num_splits = (all_points_num + max_points_num - 1) // max_points_num - query_points = torch.chunk(query_points, num_splits, dim=1) - else: + if gpu_points_num <= max_points_num: query_points = [query_points] + else: + chunk_num = math.ceil(all_points_num / (max_points_num * gpu_count)) + query_points = torch.chunk(query_points, chunk_num, dim=1) pred_track, pred_vis, _ = predict_tracks_in_chunks( - tracker, images_feed, query_points, fmaps_feed, fine_tracking=fine_tracking + tracker, images_feed, query_points, fmaps_feed, fine_tracking=fine_tracking, parallel=gpu_count > 1 ) pred_track, pred_vis = switch_tensor_order([pred_track, pred_vis], reorder_index, dim=1) @@ -236,6 +251,8 @@ def _augment_non_visible_frames( pred_points_3d: list, # ← running list of np.ndarrays for 3D points pred_colors: list, # ← running list of np.ndarrays for colors images: torch.Tensor, + extrinsics, + intrinsics, conf, points_3d, fmaps_for_tracker, @@ -243,9 +260,9 @@ def _augment_non_visible_frames( tracker, max_points_num: int, fine_tracking: bool, - *, - min_vis: int = 500, - non_vis_thresh: float = 0.1, + track_vis_thresh: float = 0.1, + reproj_error_thresh: float = 8.0, + min_inlier_per_frame: int = 64, device: torch.device = None, ): """ @@ -258,6 +275,8 @@ def _augment_non_visible_frames( pred_points_3d: List of numpy arrays containing 3D points. pred_colors: List of numpy arrays containing point colors. images: Tensor of shape [S, 3, H, W] containing the input images. + extrinsics: Array of shape [S, 3, 4] containing the extrinsic parameters for each frame. + intrinsics: Array of shape [S, 3, 3] containing the intrinsic parameters for each frame. conf: Tensor of shape [S, 1, H, W] containing confidence scores points_3d: Tensor containing 3D points fmaps_for_tracker: Feature maps for the tracker @@ -265,8 +284,9 @@ def _augment_non_visible_frames( tracker: VGG-SFM tracker max_points_num: Maximum number of points to process at once fine_tracking: Whether to use fine tracking - min_vis: Minimum visibility threshold - non_vis_thresh: Non-visibility threshold + track_vis_thresh: Visibility threshold for tracks filtering + reproj_error_thresh: Reprojection error threshold for track filtering + min_inlier_per_frame: Minimum number of inliers per frame device: Device to use for computation Returns: @@ -277,28 +297,31 @@ def _augment_non_visible_frames( cur_extractors = keypoint_extractors # may be replaced on the final trial while True: - # Visibility per frame - vis_array = np.concatenate(pred_vis_scores, axis=1) - # Count frames with sufficient visibility using numpy - sufficient_vis_count = (vis_array > non_vis_thresh).sum(axis=-1) - non_vis_frames = np.where(sufficient_vis_count < min_vis)[0].tolist() + vis_mask = np.concatenate(pred_vis_scores, axis=-1) > track_vis_thresh + + projected_points_2d, projected_points_cam = project_3D_points_np(np.concatenate(pred_points_3d, axis=0), extrinsics, intrinsics) + projected_points_2d[projected_points_cam[:, -1] <= 0] = 1e6 + projected_diff = np.linalg.norm(projected_points_2d - np.concatenate(pred_tracks, axis=1), axis=-1) + reproj_mask = projected_diff < reproj_error_thresh + mask = np.logical_and(vis_mask, reproj_mask) + non_inlier_frames = np.where(mask.sum(axis=1) < min_inlier_per_frame)[0].tolist() - if len(non_vis_frames) == 0: + if len(non_inlier_frames) == 0: break - print("Processing non visible frames:", non_vis_frames) + print("Processing non enough inlier frames:", non_inlier_frames) # Decide the frames & extractor for this round - if non_vis_frames[0] == last_query: + if non_inlier_frames[0] == last_query: # Same frame failed twice - final "all-in" attempt final_trial = True cur_extractors = initialize_feature_extractors(2048, extractor_method="sp+sift+aliked", device=device) - query_frame_list = non_vis_frames # blast them all at once + query_frame_list = non_inlier_frames # blast them all at once else: - query_frame_list = [non_vis_frames[0]] # Process one at a time + query_frame_list = [non_inlier_frames[0]] # Process one at a time - last_query = non_vis_frames[0] + last_query = non_inlier_frames[0] # Run the tracker for every selected frame for query_index in query_frame_list: diff --git a/vggt/dependency/vggsfm_utils.py b/vggt/dependency/vggsfm_utils.py index 3f7d9ba6..72e94393 100644 --- a/vggt/dependency/vggsfm_utils.py +++ b/vggt/dependency/vggsfm_utils.py @@ -5,14 +5,13 @@ # LICENSE file in the root directory of this source tree. import logging +import math import warnings -from typing import Dict, List, Optional, Tuple, Union -import numpy as np -import pycolmap import torch import torch.nn.functional as F from lightglue import ALIKED, SIFT, SuperPoint +from torch.nn import DataParallel from .vggsfm_tracker import TrackerPredictor @@ -95,17 +94,19 @@ def generate_rank_by_dino( frame_feat_norm = F.normalize(frame_feat, p=2, dim=1) similarity_matrix = torch.mm(frame_feat_norm, frame_feat_norm.transpose(-1, -2)) - distance_matrix = 100 - similarity_matrix.clone() + with torch.amp.autocast("cuda", enabled=False): + similarity_matrix = similarity_matrix.float() + distance_matrix = 100 - similarity_matrix - # Ignore self-pairing - similarity_matrix.fill_diagonal_(-100) - similarity_sum = similarity_matrix.sum(dim=1) + # Ignore self-pairing + similarity_matrix.fill_diagonal_(-100) + similarity_sum = similarity_matrix.sum(dim=1) - # Find the most common frame - most_common_frame_index = torch.argmax(similarity_sum).item() + # Find the most common frame + most_common_frame_index = torch.argmax(similarity_sum).item() - # Conduct FPS sampling starting from the most common frame - fps_idx = farthest_point_sampling(distance_matrix, query_frame_num, most_common_frame_index) + # Conduct FPS sampling starting from the most common frame + fps_idx = farthest_point_sampling(distance_matrix, query_frame_num, most_common_frame_index) # Clean up all tensors and models to free memory del frame_feat, frame_feat_norm, similarity_matrix, distance_matrix @@ -253,7 +254,7 @@ def extract_keypoints(query_image, extractors, round_keypoints=True): def predict_tracks_in_chunks( - track_predictor, images_feed, query_points_list, fmaps_feed, fine_tracking, num_splits=None, fine_chunk=40960 + track_predictor, images_feed, query_points_list, fmaps_feed, fine_tracking, num_splits=None, fine_chunk=40960, parallel=True, ): """ Process a list of query points to avoid memory issues. @@ -265,6 +266,7 @@ def predict_tracks_in_chunks( fmaps_feed (torch.Tensor): A tensor of feature maps for the tracker. fine_tracking (bool): Whether to perform fine tracking. num_splits (int, optional): Ignored when query_points_list is provided. Kept for backward compatibility. + parallel: Using multiple GPUs to speed up predicting tracks. Default is True. Returns: tuple: A tuple containing the concatenated predicted tracks, visibility, and scores. @@ -284,11 +286,30 @@ def predict_tracks_in_chunks( pred_vis_list = [] pred_score_list = [] + track_predictor = DataParallel(track_predictor) if parallel else track_predictor + for split_points in query_points_list: # Feed into track predictor for each split + if parallel: + gpu_count = torch.cuda.device_count() + points_num = split_points.shape[1] + batch_points_num = math.ceil(points_num / gpu_count) + padding_num = batch_points_num * gpu_count - points_num + mask = torch.arange(batch_points_num * gpu_count) < points_num + split_points = F.pad(split_points, pad=(0, 0, 0, padding_num)) + split_points = split_points.reshape(gpu_count, batch_points_num, -1) + images_feed = images_feed.expand(gpu_count, *images_feed.shape[1:]) + fmaps_feed = fmaps_feed.expand(gpu_count, *fmaps_feed.shape[1:]) + fine_pred_track, _, pred_vis, pred_score = track_predictor( images_feed, split_points, fmaps=fmaps_feed, fine_tracking=fine_tracking, fine_chunk=fine_chunk ) + + if parallel: + B, S, N, D = fine_pred_track.shape + fine_pred_track = fine_pred_track.permute(1, 0, 2, 3).reshape(1, S, -1, D)[:, :, mask, :] + pred_vis = pred_vis.permute(1, 0, 2).reshape(1, S, -1)[:, :, mask] + fine_pred_track_list.append(fine_pred_track) pred_vis_list.append(pred_vis) pred_score_list.append(pred_score)