@@ -253,26 +253,37 @@ def create_splats_with_optimizers(
253253 quats = torch .rand ((N , 4 )) # [N, 4]
254254 opacities = torch .logit (torch .full ((N ,), init_opacity )) # [N,]
255255
256+ global_lr_factor = 1.
257+
258+ # Adapt means_lr parameter so that a same value has a similar effect for the default strategy and MCMC
259+ means_lr_factor = (1. if isinstance (strategy , MCMCStrategy ) else 0.001 ) * global_lr_factor
260+ logging .info (f"means_lr_factor: { means_lr_factor } " )
261+ scales_lr_factor = global_lr_factor
262+ quats_lr_factor = global_lr_factor
263+ opacities_lr_factor = global_lr_factor
264+ sh0_lr_factor = global_lr_factor
265+ shN_lr_factor = global_lr_factor
266+
256267 params = [
257268 # name, value, lr
258- ("means" , torch .nn .Parameter (points ), means_lr * scene_scale ),
259- ("scales" , torch .nn .Parameter (scales ), scales_lr ),
260- ("quats" , torch .nn .Parameter (quats ), quats_lr ),
261- ("opacities" , torch .nn .Parameter (opacities ), opacities_lr ),
269+ ("means" , torch .nn .Parameter (points ), means_lr_factor * means_lr * scene_scale ),
270+ ("scales" , torch .nn .Parameter (scales ), scales_lr_factor * scales_lr ),
271+ ("quats" , torch .nn .Parameter (quats ), quats_lr_factor * quats_lr ),
272+ ("opacities" , torch .nn .Parameter (opacities ), opacities_lr_factor * opacities_lr ),
262273 ]
263274
264275 if feature_dim is None :
265276 # color is SH coefficients.
266277 colors = torch .zeros ((N , (sh_degree + 1 ) ** 2 , 3 )) # [N, K, 3]
267278 colors [:, 0 , :] = rgb_to_sh (rgbs )
268- params .append (("sh0" , torch .nn .Parameter (colors [:, :1 , :]), sh0_lr ))
269- params .append (("shN" , torch .nn .Parameter (colors [:, 1 :, :]), shN_lr ))
279+ params .append (("sh0" , torch .nn .Parameter (colors [:, :1 , :]), sh0_lr_factor * sh0_lr ))
280+ params .append (("shN" , torch .nn .Parameter (colors [:, 1 :, :]), shN_lr_factor * shN_lr ))
270281 else :
271282 # features will be used for appearance and view-dependent shading
272283 features = torch .rand (N , feature_dim ) # [N, feature_dim]
273- params .append (("features" , torch .nn .Parameter (features ), sh0_lr ))
284+ params .append (("features" , torch .nn .Parameter (features ), sh0_lr_factor * sh0_lr ))
274285 colors = torch .logit (rgbs ) # [N, 3]
275- params .append (("colors" , torch .nn .Parameter (colors ), sh0_lr ))
286+ params .append (("colors" , torch .nn .Parameter (colors ), sh0_lr_factor * sh0_lr ))
276287
277288 splats = torch .nn .ParameterDict ({n : v for n , v , _ in params }).to (device )
278289 # Scale learning rate based on batch size, reference:
0 commit comments