| 73 | return ssim_map.mean(1).mean(1).mean(1) |
| 74 | |
| 75 | def update_params_and_optimizer(new_params, params, optimizer): |
| 76 | for k, v in new_params.items(): |
| 77 | group = [x for x in optimizer.param_groups if x["name"] == k][0] |
| 78 | stored_state = optimizer.state.get(group['params'][0], None) |
| 79 | |
| 80 | stored_state["exp_avg"] = torch.zeros_like(v) |
| 81 | stored_state["exp_avg_sq"] = torch.zeros_like(v) |
| 82 | del optimizer.state[group['params'][0]] |
| 83 | |
| 84 | group["params"][0] = torch.nn.Parameter(v.requires_grad_(True)) |
| 85 | optimizer.state[group['params'][0]] = stored_state |
| 86 | params[k] = group["params"][0] |
| 87 | return params |
| 88 | |
| 89 | def remove_points(to_remove, params, variables, optimizer): |
| 90 | to_keep = ~to_remove |