| 87 | return params |
| 88 | |
| 89 | def remove_points(to_remove, params, variables, optimizer): |
| 90 | to_keep = ~to_remove |
| 91 | keys = [k for k in params.keys() if k not in ['cam_unnorm_rots', 'cam_trans']] |
| 92 | for k in keys: |
| 93 | group = [g for g in optimizer.param_groups if g['name'] == k][0] |
| 94 | stored_state = optimizer.state.get(group['params'][0], None) |
| 95 | if stored_state is not None: |
| 96 | stored_state["exp_avg"] = stored_state["exp_avg"][to_keep] |
| 97 | stored_state["exp_avg_sq"] = stored_state["exp_avg_sq"][to_keep] |
| 98 | del optimizer.state[group['params'][0]] |
| 99 | group["params"][0] = torch.nn.Parameter((group["params"][0][to_keep].requires_grad_(True))) |
| 100 | optimizer.state[group['params'][0]] = stored_state |
| 101 | params[k] = group["params"][0] |
| 102 | else: |
| 103 | group["params"][0] = torch.nn.Parameter(group["params"][0][to_keep].requires_grad_(True)) |
| 104 | params[k] = group["params"][0] |
| 105 | variables['means2D_gradient_accum'] = variables['means2D_gradient_accum'][to_keep] |
| 106 | variables['denom'] = variables['denom'][to_keep] |
| 107 | variables['max_2D_radius'] = variables['max_2D_radius'][to_keep] |
| 108 | if 'timestep' in variables.keys(): |
| 109 | variables['timestep'] = variables['timestep'][to_keep] |
| 110 | return params, variables |
| 111 | |
| 112 | |
| 113 | def inverse_sigmoid(x): |