MCPcopy Create free account
hub / github.com/IRMVLab/SemGauss-SLAM / report_progress

Function report_progress

utils/eval_utils.py:79–170  ·  view source on GitHub ↗
(params, data, i, progress_bar, iter_time_idx, sil_thres, every_i=1, qual_every_i=1, 
                    tracking=False, mapping=False, online_time_idx=None)

Source from the content-addressed store, hash-verified

77 return avg_trans_error
78
79def report_progress(params, data, i, progress_bar, iter_time_idx, sil_thres, every_i=1, qual_every_i=1,
80 tracking=False, mapping=False, online_time_idx=None):
81 if i % every_i == 0 or i == 1:
82 if tracking:
83 # Get list of gt poses
84 gt_w2c_list = data['iter_gt_w2c_list']
85 valid_gt_w2c_list = []
86
87 # Get latest trajectory
88 latest_est_w2c = data['w2c']
89 latest_est_w2c_list = []
90 latest_est_w2c_list.append(latest_est_w2c)
91 valid_gt_w2c_list.append(gt_w2c_list[0])
92 for idx in range(1, iter_time_idx+1):
93 # Check if gt pose is not nan for this time step
94 if torch.isnan(gt_w2c_list[idx]).sum() > 0:
95 continue
96 interm_cam_rot = F.normalize(params['cam_unnorm_rots'][..., idx].detach())
97 interm_cam_trans = params['cam_trans'][..., idx].detach()
98 intermrel_w2c = torch.eye(4).cuda().float()
99 intermrel_w2c[:3, :3] = build_rotation(interm_cam_rot)
100 intermrel_w2c[:3, 3] = interm_cam_trans
101 latest_est_w2c = intermrel_w2c
102 latest_est_w2c_list.append(latest_est_w2c)
103 valid_gt_w2c_list.append(gt_w2c_list[idx])
104
105 # Get latest gt pose
106 gt_w2c_list = valid_gt_w2c_list
107 iter_gt_w2c = gt_w2c_list[-1]
108 # Get euclidean distance error between latest and gt pose
109 iter_pt_error = torch.sqrt((latest_est_w2c[0,3] - iter_gt_w2c[0,3])**2 + (latest_est_w2c[1,3] - iter_gt_w2c[1,3])**2 + (latest_est_w2c[2,3] - iter_gt_w2c[2,3])**2)
110 if iter_time_idx > 0:
111 # Calculate relative pose error
112 rel_gt_w2c = relative_transformation(gt_w2c_list[-2], gt_w2c_list[-1])
113 rel_est_w2c = relative_transformation(latest_est_w2c_list[-2], latest_est_w2c_list[-1])
114 rel_pt_error = torch.sqrt((rel_gt_w2c[0,3] - rel_est_w2c[0,3])**2 + (rel_gt_w2c[1,3] - rel_est_w2c[1,3])**2 + (rel_gt_w2c[2,3] - rel_est_w2c[2,3])**2)
115 else:
116 rel_pt_error = torch.zeros(1).float()
117
118 # Calculate ATE RMSE
119 ate_rmse = evaluate_ate(gt_w2c_list, latest_est_w2c_list)
120 ate_rmse = np.round(ate_rmse, decimals=6)
121
122 # Get current frame Gaussians
123 transformed_pts = transform_to_frame(params, iter_time_idx,
124 gaussians_grad=False,
125 camera_grad=False)
126
127 # Initialize Render Variables
128 # semantic
129 rendervar = transformed_params2rendervar(params, data['w2c'], transformed_pts)
130 im_dep, radius, _, semantics = Renderer(raster_settings=data['cam'])(**rendervar)
131 im = im_dep[0:3, :, :]
132 depth_sil = im_dep[3:, :, :]
133 rastered_depth = depth_sil[0, :, :].unsqueeze(0)
134 valid_depth_mask = (data['depth'] > 0)
135 silhouette = depth_sil[1, :, :]
136 presence_sil_mask = (silhouette > sil_thres)

Callers 1

dense_semantic_slamFunction · 0.90

Calls 8

build_rotationFunction · 0.90
relative_transformationFunction · 0.90
transform_to_frameFunction · 0.90
calc_psnrFunction · 0.90
evaluate_ateFunction · 0.85
printFunction · 0.85
updateMethod · 0.45

Tested by

no test coverage detected