(self, in_channels=3, num_pos_feats=768, n_freqs=8, logscale=True)
| 44 | """ |
| 45 | |
| 46 | def __init__(self, in_channels=3, num_pos_feats=768, n_freqs=8, logscale=True): |
| 47 | super(Ego3DPositionEmbeddingMLP, self).__init__() |
| 48 | self.n_freqs = n_freqs |
| 49 | self.freq_out_channels = in_channels * (2 * n_freqs + 1) |
| 50 | if logscale: |
| 51 | freq_bands = 2 ** torch.linspace(0, n_freqs - 1, n_freqs) |
| 52 | else: |
| 53 | freq_bands = torch.linspace(1, 2 ** (n_freqs - 1), n_freqs) |
| 54 | |
| 55 | center = torch.tensor([0., 0., 2.]).repeat(in_channels // 3) |
| 56 | self.register_buffer("freq_bands", freq_bands, persistent=False) |
| 57 | self.register_buffer("center", center, persistent=False) |
| 58 | |
| 59 | self.position_embedding_head = nn.Sequential( |
| 60 | nn.Linear(self.freq_out_channels, num_pos_feats), |
| 61 | nn.LayerNorm(num_pos_feats), |
| 62 | nn.ReLU(), |
| 63 | nn.Linear(num_pos_feats, num_pos_feats), |
| 64 | ) |
| 65 | self._reset_parameters() |
| 66 | |
| 67 | def _reset_parameters(self): |
| 68 | """init with small weights to maintain stable training.""" |
no test coverage detected