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)
| 150 | self.init(dim_tokens=dim_tokens) |
| 151 | |
| 152 | def init(self, dim_tokens: int = 768, init_std=0.02): |
| 153 | """ |
| 154 | Initialize parts of module that are dependent on dimension of tokens. |
| 155 | Should be called when setting up FourM. |
| 156 | |
| 157 | Args: |
| 158 | dim_tokens: Dimension of tokens |
| 159 | init_std: Standard deviation of init |
| 160 | """ |
| 161 | self.dim_tokens = dim_tokens |
| 162 | |
| 163 | # Task embedding identifying from which task a given token comes from |
| 164 | # Fixed-size positional embeddings. Can be interpolated to different input sizes |
| 165 | h_posemb = self.image_size[0] // self.patch_size[0] |
| 166 | w_posemb = self.image_size[1] // self.patch_size[1] |
| 167 | if self.sincos_pos_emb: |
| 168 | pos_emb = build_2d_sincos_posemb(h=h_posemb, w=w_posemb, embed_dim=self.dim_tokens) |
| 169 | self.register_buffer("pos_emb", pos_emb) # self.pos_emb is now a buffer for FSDP |
| 170 | else: |
| 171 | self.pos_emb = nn.Parameter(torch.zeros(1, (h_posemb * w_posemb), self.dim_tokens)) |
| 172 | nn.init.normal_(self.pos_emb, std=init_std) |
| 173 | |
| 174 | self.mod_emb = nn.Parameter(torch.zeros(1, 1, self.dim_tokens)) |
| 175 | nn.init.normal_(self.mod_emb, std=init_std) |
| 176 | |
| 177 | # Token embedding |
| 178 | self.token_emb = nn.Embedding(num_embeddings=self.vocab_size, embedding_dim=self.dim_tokens) |
| 179 | |
| 180 | @torch.jit.ignore |
| 181 | def no_weight_decay(self): |
no test coverage detected