| 163 | return return_dict |
| 164 | |
| 165 | def get_codebook_entry(self, indices, shape=None): |
| 166 | # shape specifying (batch, height, width, channel) |
| 167 | indices_shape = indices.shape |
| 168 | indices_flatten = rearrange(indices, '... -> (...)') |
| 169 | |
| 170 | # get quantized latent vectors |
| 171 | emb = self.get_emb() |
| 172 | z_q = F.embedding(indices_flatten, emb) |
| 173 | # z_q = self.embedding(indices_flatten) |
| 174 | if self.l2_normalized: |
| 175 | z_q = F.normalize(z_q, p=2, dim=-1) |
| 176 | |
| 177 | if shape is not None: |
| 178 | z_q = z_q.reshape(shape) |
| 179 | else: |
| 180 | z_q = z_q.reshape([*indices_shape, self.dim]) |
| 181 | return z_q |
| 182 | |
| 183 | def decode(self, indices): |
| 184 | return self.get_codebook_entry(indices) |