MCPcopy Create free account
hub / github.com/OpenImagingLab/4DSloMo / training_setup

Method training_setup

scene/gaussian_model.py:326–352  ·  view source on GitHub ↗
(self, training_args)

Source from the content-addressed store, hash-verified

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 '''

Callers 2

trainingFunction · 0.95
restoreMethod · 0.95

Calls 1

get_expon_lr_funcFunction · 0.90

Tested by

no test coverage detected