| 14 | |
| 15 | |
| 16 | class VectorQuantize(nn.Module): |
| 17 | def __init__( |
| 18 | self, |
| 19 | input_dim: int, |
| 20 | codebook_size: int, |
| 21 | codebook_dim: int, |
| 22 | ): |
| 23 | super().__init__() |
| 24 | self.input_dim = input_dim |
| 25 | self.codebook_size = codebook_size |
| 26 | self.codebook_dim = codebook_dim |
| 27 | |
| 28 | self.in_project = ( |
| 29 | WNConv1d( |
| 30 | self.input_dim, self.codebook_dim, kernel_size=1 |
| 31 | ) # (B, D, T) -> (B, D', T) |
| 32 | if self.input_dim != self.codebook_dim |
| 33 | else nn.Identity() |
| 34 | ) |
| 35 | self.out_project = ( |
| 36 | WNConv1d( |
| 37 | self.codebook_dim, self.input_dim, kernel_size=1 |
| 38 | ) # (B, D', T) -> (B, D, T) |
| 39 | if self.input_dim != self.codebook_dim |
| 40 | else nn.Identity() |
| 41 | ) |
| 42 | |
| 43 | # Initialize codebook and EMA buffers |
| 44 | self.register_buffer( |
| 45 | "codebook", torch.zeros(codebook_size, codebook_dim).float() |
| 46 | ) # (codebook_size, D'), ensure fp32 |
| 47 | # Place holder, not used in inference |
| 48 | self.register_buffer("inited", torch.tensor([True], dtype=torch.bool)) # (1) |
| 49 | self.register_buffer( |
| 50 | "cluster_size", torch.zeros(codebook_size).float() |
| 51 | ) # (codebook_size), ensure fp32 |
| 52 | self.register_buffer( |
| 53 | "embed_avg", self.codebook.clone().float() |
| 54 | ) # (codebook_size, D'), ensure fp32 |
| 55 | |
| 56 | def decode_code(self, embed_id): # embed_id: (B, T) |
| 57 | embed = ( |
| 58 | F.embedding(embed_id, self.codebook).transpose(1, 2).float() |
| 59 | ) # (B, D', T), ensure fp32 |
| 60 | return embed |
| 61 | |
| 62 | def encode_code(self, z: torch.Tensor): # z: (B, D, T) |
| 63 | # logging.info(f"{self.cluster_size = }, {self.codebook = }, {self.embed_avg = }, {self.inited = }") |
| 64 | z = z.float() # Ensure fp32 |
| 65 | z_e = self.in_project(z).float() # (B, D', T), ensure fp32 |
| 66 | |
| 67 | # Rearrange for quantization |
| 68 | encodings = rearrange(z_e, "b d t -> (b t) d").float() # (B*T, D'), ensure fp32 |
| 69 | |
| 70 | # Quantization |
| 71 | dist = ( |
| 72 | encodings.pow(2).sum(1, keepdim=True) # (B*T, 1) |
| 73 | - 2 * encodings @ self.codebook.float().t() # (B*T, codebook_size) |