MCPcopy Create free account
hub / github.com/JasonLSC/GSCodec_Studio / step_post_backward

Method step_post_backward

gsplat/strategy/STG_Strategy.py:91–180  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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. "

Callers

nothing calls this directly

Calls 7

_zero_omegabymotionMethod · 0.95
removeminmaxMethod · 0.95
_update_stateMethod · 0.95
_grow_gsMethod · 0.95
_prune_gsMethod · 0.95
removeFunction · 0.85
reset_opaFunction · 0.85

Tested by

no test coverage detected