(
self, n_e, vq_embed_dim, beta, remap=None, unknown_index="random", sane_index_shape=False, legacy=True
)
| 290 | # backwards compatibility we use the buggy version by default, but you can |
| 291 | # specify legacy=False to fix it. |
| 292 | def __init__( |
| 293 | self, n_e, vq_embed_dim, beta, remap=None, unknown_index="random", sane_index_shape=False, legacy=True |
| 294 | ): |
| 295 | super().__init__() |
| 296 | self.n_e = n_e |
| 297 | self.vq_embed_dim = vq_embed_dim |
| 298 | self.beta = beta |
| 299 | self.legacy = legacy |
| 300 | |
| 301 | self.embedding = nn.Embedding(self.n_e, self.vq_embed_dim) |
| 302 | self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e) |
| 303 | |
| 304 | self.remap = remap |
| 305 | if self.remap is not None: |
| 306 | self.register_buffer("used", torch.tensor(np.load(self.remap))) |
| 307 | self.re_embed = self.used.shape[0] |
| 308 | self.unknown_index = unknown_index # "random" or "extra" or integer |
| 309 | if self.unknown_index == "extra": |
| 310 | self.unknown_index = self.re_embed |
| 311 | self.re_embed = self.re_embed + 1 |
| 312 | print( |
| 313 | f"Remapping {self.n_e} indices to {self.re_embed} indices. " |
| 314 | f"Using {self.unknown_index} for unknown indices." |
| 315 | ) |
| 316 | else: |
| 317 | self.re_embed = n_e |
| 318 | |
| 319 | self.sane_index_shape = sane_index_shape |
| 320 | |
| 321 | def remap_to_used(self, inds): |
| 322 | ishape = inds.shape |
nothing calls this directly
no test coverage detected