Initialize parts of encoder 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)
| 245 | self.init(dim_tokens=dim_tokens) |
| 246 | |
| 247 | def init(self, dim_tokens: int = 768, init_std=0.02): |
| 248 | """ |
| 249 | Initialize parts of encoder that are dependent on dimension of tokens. |
| 250 | Should be called when setting up FourM. |
| 251 | |
| 252 | Args: |
| 253 | dim_tokens: Dimension of tokens |
| 254 | init_std: Standard deviation of init |
| 255 | """ |
| 256 | self.dim_tokens = dim_tokens |
| 257 | |
| 258 | # Task embedding identifying from which task a given token comes from |
| 259 | # Fixed-size positional embeddings. Can be interpolated to different input sizes |
| 260 | h_posemb = self.image_size[0] // self.patch_size[0] |
| 261 | w_posemb = self.image_size[1] // self.patch_size[1] |
| 262 | if self.sincos_pos_emb: |
| 263 | pos_emb = build_2d_sincos_posemb(h=h_posemb, w=w_posemb, embed_dim=self.dim_tokens) |
| 264 | self.register_buffer("pos_emb", pos_emb) # self.pos_emb is now a buffer for FSDP |
| 265 | else: |
| 266 | self.pos_emb = nn.Parameter(torch.zeros(1, (h_posemb * w_posemb), self.dim_tokens)) |
| 267 | nn.init.normal_(self.pos_emb, std=init_std) |
| 268 | |
| 269 | self.mod_emb = nn.Parameter(torch.zeros(1, 1, self.dim_tokens)) |
| 270 | nn.init.normal_(self.mod_emb, std=init_std) |
| 271 | |
| 272 | # Image -> tokens projection |
| 273 | # No bias term here, so modality embedding fully comes from self.mod_emb |
| 274 | self.proj = nn.Linear(self.num_channels * self.patch_size[0] * self.patch_size[1], self.dim_tokens, bias=False) |
| 275 | |
| 276 | @torch.jit.ignore |
| 277 | def no_weight_decay(self): |
no test coverage detected