Inputs the output of the encoder network z and maps it to a discrete one-hot vector that is the index of the closest embedding vector e_j z (continuous) -> z_q (discrete) z.shape = (batch, channel, height, width) quantization pipeline: 1. get enco
(self, z)
| 32 | self.embedding.weight.data.uniform_(-1.0 / self.n_e, 1.0 / self.n_e) |
| 33 | |
| 34 | def forward(self, z): |
| 35 | """ |
| 36 | Inputs the output of the encoder network z and maps it to a discrete |
| 37 | one-hot vector that is the index of the closest embedding vector e_j |
| 38 | z (continuous) -> z_q (discrete) |
| 39 | z.shape = (batch, channel, height, width) |
| 40 | quantization pipeline: |
| 41 | 1. get encoder input (B,C,H,W) |
| 42 | 2. flatten input to (B*H*W,C) |
| 43 | """ |
| 44 | # reshape z -> (batch, height, width, channel) and flatten |
| 45 | z = z.permute(0, 2, 3, 1).contiguous() |
| 46 | z_flattened = z.view(-1, self.e_dim) |
| 47 | # distances from z to embeddings e_j (z - e)^2 = z^2 + e^2 - 2 e * z |
| 48 | |
| 49 | d = torch.sum(z_flattened ** 2, dim=1, keepdim=True) + \ |
| 50 | torch.sum(self.embedding.weight**2, dim=1) - 2 * \ |
| 51 | torch.matmul(z_flattened, self.embedding.weight.t()) |
| 52 | |
| 53 | ## could possible replace this here |
| 54 | # #\start... |
| 55 | # find closest encodings |
| 56 | min_encoding_indices = torch.argmin(d, dim=1).unsqueeze(1) |
| 57 | |
| 58 | min_encodings = torch.zeros( |
| 59 | min_encoding_indices.shape[0], self.n_e).to(z) |
| 60 | min_encodings.scatter_(1, min_encoding_indices, 1) |
| 61 | |
| 62 | # dtype min encodings: torch.float32 |
| 63 | # min_encodings shape: torch.Size([2048, 512]) |
| 64 | # min_encoding_indices.shape: torch.Size([2048, 1]) |
| 65 | |
| 66 | # get quantized latent vectors |
| 67 | z_q = torch.matmul(min_encodings, self.embedding.weight).view(z.shape) |
| 68 | #.........\end |
| 69 | |
| 70 | # with: |
| 71 | # .........\start |
| 72 | #min_encoding_indices = torch.argmin(d, dim=1) |
| 73 | #z_q = self.embedding(min_encoding_indices) |
| 74 | # ......\end......... (TODO) |
| 75 | |
| 76 | # compute loss for embedding |
| 77 | loss = torch.mean((z_q.detach()-z)**2) + self.beta * \ |
| 78 | torch.mean((z_q - z.detach()) ** 2) |
| 79 | |
| 80 | # preserve gradients |
| 81 | z_q = z + (z_q - z).detach() |
| 82 | |
| 83 | # perplexity |
| 84 | e_mean = torch.mean(min_encodings, dim=0) |
| 85 | perplexity = torch.exp(-torch.sum(e_mean * torch.log(e_mean + 1e-10))) |
| 86 | |
| 87 | # reshape back to match original input shape |
| 88 | z_q = z_q.permute(0, 3, 1, 2).contiguous() |
| 89 | |
| 90 | return z_q, loss, (perplexity, min_encodings, min_encoding_indices) |
| 91 |
nothing calls this directly
no outgoing calls
no test coverage detected