| 369 | return gsparams, prior_params |
| 370 | |
| 371 | class GSDynamicDecoder(nn.Module): |
| 372 | def __init__(self, opt: Options, transformer_dim: int, mlp_dim=None, bias=True): |
| 373 | super(GSDynamicDecoder, self).__init__() |
| 374 | self.opt = opt |
| 375 | self.embed_dim = transformer_dim |
| 376 | self.mlp_dim = mlp_dim if mlp_dim is not None else transformer_dim |
| 377 | self.key_dims = {"xyz_dynamic": 3 * (opt.forder), "opacity_dynamic": 2} |
| 378 | self.gs_layer = nn.ModuleDict() |
| 379 | self.prior = Truncated_Gaussian_Model(n_sample=1, nr_mix=1) |
| 380 | self.pm = opt.pm_dynamic |
| 381 | self.register_buffer("dynamic_scalar", torch.tensor([0.5, 0.1, 0.5])) |
| 382 | |
| 383 | for key in ["xyz_dynamic", "opacity_dynamic"]: |
| 384 | if key == "xyz_dynamic": |
| 385 | layer = MLP(self.mlp_dim*2, self.key_dims[key], n_neurons=self.mlp_dim, n_hidden_layers=2, activation="silu", output_activation=None, bias=bias) |
| 386 | if self.pm: |
| 387 | pred_scale = nn.Linear(self.mlp_dim*2, self.key_dims[key], bias=False) |
| 388 | torch.nn.init.xavier_normal_(pred_scale.weight, 0.01) |
| 389 | self.gs_layer[f'{key}_scale'] = pred_scale |
| 390 | elif key == "opacity_dynamic": |
| 391 | layer = MLP(self.mlp_dim*2, self.key_dims[key], n_neurons=self.mlp_dim, n_hidden_layers=2, activation="silu", output_activation=None, bias=bias) |
| 392 | else: |
| 393 | raise NotImplementedError |
| 394 | self.gs_layer[key] = layer |
| 395 | |
| 396 | @autocast('cuda', enabled=False) |
| 397 | def forward(self, feats, timestamp=None): |
| 398 | """ |
| 399 | Perform dynamic predictions. |
| 400 | """ |
| 401 | feats = feats.type(torch.float32) |
| 402 | feats = rearrange(feats, 'b v n d -> (b v) n d') |
| 403 | gsparams = {} |
| 404 | prior_params = {} |
| 405 | for key in ["xyz_dynamic", "opacity_dynamic"]: |
| 406 | v = feats |
| 407 | if f'{key}_scale' in self.gs_layer and key == "xyz_dynamic": |
| 408 | logits_pred = torch.ones(feats.shape[0], feats.shape[1], 1).to(feats.device).float() |
| 409 | means = self.gs_layer[key](v) |
| 410 | log_scales = self.gs_layer[f'{key}_scale'](v) |
| 411 | logits, means, log_scales = self.prior.expand_params(logits_pred, means, log_scales, mean_activation='tanh') |
| 412 | prior_params[key] = {"logits": logits, "means": means, "log_scales": log_scales} # [B, N*r, dim, nr_mix] |
| 413 | val, probs = self.prior.sample(logits, means, log_scales) |
| 414 | val = val.reshape(*val.shape[:2], -1, 3) # [B, N, L * forder, 3] |
| 415 | val = val * self.dynamic_scalar |
| 416 | gsparams[key] = val |
| 417 | elif f'{key}_scale' in self.gs_layer and key == "opacity_dynamic": |
| 418 | logits_pred = torch.ones(feats.shape[0], feats.shape[1], 1).to(feats.device).float() |
| 419 | val = self.gs_layer[key](v) |
| 420 | log_scales = self.gs_layer[f'{key}_scale'](v) |
| 421 | scalar = torch.exp(val[..., 0:1]) |
| 422 | means = val[..., 1:2] |
| 423 | logits, means, log_scales = self.prior.expand_params(logits_pred, means, log_scales, mean_activation='tanh') |
| 424 | prior_params[key] = {"logits": logits, "means": means, "log_scales": log_scales} # [B, N*r, dim, nr_mix] |
| 425 | val, probs = self.prior.sample(logits, means, log_scales) |
| 426 | val = 0.5 + 0.5 * val # t1 in [0, 1] |
| 427 | gsparams[key] = torch.cat([scalar, val], dim=-1) |
| 428 | else: |