MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

sat/sgm/modules/autoencoding/vqvae/quantize.py:137–176  ·  view source on GitHub ↗
(
        self,
        num_hiddens,
        embedding_dim,
        n_embed,
        straight_through=True,
        kl_weight=5e-4,
        temp_init=1.0,
        use_vqinterface=True,
        remap=None,
        unknown_index="random",
    )

Source from the content-addressed store, hash-verified

135 """
136
137 def __init__(
138 self,
139 num_hiddens,
140 embedding_dim,
141 n_embed,
142 straight_through=True,
143 kl_weight=5e-4,
144 temp_init=1.0,
145 use_vqinterface=True,
146 remap=None,
147 unknown_index="random",
148 ):
149 super().__init__()
150
151 self.embedding_dim = embedding_dim
152 self.n_embed = n_embed
153
154 self.straight_through = straight_through
155 self.temperature = temp_init
156 self.kl_weight = kl_weight
157
158 self.proj = nn.Conv2d(num_hiddens, n_embed, 1)
159 self.embed = nn.Embedding(n_embed, embedding_dim)
160
161 self.use_vqinterface = use_vqinterface
162
163 self.remap = remap
164 if self.remap is not None:
165 self.register_buffer("used", torch.tensor(np.load(self.remap)))
166 self.re_embed = self.used.shape[0]
167 self.unknown_index = unknown_index # "random" or "extra" or integer
168 if self.unknown_index == "extra":
169 self.unknown_index = self.re_embed
170 self.re_embed = self.re_embed + 1
171 print(
172 f"Remapping {self.n_embed} indices to {self.re_embed} indices. "
173 f"Using {self.unknown_index} for unknown indices."
174 )
175 else:
176 self.re_embed = n_embed
177
178 def remap_to_used(self, inds):
179 ishape = inds.shape

Callers 1

__init__Method · 0.45

Calls 2

printFunction · 0.50
loadMethod · 0.45

Tested by

no test coverage detected