Initialize parts of module that are dependent on dimension of tokens. Should be called when setting up FourM. Args: dim_tokens: Dimension of tokens init_std: Standard deviation of init
(self, dim_tokens: int = 768, init_std=0.02)
| 185 | self.init(dim_tokens=dim_tokens) |
| 186 | |
| 187 | def init(self, dim_tokens: int = 768, init_std=0.02): |
| 188 | """ |
| 189 | Initialize parts of module that are dependent on dimension of tokens. |
| 190 | Should be called when setting up FourM. |
| 191 | |
| 192 | Args: |
| 193 | dim_tokens: Dimension of tokens |
| 194 | init_std: Standard deviation of init |
| 195 | """ |
| 196 | self.dim_tokens = dim_tokens |
| 197 | |
| 198 | # Task embedding identifying from which task a given token comes from |
| 199 | # Fixed-size positional embeddings. Can be interpolated to different input sizes |
| 200 | h_posemb = self.image_size[0] // self.patch_size[0] |
| 201 | w_posemb = self.image_size[1] // self.patch_size[1] |
| 202 | if self.sincos_pos_emb: |
| 203 | pos_emb = build_2d_sincos_posemb(h=h_posemb, w=w_posemb, embed_dim=self.dim_tokens) |
| 204 | self.register_buffer("pos_emb", pos_emb) |
| 205 | else: |
| 206 | self.pos_emb = nn.Parameter(torch.zeros(1, (h_posemb * w_posemb), self.dim_tokens)) |
| 207 | nn.init.normal_(self.pos_emb, std=init_std) |
| 208 | |
| 209 | self.mod_emb = nn.Parameter(torch.zeros(1, 1, self.dim_tokens)) |
| 210 | nn.init.normal_(self.mod_emb, std=init_std) |
| 211 | |
| 212 | # Token embedding (not needed if only masked tokens are given as input, but can be useful to train Token Critic) |
| 213 | self.token_emb = nn.Embedding(num_embeddings=self.vocab_size, embedding_dim=self.dim_tokens) |
| 214 | |
| 215 | # Output projection layer |
| 216 | self.to_logits = nn.Linear(self.dim_tokens, self.vocab_size, bias=False) |
| 217 | |
| 218 | if self.share_embedding: |
| 219 | # Share input and output embedding weights |
| 220 | self.to_logits.weight = self.token_emb.weight |
| 221 | |
| 222 | @torch.jit.ignore |
| 223 | def no_weight_decay(self): |
no test coverage detected