Embedding module for spatial inputs, like images or feature maps. Creates tokens from patches over the image. This adapter / embedding differs from the one of MultiMAE by taking as input a dict and separating positional embeddings and modality embeddings from the input projection
| 212 | |
| 213 | |
| 214 | class ImageEncoderEmbedding(nn.Module): |
| 215 | """Embedding module for spatial inputs, like images or feature maps. |
| 216 | Creates tokens from patches over the image. |
| 217 | |
| 218 | This adapter / embedding differs from the one of MultiMAE by taking as input a dict and |
| 219 | separating positional embeddings and modality embeddings from the input projection |
| 220 | Input projection is 'x', posemb + modemb is 'emb' |
| 221 | |
| 222 | Args: |
| 223 | num_channels: Number of input channels of the image/feature map |
| 224 | patch_size: Int or tuple of the patch size over the full image size. |
| 225 | dim_tokens: Dimension of output tokens. Can be set using init method. |
| 226 | sincos_pos_emb: Set to True (default) to use fixed 2D sin-cos positional embeddings |
| 227 | image_size: Default image size. Used to initialize size of positional embeddings. |
| 228 | """ |
| 229 | def __init__(self, |
| 230 | num_channels: int, |
| 231 | patch_size: Union[int, Tuple[int,int]], |
| 232 | dim_tokens: Optional[int] = None, |
| 233 | sincos_pos_emb: bool = True, |
| 234 | image_size: Union[int, Tuple[int]] = 224): |
| 235 | |
| 236 | super().__init__() |
| 237 | self.num_channels = num_channels |
| 238 | self.patch_size = pair(patch_size) |
| 239 | self.dim_tokens = dim_tokens |
| 240 | self.sincos_pos_emb = sincos_pos_emb |
| 241 | self.image_size = pair(image_size) |
| 242 | self.num_patches = (self.image_size[0] // patch_size) * (self.image_size[1] // patch_size) |
| 243 | |
| 244 | if self.dim_tokens is not None: |
| 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 |