(self, training_args)
| 324 | self._rotation_r = nn.Parameter(rots_r.requires_grad_(True)) |
| 325 | |
| 326 | def training_setup(self, training_args): |
| 327 | self.percent_dense = training_args.percent_dense |
| 328 | self.xyz_gradient_accum = torch.zeros((self.get_xyz.shape[0], 1), device="cuda") |
| 329 | self.denom = torch.zeros((self.get_xyz.shape[0], 1), device="cuda") |
| 330 | |
| 331 | l = [ |
| 332 | {'params': [self._xyz], 'lr': training_args.position_lr_init * self.spatial_lr_scale, "name": "xyz"}, |
| 333 | {'params': [self._features_dc], 'lr': training_args.feature_lr, "name": "f_dc"}, |
| 334 | {'params': [self._features_rest], 'lr': training_args.feature_lr / 20.0, "name": "f_rest"}, |
| 335 | {'params': [self._opacity], 'lr': training_args.opacity_lr, "name": "opacity"}, |
| 336 | {'params': [self._scaling], 'lr': training_args.scaling_lr, "name": "scaling"}, |
| 337 | {'params': [self._rotation], 'lr': training_args.rotation_lr, "name": "rotation"} |
| 338 | ] |
| 339 | if self.gaussian_dim == 4: # TODO: tune time_lr_scale |
| 340 | if training_args.position_t_lr_init < 0: |
| 341 | training_args.position_t_lr_init = training_args.position_lr_init |
| 342 | self.t_gradient_accum = torch.zeros((self.get_xyz.shape[0], 1), device="cuda") |
| 343 | l.append({'params': [self._t], 'lr': training_args.position_t_lr_init * self.spatial_lr_scale, "name": "t"}) |
| 344 | l.append({'params': [self._scaling_t], 'lr': training_args.scaling_lr, "name": "scaling_t"}) |
| 345 | if self.rot_4d: |
| 346 | l.append({'params': [self._rotation_r], 'lr': training_args.rotation_lr, "name": "rotation_r"}) |
| 347 | |
| 348 | self.optimizer = torch.optim.Adam(l, lr=0.0, eps=1e-15) |
| 349 | self.xyz_scheduler_args = get_expon_lr_func(lr_init=training_args.position_lr_init*self.spatial_lr_scale, |
| 350 | lr_final=training_args.position_lr_final*self.spatial_lr_scale, |
| 351 | lr_delay_mult=training_args.position_lr_delay_mult, |
| 352 | max_steps=training_args.position_lr_max_steps) |
| 353 | |
| 354 | def update_learning_rate(self, iteration): |
| 355 | ''' Learning rate scheduling per step ''' |
no test coverage detected