Callback function to be executed after the `loss.backward()` call.
(
self,
params: Union[Dict[str, torch.nn.Parameter], torch.nn.ParameterDict],
optimizers: Dict[str, torch.optim.Optimizer],
state: Dict[str, Any],
step: int,
info: Dict[str, Any],
flag: int,
desicnt: int,
maxbounds: float,
minbounds: float,
packed: bool = False,
)
| 89 | info["means2d"].retain_grad() |
| 90 | |
| 91 | def step_post_backward( |
| 92 | self, |
| 93 | params: Union[Dict[str, torch.nn.Parameter], torch.nn.ParameterDict], |
| 94 | optimizers: Dict[str, torch.optim.Optimizer], |
| 95 | state: Dict[str, Any], |
| 96 | step: int, |
| 97 | info: Dict[str, Any], |
| 98 | flag: int, |
| 99 | desicnt: int, |
| 100 | maxbounds: float, |
| 101 | minbounds: float, |
| 102 | packed: bool = False, |
| 103 | ): |
| 104 | """Callback function to be executed after the `loss.backward()` call.""" |
| 105 | if step >= self.refine_stop_iter: |
| 106 | # freeze weights of omega |
| 107 | params["omega"].grad = params["omega"].grad * self.omegamask # TODO check if this is proceed as expected |
| 108 | self.rotationmask = torch.logical_not(self.omegamask) |
| 109 | # freeze weights of rotation |
| 110 | params["quats"].grad = params["quats"].grad * self.rotationmask # TODO check if this is proceed as expected |
| 111 | if step % 1000 == 500 : |
| 112 | zmask = params["means"][:,2] < 4.5 |
| 113 | remove(params=params, optimizers=optimizers, state=state, mask=zmask) |
| 114 | self.omegamask = self._zero_omegabymotion(params, optimizers) # calculate omegamask again to adjust the change of gaussian numbers |
| 115 | torch.cuda.empty_cache() |
| 116 | if step == 10000: |
| 117 | self.removeminmax(params=params, optimizers=optimizers, state=state, maxbounds=maxbounds, minbounds=minbounds) |
| 118 | self.omegamask = self._zero_omegabymotion(params, optimizers) # calculate omegamask again to adjust the change of gaussian numbers |
| 119 | return flag |
| 120 | |
| 121 | self._update_state(params, state, info, packed=packed) |
| 122 | |
| 123 | # TODO: need to consider more strategy, there are totally 3 types of strategy in STG (densify = 1,2,3) |
| 124 | # sicheng: in original STG, n3d scenes in night use densify=1, scenes in day use densify=2 |
| 125 | # here is a implementation of densify=1 |
| 126 | # omega & rotation mask |
| 127 | if step == 8001 : |
| 128 | omegamask = self._zero_omegabymotion(params, optimizers) |
| 129 | self.omegamask = omegamask |
| 130 | # record process |
| 131 | elif step > 8001: |
| 132 | # freeze weights of omega |
| 133 | params["omega"].grad = params["omega"].grad * self.omegamask # this is likely wrong |
| 134 | self.rotationmask = torch.logical_not(self.omegamask) |
| 135 | # freeze weights of rotation |
| 136 | params["quats"].grad = params["quats"].grad * self.rotationmask # this is likely wrong |
| 137 | |
| 138 | if ( |
| 139 | step > self.refine_start_iter |
| 140 | and step % self.refine_every == 0 |
| 141 | # and step % self.reset_every >= self.pause_refine_after_reset |
| 142 | ): |
| 143 | if flag < desicnt: |
| 144 | # grow GSs |
| 145 | n_dupli, n_split = self._grow_gs(params, optimizers, state, step) |
| 146 | if self.verbose: |
| 147 | print( |
| 148 | f"Step {step}: {n_dupli} GSs duplicated, {n_split} GSs split. " |
nothing calls this directly
no test coverage detected