WindFormer / windformer.py
ESA-philab:service:esawaai-upload's picture
Sync WindFormer model weights and configuration
643929e verified
Raw History Blame Contribute Delete
10.1 kB
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)