Download model/afno.py from OneScience-Group/W-MAE: direct link, hf CLI and curl.
- Browser
- Download file 5.79 kB
-
https://huggingface.co/OneScience-Group/W-MAE/resolve/main/model/afno.py
- Command line
-
hf download hf://OneScience-Group/W-MAE/model/afno.py
-
curl -L -o afno.py https://huggingface.co/OneScience-Group/W-MAE/resolve/main/model/afno.py
5.79 kB
| from __future__ import annotations | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| class Mlp(nn.Module): | |
| def __init__( | |
| self, | |
| in_features: int, | |
| hidden_features: int | None = None, | |
| out_features: int | None = None, | |
| act_layer: type[nn.Module] = nn.GELU, | |
| drop: float = 0.0, | |
| ) -> None: | |
| super().__init__() | |
| hidden_features = hidden_features or in_features | |
| out_features = out_features or in_features | |
| self.fc1 = nn.Linear(in_features, hidden_features) | |
| self.act = act_layer() | |
| self.fc2 = nn.Linear(hidden_features, out_features) | |
| self.drop = nn.Dropout(drop) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.fc1(x) | |
| x = self.act(x) | |
| x = self.drop(x) | |
| x = self.fc2(x) | |
| return self.drop(x) | |
| class AFNO2D(nn.Module): | |
| """Adaptive Fourier mixing over a channel-last 2D patch grid.""" | |
| def __init__( | |
| self, | |
| hidden_size: int, | |
| num_blocks: int = 8, | |
| sparsity_threshold: float = 0.01, | |
| hard_thresholding_fraction: float = 1.0, | |
| hidden_size_factor: int = 1, | |
| ) -> None: | |
| super().__init__() | |
| if hidden_size % num_blocks != 0: | |
| raise ValueError(f"hidden_size={hidden_size} must be divisible by num_blocks={num_blocks}.") | |
| if not 0.0 < hard_thresholding_fraction <= 1.0: | |
| raise ValueError("hard_thresholding_fraction must be in (0, 1].") | |
| self.hidden_size = hidden_size | |
| self.sparsity_threshold = sparsity_threshold | |
| self.num_blocks = num_blocks | |
| self.block_size = hidden_size // num_blocks | |
| self.hard_thresholding_fraction = hard_thresholding_fraction | |
| self.hidden_size_factor = hidden_size_factor | |
| scale = 0.02 | |
| self.w1 = nn.Parameter( | |
| scale * torch.randn(2, num_blocks, self.block_size, self.block_size * hidden_size_factor) | |
| ) | |
| self.b1 = nn.Parameter(scale * torch.randn(2, num_blocks, self.block_size * hidden_size_factor)) | |
| self.w2 = nn.Parameter( | |
| scale * torch.randn(2, num_blocks, self.block_size * hidden_size_factor, self.block_size) | |
| ) | |
| self.b2 = nn.Parameter(scale * torch.randn(2, num_blocks, self.block_size)) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| if x.ndim != 4 or x.shape[-1] != self.hidden_size: | |
| raise ValueError(f"AFNO2D expects [B,H,W,{self.hidden_size}], got {tuple(x.shape)}.") | |
| bias = x | |
| dtype = x.dtype | |
| x = torch.fft.rfft2(x.float(), dim=(1, 2), norm="ortho") | |
| batch, height, freq_width, _ = x.shape | |
| x = x.reshape(batch, height, freq_width, self.num_blocks, self.block_size) | |
| hidden_block = self.block_size * self.hidden_size_factor | |
| o1_real = torch.zeros( | |
| batch, height, freq_width, self.num_blocks, hidden_block, device=x.device, dtype=x.real.dtype | |
| ) | |
| o1_imag = torch.zeros_like(o1_real) | |
| o2_real = torch.zeros_like(x.real) | |
| o2_imag = torch.zeros_like(x.real) | |
| total_modes = height // 2 + 1 | |
| kept_modes = max(1, int(total_modes * self.hard_thresholding_fraction)) | |
| row_slice = slice(total_modes - kept_modes, total_modes + kept_modes) | |
| col_slice = slice(0, min(kept_modes, freq_width)) | |
| selected = x[:, row_slice, col_slice] | |
| o1_real[:, row_slice, col_slice] = F.relu( | |
| torch.einsum("...bi,bio->...bo", selected.real, self.w1[0]) | |
| - torch.einsum("...bi,bio->...bo", selected.imag, self.w1[1]) | |
| + self.b1[0] | |
| ) | |
| o1_imag[:, row_slice, col_slice] = F.relu( | |
| torch.einsum("...bi,bio->...bo", selected.imag, self.w1[0]) | |
| + torch.einsum("...bi,bio->...bo", selected.real, self.w1[1]) | |
| + self.b1[1] | |
| ) | |
| hidden_real = o1_real[:, row_slice, col_slice] | |
| hidden_imag = o1_imag[:, row_slice, col_slice] | |
| o2_real[:, row_slice, col_slice] = ( | |
| torch.einsum("...bi,bio->...bo", hidden_real, self.w2[0]) | |
| - torch.einsum("...bi,bio->...bo", hidden_imag, self.w2[1]) | |
| + self.b2[0] | |
| ) | |
| o2_imag[:, row_slice, col_slice] = ( | |
| torch.einsum("...bi,bio->...bo", hidden_imag, self.w2[0]) | |
| + torch.einsum("...bi,bio->...bo", hidden_real, self.w2[1]) | |
| + self.b2[1] | |
| ) | |
| x = torch.stack((o2_real, o2_imag), dim=-1) | |
| x = F.softshrink(x, lambd=self.sparsity_threshold) | |
| x = torch.view_as_complex(x) | |
| x = x.reshape(batch, height, freq_width, self.hidden_size) | |
| x = torch.fft.irfft2(x, s=bias.shape[1:3], dim=(1, 2), norm="ortho") | |
| return x.to(dtype=dtype) + bias | |
| class Block(nn.Module): | |
| def __init__( | |
| self, | |
| dim: int, | |
| mlp_ratio: float = 4.0, | |
| drop: float = 0.0, | |
| act_layer: type[nn.Module] = nn.GELU, | |
| norm_layer: type[nn.Module] = nn.LayerNorm, | |
| double_skip: bool = True, | |
| num_blocks: int = 8, | |
| sparsity_threshold: float = 0.01, | |
| hard_thresholding_fraction: float = 1.0, | |
| ) -> None: | |
| super().__init__() | |
| self.norm1 = norm_layer(dim) | |
| self.filter = AFNO2D(dim, num_blocks, sparsity_threshold, hard_thresholding_fraction) | |
| self.drop_path = nn.Identity() | |
| self.norm2 = norm_layer(dim) | |
| self.mlp = Mlp(dim, int(dim * mlp_ratio), act_layer=act_layer, drop=drop) | |
| self.double_skip = double_skip | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| residual = x | |
| x = self.filter(self.norm1(x)) | |
| if self.double_skip: | |
| x = x + residual | |
| residual = x | |
| x = self.drop_path(self.mlp(self.norm2(x))) | |
| return x + residual | |