(self, z: torch.Tensor)
| 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) |
| 74 | + self.codebook.float().pow(2).sum(1, keepdim=True).t() |
| 75 | ) # (1, codebook_size) |
| 76 | |
| 77 | # dist: (B*T, codebook_size) |
| 78 | indices = (-dist).max(1)[1] # (B*T) |
| 79 | indices = rearrange(indices, "(b t) -> b t", b=z.size(0)) # (B, T) |
| 80 | |
| 81 | # Get quantized vectors |
| 82 | z_q = self.decode_code(indices).float() # (B, D', T), ensure fp32 |
| 83 | |
| 84 | # Straight-through estimator |
| 85 | z_q = z_e + (z_q - z_e).detach() # (B, D', T) |
| 86 | z_q = self.out_project(z_q).float() # (B, D, T), ensure fp32 |
| 87 | |
| 88 | # z_q: (B, D, T), commit_loss: (B), indices: (B, T), z: (B, D', T) |
| 89 | return z_q, indices |
| 90 | |
| 91 | |
| 92 | class ResidualVQ(nn.Module): |
no test coverage detected