(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)
| 77 | return avg_trans_error |
| 78 | |
| 79 | def 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) |
no test coverage detected