| 218 | self.sampled_val = {} |
| 219 | |
| 220 | def add_layer(self, layer_dict, key, dim, init_val=0., bias=True): |
| 221 | if self.nr_mix > 1: |
| 222 | pred_prob = nn.Linear(self.mlp_dim, self.nr_mix * self.ratio, bias=bias) |
| 223 | torch.nn.init.xavier_uniform_(pred_prob.weight) |
| 224 | layer_dict[f"{key}_prob"] = pred_prob |
| 225 | |
| 226 | pred_mean = nn.Linear(self.mlp_dim, self.nr_mix * dim * self.ratio, bias=bias) |
| 227 | torch.nn.init.xavier_normal_(pred_mean.weight, 0.01) |
| 228 | layer_dict[f"{key}_mean"] = pred_mean |
| 229 | |
| 230 | pred_scale = nn.Linear(self.mlp_dim, self.nr_mix * dim * self.ratio, bias=bias) |
| 231 | torch.nn.init.xavier_normal_(pred_scale.weight, 0.01) |
| 232 | layer_dict[f"{key}_scale"] = pred_scale |
| 233 | |
| 234 | def key_activation(self, v: torch.Tensor, key=''): |
| 235 | # [B, N, D] |