Download windformer.py from ESA-philab/WindFormer: direct link, hf CLI and curl.
- Browser
- Download file 10.1 kB
-
https://huggingface.co/ESA-philab/WindFormer/resolve/main/windformer.py
- Command line
-
hf download hf://ESA-philab/WindFormer/windformer.py
-
curl -L -o windformer.py https://huggingface.co/ESA-philab/WindFormer/resolve/main/windformer.py
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) | |