MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / PatchEmbed

Class PatchEmbed

semantic_sam/backbone/focal_dw.py:368–431  ·  view source on GitHub ↗

Image to Patch Embedding Args: 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. norm_layer (nn.Module, optional): Normalization

Source from the content-addressed store, hash-verified

366# return flops
367
368class PatchEmbed(nn.Module):
369 """ Image to Patch Embedding
370
371 Args:
372 patch_size (int): Patch token size. Default: 4.
373 in_chans (int): Number of input image channels. Default: 3.
374 embed_dim (int): Number of linear projection output channels. Default: 96.
375 norm_layer (nn.Module, optional): Normalization layer. Default: None
376 use_conv_embed (bool): Whether use overlapped convolution for patch embedding. Default: False
377 is_stem (bool): Is the stem block or not.
378 """
379
380 def __init__(self, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None, use_conv_embed=False, is_stem=False, use_pre_norm=False):
381 super().__init__()
382 patch_size = to_2tuple(patch_size)
383 self.patch_size = patch_size
384
385 self.in_chans = in_chans
386 self.embed_dim = embed_dim
387 self.use_pre_norm = use_pre_norm
388
389 if use_conv_embed:
390 # if we choose to use conv embedding, then we treat the stem and non-stem differently
391 if is_stem:
392 kernel_size = 7; padding = 3; stride = 4
393 else:
394 kernel_size = 3; padding = 1; stride = 2
395 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding)
396 else:
397 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
398
399 if self.use_pre_norm:
400 if norm_layer is not None:
401 self.norm = norm_layer(in_chans)
402 else:
403 self.norm = None
404 else:
405 if norm_layer is not None:
406 self.norm = norm_layer(embed_dim)
407 else:
408 self.norm = None
409
410 def forward(self, x):
411 """Forward function."""
412 B, C, H, W = x.size()
413 if W % self.patch_size[1] != 0:
414 x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))
415 if H % self.patch_size[0] != 0:
416 x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))
417
418 if self.use_pre_norm:
419 if self.norm is not None:
420 x = x.flatten(2).transpose(1, 2) # B Ph*Pw C
421 x = self.norm(x).transpose(1, 2).view(B, C, H, W)
422 x = self.proj(x)
423 else:
424 x = self.proj(x) # B C Wh Ww
425 if self.norm is not None:

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected