import torch import torch.nn as nn import torch.nn.functional as F import numbers from einops import rearrange ########################################################################## ## Layer Norm def to_3d(x): return rearrange(x, 'b c h w -> b (h w) c') def to_4d(x, h, w): return rearrange(x, 'b (h w) c -> b c h w', h=h, w=w) class BiasFree_LayerNorm(nn.Module): def __init__(self, normalized_shape): super(BiasFree_LayerNorm, self).__init__() if isinstance(normalized_shape, numbers.Integral): normalized_shape = (normalized_shape,) normalized_shape = torch.Size(normalized_shape) assert len(normalized_shape) == 1 self.weight = nn.Parameter(torch.ones(normalized_shape)) self.normalized_shape = normalized_shape def forward(self, x): sigma = x.var(-1, keepdim=True, unbiased=False) return x / torch.sqrt(sigma + 1e-5) * self.weight class WithBias_LayerNorm(nn.Module): def __init__(self, normalized_shape): super(WithBias_LayerNorm, self).__init__() if isinstance(normalized_shape, numbers.Integral): normalized_shape = (normalized_shape,) normalized_shape = torch.Size(normalized_shape) assert len(normalized_shape) == 1 self.weight = nn.Parameter(torch.ones(normalized_shape)) self.bias = nn.Parameter(torch.zeros(normalized_shape)) self.normalized_shape = normalized_shape def forward(self, x): mu = x.mean(-1, keepdim=True) sigma = x.var(-1, keepdim=True, unbiased=False) return (x - mu) / torch.sqrt(sigma + 1e-5) * self.weight + self.bias class LayerNorm(nn.Module): def __init__(self, dim, LayerNorm_type): super(LayerNorm, self).__init__() if LayerNorm_type == 'BiasFree': self.body = BiasFree_LayerNorm(dim) else: self.body = WithBias_LayerNorm(dim) def forward(self, x): h, w = x.shape[-2:] return to_4d(self.body(to_3d(x)), h, w) ########################################################################## ## Gated-Dconv Feed-Forward Network (GDFN) class FeedForward(nn.Module): def __init__(self, dim, ffn_expansion_factor, bias): super(FeedForward, self).__init__() hidden_features = int(dim * ffn_expansion_factor) self.project_in = nn.Conv2d(dim, hidden_features * 2, kernel_size=1, bias=bias) self.dwconv = nn.Conv2d(hidden_features * 2, hidden_features * 2, kernel_size=3, stride=1, padding=1, groups=hidden_features * 2, bias=bias) self.project_out = nn.Conv2d(hidden_features, dim, kernel_size=1, bias=bias) def forward(self, x): x = self.project_in(x) x1, x2 = self.dwconv(x).chunk(2, dim=1) x = F.gelu(x1) * x2 x = self.project_out(x) return x ########################################################################## ## Multi-DConv Head Transposed Self-Attention (MDTA) class Attention(nn.Module): def __init__(self, dim, num_heads, bias): super(Attention, self).__init__() self.num_heads = num_heads self.temperature = nn.Parameter(torch.ones(num_heads, 1, 1)) self.qkv = nn.Conv2d(dim, dim * 3, kernel_size=1, bias=bias) self.qkv_dwconv = nn.Conv2d(dim * 3, dim * 3, kernel_size=3, stride=1, padding=1, groups=dim * 3, bias=bias) self.project_out = nn.Conv2d(dim, dim, kernel_size=1, bias=bias) def forward(self, x): b, c, h, w = x.shape qkv = self.qkv_dwconv(self.qkv(x)) q, k, v = qkv.chunk(3, dim=1) q = rearrange(q, 'b (head c) h w -> b head c (h w)', head=self.num_heads) k = rearrange(k, 'b (head c) h w -> b head c (h w)', head=self.num_heads) v = rearrange(v, 'b (head c) h w -> b head c (h w)', head=self.num_heads) q = torch.nn.functional.normalize(q, dim=-1) k = torch.nn.functional.normalize(k, dim=-1) attn = (q @ k.transpose(-2, -1)) * self.temperature attn = attn.softmax(dim=-1) out = (attn @ v) out = rearrange(out, 'b head c (h w) -> b (head c) h w', head=self.num_heads, h=h, w=w) out = self.project_out(out) return out ########################################################################## class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, ffn_expansion_factor, bias, LayerNorm_type): super(TransformerBlock, self).__init__() self.norm1 = LayerNorm(dim, LayerNorm_type) self.attn = Attention(dim, num_heads, bias) self.norm2 = LayerNorm(dim, LayerNorm_type) self.ffn = FeedForward(dim, ffn_expansion_factor, bias) def forward(self, x): x = x + self.attn(self.norm1(x)) x = x + self.ffn(self.norm2(x)) return x ########################################################################## ## Overlapped image patch embedding with 3x3 Conv class OverlapPatchEmbed(nn.Module): def __init__(self, in_c=3, embed_dim=48, bias=False): super(OverlapPatchEmbed, self).__init__() self.proj1 = nn.Conv2d(in_c, embed_dim, kernel_size=3, stride=1, padding=1, bias=bias) self.activation = nn.LeakyReLU(0.1) self.bn1 = nn.BatchNorm2d(embed_dim) self.proj2 = nn.Conv2d(embed_dim, embed_dim, kernel_size=3, stride=1, padding=1, bias=bias) def forward(self, x): x = self.proj1(x) x = self.bn1(x) x = self.activation(x) x = self.proj2(x) return x ########################################################################## ## Resizing modules class Downsample(nn.Module): def __init__(self, n_feat): super(Downsample, self).__init__() self.body = nn.Sequential(nn.Conv2d(n_feat, n_feat // 2, kernel_size=3, stride=1, padding=1, bias=False), nn.PixelUnshuffle(2)) def forward(self, x): return self.body(x) class Upsample(nn.Module): def __init__(self, n_feat): super(Upsample, self).__init__() self.body = nn.Sequential(nn.Conv2d(n_feat, n_feat * 2, kernel_size=3, stride=1, padding=1, bias=False), nn.PixelShuffle(2)) def forward(self, x): return self.body(x) ########################################################################## ##---------- WindFormer ----------------------- class WindFormer(nn.Module): def __init__(self, inp_channels=3, out_channels=1, dim=64, # 16 before num_blocks=[12], num_refinement_blocks=3, heads=[8], ffn_expansion_factor=2.66, bias=True, LayerNorm_type='BiasFree', ## Other option 'BiasFree' fusion=False ): super(WindFormer, self).__init__() self.patch_embed = OverlapPatchEmbed(inp_channels, dim) self.encoder_level1 = nn.Sequential(*[ TransformerBlock(dim=dim, num_heads=heads[0], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type) for i in range(num_blocks[0])]) self.output = nn.Conv2d(dim, out_channels, kernel_size=3, stride=1, padding=1, bias=bias) if fusion: self.output_ocn = nn.Conv2d(dim, out_channels, kernel_size=3, stride=1, padding=1, bias=bias) else: self.output_ocn = None # self.dropout = nn.Dropout2d(0.10) def forward(self, inp_img): x = self.patch_embed(inp_img) # x = self.dropout(x) # x = F.dropout(x, p=0.05, training=True) x = self.encoder_level1(x) x_nora3 = self.output(x) if self.output_ocn is not None: x_ocn = self.output_ocn(x) x = torch.cat((x_nora3, x_ocn), 1) else: x = x_nora3 x = torch.sigmoid(x) return x class WindFormerDist(nn.Module): def __init__(self, inp_channels=4, out_channels=1, dim=64, # 16 before num_blocks=[12], num_refinement_blocks=3, heads=[8], ffn_expansion_factor=2.66, bias=True, LayerNorm_type='BiasFree', ## Other option 'BiasFree' fusion=False ): super(WindFormerDist, self).__init__() self.patch_embed = OverlapPatchEmbed(inp_channels, dim) self.encoder_level1 = nn.Sequential(*[ TransformerBlock(dim=dim, num_heads=heads[0], ffn_expansion_factor=ffn_expansion_factor, bias=bias, LayerNorm_type=LayerNorm_type) for i in range(num_blocks[0])]) self.output = nn.Conv2d(dim, 20 * out_channels, kernel_size=3, stride=1, padding=1, bias=bias) if fusion: self.output_ocn = nn.Conv2d(dim, out_channels, kernel_size=3, stride=1, padding=1, bias=bias) else: self.output_ocn = None self.output_mean = nn.Conv2d(20 * out_channels, out_channels, kernel_size=1, stride=1, bias=True) self.output_logvar = nn.Conv2d(20 * out_channels, out_channels, kernel_size=1, stride=1, bias=True) # self.dropout = nn.Dropout2d(0.10) def forward(self, inp_img): x = self.patch_embed(inp_img) # x = self.dropout(x) x = self.encoder_level1(x) x_nora3 = self.output(x) if self.output_ocn is not None: x_ocn = self.output_ocn(x) x = torch.cat((x_nora3, x_ocn), 1) else: x = x_nora3 x_mean = torch.sigmoid(self.output_mean(x)) x_logvar = self.output_logvar(x) return torch.cat((x_mean, x_logvar), 1)