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

Class PatchMerging

Image_Classification/src/models/swin.py:156–168  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

154
155
156class PatchMerging(nn.Module):
157 def __init__(self, in_channels, out_channels, downscaling_factor):
158 super().__init__()
159 self.downscaling_factor = downscaling_factor
160 self.patch_merge = nn.Unfold(kernel_size=downscaling_factor, stride=downscaling_factor, padding=0)
161 self.linear = nn.Linear(in_channels * downscaling_factor ** 2, out_channels)
162
163 def forward(self, x):
164 b, c, h, w = x.shape
165 new_h, new_w = h // self.downscaling_factor, w // self.downscaling_factor
166 x = self.patch_merge(x).view(b, -1, new_h, new_w).permute(0, 2, 3, 1)
167 x = self.linear(x)
168 return x
169
170
171class StageModule(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected