MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / PatchEmbed

Class PatchEmbed

Image_Classification/src/models/focalnet.py:315–373  ·  view source on GitHub ↗

r""" Image to Patch Embedding Args: img_size (int): Image size. Default: 224. patch_size (int): Patch token size. Default: 4. in_chans (int): Number of input image channels. Default: 3. embed_dim (int): Number of linear projection output channels. Default: 96.

Source from the content-addressed store, hash-verified

313
314
315class PatchEmbed(nn.Module):
316 r""" Image to Patch Embedding
317
318 Args:
319 img_size (int): Image size. Default: 224.
320 patch_size (int): Patch token size. Default: 4.
321 in_chans (int): Number of input image channels. Default: 3.
322 embed_dim (int): Number of linear projection output channels. Default: 96.
323 norm_layer (nn.Module, optional): Normalization layer. Default: None
324 """
325
326 def __init__(self, img_size=(224, 224), patch_size=4, in_chans=3, embed_dim=96, use_conv_embed=False,
327 norm_layer=None, is_stem=False):
328 super().__init__()
329 patch_size = to_2tuple(patch_size)
330 patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]]
331 self.img_size = img_size
332 self.patch_size = patch_size
333 self.patches_resolution = patches_resolution
334 self.num_patches = patches_resolution[0] * patches_resolution[1]
335
336 self.in_chans = in_chans
337 self.embed_dim = embed_dim
338
339 if use_conv_embed:
340 # if we choose to use conv embedding, then we treat the stem and non-stem differently
341 if is_stem:
342 kernel_size = 7;
343 padding = 2;
344 stride = 4
345 else:
346 kernel_size = 3;
347 padding = 1;
348 stride = 2
349 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding)
350 else:
351 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
352
353 if norm_layer is not None:
354 self.norm = norm_layer(embed_dim)
355 else:
356 self.norm = None
357
358 def forward(self, x):
359 B, C, H, W = x.shape
360
361 x = self.proj(x)
362 H, W = x.shape[2:]
363 x = x.flatten(2).transpose(1, 2) # B Ph*Pw C
364 if self.norm is not None:
365 x = self.norm(x)
366 return x, H, W
367
368 def flops(self):
369 Ho, Wo = self.patches_resolution
370 flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1])
371 if self.norm is not None:
372 flops += Ho * Wo * self.embed_dim

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected