(self, n_e, e_dim)
| 685 | |
| 686 | class Quantizer_module(torch.nn.Module): |
| 687 | def __init__(self, n_e, e_dim): |
| 688 | super(Quantizer_module, self).__init__() |
| 689 | self.embedding = nn.Embedding(n_e, e_dim) |
| 690 | self.embedding.weight.data.uniform_(-1.0 / n_e, 1.0 / n_e) |
| 691 | |
| 692 | def forward(self, x): |
| 693 | d = ( |