Skip to content

Commit 56e21bf

Browse files
Adapt internally the parameters so that the exposed parameters have similar effects for the default strategy and MCMC
1 parent 5923849 commit 56e21bf

1 file changed

Lines changed: 19 additions & 8 deletions

File tree

gsplatInterface/trainer.py

Lines changed: 19 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)