MCPcopy Create free account
hub / github.com/DSL-Lab/StreamSplat / GSDynamicDecoder

Class GSDynamicDecoder

model/model_utils.py:371–448  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

369 return gsparams, prior_params
370
371class 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:

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected