| 654 | # backwards compatibility we use the buggy version by default, but you can |
| 655 | # specify legacy=False to fix it. |
| 656 | def __init__( |
| 657 | self, |
| 658 | n_e: int, |
| 659 | vq_embed_dim: int, |
| 660 | beta: float, |
| 661 | remap=None, |
| 662 | unknown_index: str = "random", |
| 663 | sane_index_shape: bool = False, |
| 664 | legacy: bool = True, |
| 665 | ): |
| 666 | super().__init__() |
| 667 | self.n_e = n_e |
| 668 | self.vq_embed_dim = vq_embed_dim |
| 669 | self.beta = beta |
| 670 | self.legacy = legacy |
| 671 | |
| 672 | self.embedding = nn.Embedding(self.n_e, self.vq_embed_dim) |
| 673 | self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e) |
| 674 | |
| 675 | self.remap = remap |
| 676 | if self.remap is not None: |
| 677 | self.register_buffer("used", torch.tensor(np.load(self.remap))) |
| 678 | self.used: torch.Tensor |
| 679 | self.re_embed = self.used.shape[0] |
| 680 | self.unknown_index = unknown_index # "random" or "extra" or integer |
| 681 | if self.unknown_index == "extra": |
| 682 | self.unknown_index = self.re_embed |
| 683 | self.re_embed = self.re_embed + 1 |
| 684 | print( |
| 685 | f"Remapping {self.n_e} indices to {self.re_embed} indices. " |
| 686 | f"Using {self.unknown_index} for unknown indices." |
| 687 | ) |
| 688 | else: |
| 689 | self.re_embed = n_e |
| 690 | |
| 691 | self.sane_index_shape = sane_index_shape |
| 692 | |
| 693 | def remap_to_used(self, inds: torch.LongTensor) -> torch.LongTensor: |
| 694 | ishape = inds.shape |