Quantized the input tensor using a fixed codebook and returns the corresponding codebook vectors Parameters ---------- z : Tensor[B x D x T] Returns ------- Tensor[B x D x T] Quantized continuous representation of input Te
(self, z)
| 359 | self.codebook = nn.Embedding(codebook_size, codebook_dim) |
| 360 | |
| 361 | def forward(self, z): |
| 362 | """Quantized the input tensor using a fixed codebook and returns |
| 363 | the corresponding codebook vectors |
| 364 | |
| 365 | Parameters |
| 366 | ---------- |
| 367 | z : Tensor[B x D x T] |
| 368 | |
| 369 | Returns |
| 370 | ------- |
| 371 | Tensor[B x D x T] |
| 372 | Quantized continuous representation of input |
| 373 | Tensor[1] |
| 374 | Commitment loss to train encoder to predict vectors closer to codebook |
| 375 | entries |
| 376 | Tensor[1] |
| 377 | Codebook loss to update the codebook |
| 378 | Tensor[B x T] |
| 379 | Codebook indices (quantized discrete representation of input) |
| 380 | Tensor[B x D x T] |
| 381 | Projected latents (continuous representation of input before quantization) |
| 382 | """ |
| 383 | |
| 384 | # Factorized codes (ViT-VQGAN) Project input into low-dimensional space |
| 385 | z_e = self.in_proj(z) # z_e : (B x D x T) |
| 386 | z_q, indices = self.decode_latents(z_e) |
| 387 | |
| 388 | commitment_loss = F.mse_loss(z_e, z_q.detach(), reduction="none").mean([1, 2]) |
| 389 | codebook_loss = F.mse_loss(z_q, z_e.detach(), reduction="none").mean([1, 2]) |
| 390 | |
| 391 | z_q = ( |
| 392 | z_e + (z_q - z_e).detach() |
| 393 | ) # noop in forward pass, straight-through gradient estimator in backward pass |
| 394 | |
| 395 | z_q = self.out_proj(z_q) |
| 396 | |
| 397 | return z_q, commitment_loss, codebook_loss, indices, z_e |
| 398 | |
| 399 | def embed_code(self, embed_id): |
| 400 | return F.embedding(embed_id, self.codebook.weight) |
nothing calls this directly
no test coverage detected