| 135 | """ |
| 136 | |
| 137 | def __init__( |
| 138 | self, |
| 139 | num_hiddens, |
| 140 | embedding_dim, |
| 141 | n_embed, |
| 142 | straight_through=True, |
| 143 | kl_weight=5e-4, |
| 144 | temp_init=1.0, |
| 145 | use_vqinterface=True, |
| 146 | remap=None, |
| 147 | unknown_index="random", |
| 148 | ): |
| 149 | super().__init__() |
| 150 | |
| 151 | self.embedding_dim = embedding_dim |
| 152 | self.n_embed = n_embed |
| 153 | |
| 154 | self.straight_through = straight_through |
| 155 | self.temperature = temp_init |
| 156 | self.kl_weight = kl_weight |
| 157 | |
| 158 | self.proj = nn.Conv2d(num_hiddens, n_embed, 1) |
| 159 | self.embed = nn.Embedding(n_embed, embedding_dim) |
| 160 | |
| 161 | self.use_vqinterface = use_vqinterface |
| 162 | |
| 163 | self.remap = remap |
| 164 | if self.remap is not None: |
| 165 | self.register_buffer("used", torch.tensor(np.load(self.remap))) |
| 166 | self.re_embed = self.used.shape[0] |
| 167 | self.unknown_index = unknown_index # "random" or "extra" or integer |
| 168 | if self.unknown_index == "extra": |
| 169 | self.unknown_index = self.re_embed |
| 170 | self.re_embed = self.re_embed + 1 |
| 171 | print( |
| 172 | f"Remapping {self.n_embed} indices to {self.re_embed} indices. " |
| 173 | f"Using {self.unknown_index} for unknown indices." |
| 174 | ) |
| 175 | else: |
| 176 | self.re_embed = n_embed |
| 177 | |
| 178 | def remap_to_used(self, inds): |
| 179 | ishape = inds.shape |