| 326 | self.active_sh_degree = self.max_sh_degree |
| 327 | |
| 328 | def replace_tensor_to_optimizer(self, tensor, name): |
| 329 | optimizable_tensors = {} |
| 330 | for group in self.optimizer.param_groups: |
| 331 | if group["name"] == name: |
| 332 | # breakpoint() |
| 333 | stored_state = self.optimizer.state.get(group['params'][0], None) |
| 334 | stored_state["exp_avg"] = torch.zeros_like(tensor) |
| 335 | stored_state["exp_avg_sq"] = torch.zeros_like(tensor) |
| 336 | |
| 337 | del self.optimizer.state[group['params'][0]] |
| 338 | group["params"][0] = nn.Parameter(tensor.requires_grad_(True)) |
| 339 | self.optimizer.state[group['params'][0]] = stored_state |
| 340 | |
| 341 | optimizable_tensors[group["name"]] = group["params"][0] |
| 342 | return optimizable_tensors |
| 343 | |
| 344 | def _prune_optimizer(self, mask): |
| 345 | optimizable_tensors = {} |