in models/swin_transformer_3d.py [0:0]
def forward(self, x):
"""Forward function."""
# padding
_, _, D, H, W = x.size()
if W % self.patch_size[2] != 0:
x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
if H % self.patch_size[1] != 0:
x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
if D % self.patch_size[0] != 0:
x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - D % self.patch_size[0]))
if self.additional_variable_channels:
x_rgb = x[:, :3, ...]
x_rem = x[:, 3:, ...]
x_rgb = self.proj(x_rgb)
if x.shape[1] > 3:
x_rem = self.run_variable_channel_forward(x_rem)
x = x_rgb + x_rem
else:
x = x_rgb
else:
x = self.proj(x) # B C D Wh Ww
if self.norm is not None:
D, Wh, Ww = x.size(2), x.size(3), x.size(4)
x = x.flatten(2).transpose(1, 2)
x = self.norm(x)
x = x.transpose(1, 2).view(-1, self.embed_dim, D, Wh, Ww)
return x