MCPcopy Create free account
hub / github.com/dek924/PerX2CT / forward

Method forward

taming/modules/vqvae/quantize.py:34–90  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected