W-MAE / model /afno.py
yzt15806542928's picture
Upload folder using huggingface_hub
80cf062 verified
Raw
History Blame Contribute Delete
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