MCPcopy Create free account
hub / github.com/apple/ml-4m / ImageEncoderEmbedding

Class ImageEncoderEmbedding

fourm/models/encoder_embeddings.py:214–309  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

212
213
214class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected