(source_path, model_path, ckpt_path, device, batch_size, image_size, schedule, lr, niter,
min_conf_thr, llffhold, n_views, co_vis_dsp, depth_thre, conf_aware_ranking=False, focal_avg=False, infer_video=False)
| 22 | |
| 23 | |
| 24 | def main(source_path, model_path, ckpt_path, device, batch_size, image_size, schedule, lr, niter, |
| 25 | min_conf_thr, llffhold, n_views, co_vis_dsp, depth_thre, conf_aware_ranking=False, focal_avg=False, infer_video=False): |
| 26 | |
| 27 | # ---------------- (1) Load model and images ---------------- |
| 28 | save_path, sparse_0_path, sparse_1_path = init_filestructure(Path(source_path), n_views) |
| 29 | model = AsymmetricMASt3R.from_pretrained(ckpt_path).to(device) |
| 30 | image_dir = Path(source_path) / 'images' |
| 31 | image_files, image_suffix = get_sorted_image_files(image_dir) |
| 32 | if infer_video: |
| 33 | train_img_files = image_files |
| 34 | else: |
| 35 | train_img_files, test_img_files = split_train_test(image_files, llffhold, n_views, verbose=True) |
| 36 | |
| 37 | # when geometry init, only use train images |
| 38 | image_files = train_img_files |
| 39 | images, org_imgs_shape = load_images(image_files, size=image_size) |
| 40 | |
| 41 | start_time = time() |
| 42 | print(f'>> Making pairs...') |
| 43 | pairs = make_pairs(images, scene_graph='complete', prefilter=None, symmetrize=True) |
| 44 | print(f'>> Inference...') |
| 45 | output = inference(pairs, model, device, batch_size=1, verbose=True) |
| 46 | print(f'>> Global alignment...') |
| 47 | scene = global_aligner(output, device=args.device, mode=GlobalAlignerMode.PointCloudOptimizer) |
| 48 | loss = scene.compute_global_alignment(init="mst", niter=300, schedule=schedule, lr=lr, focal_avg=args.focal_avg) |
| 49 | |
| 50 | # Extract scene information |
| 51 | extrinsics_w2c = inv(to_numpy(scene.get_im_poses())) |
| 52 | intrinsics = to_numpy(scene.get_intrinsics()) |
| 53 | focals = to_numpy(scene.get_focals()) |
| 54 | imgs = np.array(scene.imgs) |
| 55 | pts3d = to_numpy(scene.get_pts3d()) |
| 56 | pts3d = np.array(pts3d) |
| 57 | depthmaps = to_numpy(scene.im_depthmaps.detach().cpu().numpy()) |
| 58 | values = [param.detach().cpu().numpy() for param in scene.im_conf] |
| 59 | confs = np.array(values) |
| 60 | |
| 61 | if conf_aware_ranking: |
| 62 | print(f'>> Confiden-aware Ranking...') |
| 63 | avg_conf_scores = confs.mean(axis=(1, 2)) |
| 64 | sorted_conf_indices = np.argsort(avg_conf_scores)[::-1] |
| 65 | sorted_conf_avg_conf_scores = avg_conf_scores[sorted_conf_indices] |
| 66 | print("Sorted indices:", sorted_conf_indices) |
| 67 | print("Sorted average confidence scores:", sorted_conf_avg_conf_scores) |
| 68 | else: |
| 69 | sorted_conf_indices = np.arange(n_views) |
| 70 | print("Sorted indices:", sorted_conf_indices) |
| 71 | |
| 72 | # Calculate the co-visibility mask |
| 73 | print(f'>> Calculate the co-visibility mask...') |
| 74 | if depth_thre > 0: |
| 75 | overlapping_masks = compute_co_vis_masks(sorted_conf_indices, depthmaps, pts3d, intrinsics, extrinsics_w2c, imgs.shape, depth_threshold=depth_thre) |
| 76 | overlapping_masks = ~overlapping_masks |
| 77 | else: |
| 78 | co_vis_dsp = False |
| 79 | overlapping_masks = None |
| 80 | end_time = time() |
| 81 | Train_Time = end_time - start_time |
no test coverage detected