MCPcopy Create free account
hub / github.com/AMAP-ML/EMF / get_codebook_entry

Method get_codebook_entry

tok/ar_dtok/bottleneck.py:165–181  ·  view source on GitHub ↗
(self, indices, shape=None)

Source from the content-addressed store, hash-verified

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)

Callers 1

decodeMethod · 0.95

Calls 1

get_embMethod · 0.95

Tested by

no test coverage detected