MCPcopy Create free account
hub / github.com/NVlabs/InstantSplat / main

Function main

init_geo.py:24–129  ·  view source on GitHub ↗
(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)

Source from the content-addressed store, hash-verified

22
23
24def 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

Callers 1

init_geo.pyFile · 0.70

Calls 15

init_filestructureFunction · 0.90
get_sorted_image_filesFunction · 0.90
split_train_testFunction · 0.90
load_imagesFunction · 0.90
make_pairsFunction · 0.90
inferenceFunction · 0.90
global_alignerFunction · 0.90
invFunction · 0.90
to_numpyFunction · 0.90
compute_co_vis_masksFunction · 0.90
save_timeFunction · 0.90

Tested by

no test coverage detected