| """Trainable FuXi 2.1 forward-graph reconstruction.""" |
|
|
| from __future__ import annotations |
|
|
| import math |
| from torch.utils.checkpoint import checkpoint as activation_checkpoint |
|
|
| import torch |
| import torch.nn.functional as F |
| from torch import nn |
|
|
|
|
| DIAGNOSTIC_INDICES = (79, 80, 81, 82, 84) |
|
|
|
|
| class UnbiasedNorm(nn.Module): |
| """Layer normalization matching the PT2 graph's unbiased variance.""" |
|
|
| def __init__(self, dim: int, conditioned: bool = False, eps: float = 1e-6) -> None: |
| super().__init__() |
| self.eps = eps |
| self.weight = nn.Parameter(torch.ones(dim)) |
| self.conditioned = conditioned |
| self.scale_shift = nn.Sequential(nn.SiLU(), nn.Linear(dim, 2 * dim)) if conditioned else None |
|
|
| def forward(self, x: torch.Tensor, condition: torch.Tensor | None = None) -> torch.Tensor: |
| variance, mean = torch.var_mean(x, dim=-1, correction=1, keepdim=True) |
| x = (x - mean) * torch.rsqrt(variance + self.eps) * self.weight |
| if self.scale_shift is not None: |
| if condition is None: |
| raise ValueError("condition is required by conditioned normalization") |
| scale, shift = self.scale_shift(condition).chunk(2, dim=-1) |
| x = x * (1 + scale[:, None, :]) + shift[:, None, :] |
| return x |
|
|
|
|
| def _rope_frequencies(height: int, width: int, head_dim: int) -> tuple[torch.Tensor, torch.Tensor]: |
| if head_dim % 2: |
| raise ValueError("head_dim must be even for rotary embeddings") |
| y, x = torch.meshgrid(torch.arange(height), torch.arange(width), indexing="ij") |
| positions = (y * width + x).flatten().float() |
| frequencies = 1.0 / (10000 ** (torch.arange(0, head_dim, 2).float() / head_dim)) |
| angles = positions[:, None] * frequencies[None, :] |
| return angles.cos(), angles.sin() |
|
|
|
|
| def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: |
| even, odd = x[..., 0::2], x[..., 1::2] |
| cos = cos[None, :, None, :].to(dtype=x.dtype, device=x.device) |
| sin = sin[None, :, None, :].to(dtype=x.dtype, device=x.device) |
| return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2) |
|
|
|
|
| def _window_partition(x: torch.Tensor, window: int) -> torch.Tensor: |
| batch, height, width, channels = x.shape |
| return x.view(batch, height // window, window, width // window, window, channels).permute(0, 1, 3, 2, 4, 5).reshape(-1, window * window, channels) |
|
|
|
|
| def _window_reverse(x: torch.Tensor, batch: int, height: int, width: int, window: int) -> torch.Tensor: |
| return x.view(batch, height // window, width // window, window, window, -1).permute(0, 1, 3, 2, 4, 5).reshape(batch, height, width, -1) |
|
|
|
|
| def _shift_mask(height: int, width: int, window: int) -> torch.Tensor: |
| shift = window // 2 |
| labels = torch.zeros(1, height, width, 1) |
| h_slices = (slice(0, -window), slice(-window, -shift), slice(-shift, None)) |
| w_slices = (slice(0, -window), slice(-window, -shift), slice(-shift, None)) |
| index = 0 |
| for h_slice in h_slices: |
| for w_slice in w_slices: |
| labels[:, h_slice, w_slice] = index |
| index += 1 |
| labels = _window_partition(labels, window).squeeze(-1) |
| mask = labels[:, None, :] - labels[:, :, None] |
| return mask.masked_fill(mask != 0, float("-inf")).masked_fill(mask == 0, 0.0) |
|
|
|
|
| class HeadGatedWindowAttention(nn.Module): |
| def __init__(self, dim: int, num_heads: int, window: int, grid_size: tuple[int, int], shifted: bool) -> None: |
| super().__init__() |
| if dim % num_heads: |
| raise ValueError("dim must be divisible by num_heads") |
| self.num_heads = num_heads |
| self.head_dim = dim // num_heads |
| self.window = window |
| self.grid_size = grid_size |
| self.shifted = shifted |
| self.wq = nn.Linear(dim, num_heads * (self.head_dim + 1), bias=False) |
| self.wk = nn.Linear(dim, dim, bias=False) |
| self.wv = nn.Linear(dim, dim, bias=False) |
| self.wo = nn.Linear(dim, dim, bias=False) |
| cos, sin = _rope_frequencies(*grid_size, self.head_dim) |
| self.register_buffer("freqs_cos", cos, persistent=False) |
| self.register_buffer("freqs_sin", sin, persistent=False) |
| self.register_buffer("attention_mask", _shift_mask(*grid_size, window) if shifted else None, persistent=False) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| batch, tokens, channels = x.shape |
| height, width = self.grid_size |
| qg = self.wq(x).view(batch, tokens, self.num_heads, self.head_dim + 1) |
| q, gate = qg[..., : self.head_dim], qg[..., -1:].sigmoid() |
| k = self.wk(x).view(batch, tokens, self.num_heads, self.head_dim) |
| v = self.wv(x).view(batch, tokens, self.num_heads, self.head_dim) |
| q = _apply_rope(q, self.freqs_cos, self.freqs_sin).reshape(batch, height, width, channels) |
| k = _apply_rope(k, self.freqs_cos, self.freqs_sin).reshape(batch, height, width, channels) |
| v = v.reshape(batch, height, width, channels) |
| gate = gate.reshape(batch, height, width, self.num_heads, 1) |
|
|
| if self.shifted: |
| shift = self.window // 2 |
| q, k, v, gate = [torch.roll(item, shifts=(-shift, -shift), dims=(1, 2)) for item in (q, k, v, gate)] |
| q, k, v = [_window_partition(item, self.window).view(-1, self.window**2, self.num_heads, self.head_dim).transpose(1, 2) for item in (q, k, v)] |
| gate = _window_partition(gate.flatten(-2), self.window).view(-1, self.window**2, self.num_heads, 1).transpose(1, 2) |
| scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim) |
| if self.attention_mask is not None: |
| windows = self.attention_mask.shape[0] |
| scores = scores.view(batch, windows, self.num_heads, self.window**2, self.window**2) |
| scores = scores + self.attention_mask[None, :, None].to(scores) |
| scores = scores.flatten(0, 1) |
| output = torch.matmul(scores.softmax(dim=-1), v) * gate |
| output = output.transpose(1, 2).reshape(-1, self.window**2, channels) |
| output = _window_reverse(output, batch, height, width, self.window) |
| if self.shifted: |
| output = torch.roll(output, shifts=(self.window // 2, self.window // 2), dims=(1, 2)) |
| return self.wo(output.reshape(batch, tokens, channels)) |
|
|
|
|
| class FuXi21Block(nn.Module): |
| def __init__(self, dim: int, mlp_dim: int, num_heads: int, window: int, grid_size: tuple[int, int], shifted: bool) -> None: |
| super().__init__() |
| self.adaln = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim)) |
| self.norm1 = UnbiasedNorm(dim) |
| self.attn = HeadGatedWindowAttention(dim, num_heads, window, grid_size, shifted) |
| self.norm2 = UnbiasedNorm(dim) |
| self.w1 = nn.Linear(dim, mlp_dim, bias=False) |
| self.w2 = nn.Linear(mlp_dim, dim, bias=False) |
| self.w3 = nn.Linear(dim, mlp_dim, bias=False) |
|
|
| def forward(self, x: torch.Tensor, condition: torch.Tensor) -> torch.Tensor: |
| attn_scale, attn_shift, attn_gate, mlp_scale, mlp_shift, mlp_gate = self.adaln(condition).chunk(6, dim=-1) |
| normalized = self.norm1(x) * (1 + attn_scale[:, None]) + attn_shift[:, None] |
| x = x + attn_gate[:, None] * self.attn(normalized) |
| normalized = self.norm2(x) * (1 + mlp_scale[:, None]) + mlp_shift[:, None] |
| mlp = self.w2(F.silu(self.w1(normalized)) * self.w3(normalized)) |
| return x + mlp_gate[:, None] * mlp |
|
|
|
|
| class PixelShuffleHead(nn.Module): |
| def __init__(self, dim: int, output_channels: int) -> None: |
| super().__init__() |
| self.conv1 = nn.Conv2d(dim, 2 * dim, 3, padding=1) |
| self.conv2 = nn.Conv2d(dim // 2, output_channels * 9, 3, padding=1) |
|
|
| def forward(self, x: torch.Tensor, output_size: tuple[int, int]) -> torch.Tensor: |
| x = F.pad(x, (0, 0, 0, 1), mode="replicate") |
| x = F.gelu(F.pixel_shuffle(self.conv1(x), 2)) |
| x = F.pixel_shuffle(self.conv2(x), 3) |
| return x[..., : output_size[0], : output_size[1]] |
|
|
|
|
| class FuXi21(nn.Module): |
| """Randomly initialized, trainable reconstruction of the FuXi 2.1 PT2 forward graph. |
| |
| The defaults reproduce the recovered architecture; reduced dimensions and grids |
| are intended for smoke tests. |
| """ |
|
|
| def __init__( |
| self, |
| static_fields: torch.Tensor, |
| channel_mask: torch.Tensor, |
| grid_size: tuple[int, int] = (721, 1440), |
| embed_dim: int = 1536, |
| depth: int = 30, |
| num_heads: int = 24, |
| mlp_dim: int = 4096, |
| patch_size: int = 6, |
| window_size: int = 20, |
| activation_checkpointing: bool = False, |
| ) -> None: |
| super().__init__() |
| height, width = grid_size |
| token_grid = (height // patch_size, width // patch_size) |
| if patch_size != 6: |
| raise ValueError("The recovered PixelShuffle decoder requires patch_size=6") |
| if any(size % window_size for size in token_grid): |
| raise ValueError(f"token grid {token_grid} must be divisible by window_size={window_size}") |
| if static_fields.shape != (6, height, width): |
| raise ValueError(f"static_fields must have shape {(6, height, width)}, got {tuple(static_fields.shape)}") |
| if channel_mask.shape != (85, height, width): |
| raise ValueError(f"channel_mask must have shape {(85, height, width)}, got {tuple(channel_mask.shape)}") |
| if embed_dim % 4: |
| raise ValueError("embed_dim must be divisible by 4 for the PixelShuffle heads") |
|
|
| self.grid_size = grid_size |
| self.token_grid = token_grid |
| self.activation_checkpointing = activation_checkpointing |
| self.register_buffer("static_fields", static_fields.detach().float()) |
| self.register_buffer("channel_mask", channel_mask.detach().float()) |
| self.patch_embed = nn.Conv2d(170, embed_dim, patch_size, stride=patch_size) |
| self.patch_norm = UnbiasedNorm(embed_dim) |
| self.const_embed = nn.Conv2d(6, embed_dim, patch_size, stride=patch_size) |
| self.const_norm = UnbiasedNorm(embed_dim) |
| self.joint_embed_layer = nn.Sequential(nn.Linear(384, embed_dim), nn.SiLU(), nn.Linear(embed_dim, embed_dim)) |
| self.layers = nn.ModuleList( |
| FuXi21Block(embed_dim, mlp_dim, num_heads, window_size, token_grid, bool(index % 2)) |
| for index in range(depth) |
| ) |
| self.norm_layer = UnbiasedNorm(embed_dim, conditioned=True) |
| self.pressure_head = nn.ConvTranspose2d(embed_dim, 65, 9, stride=6, padding=1) |
| self.surface_head = PixelShuffleHead(embed_dim, 15) |
| self.derived_head = PixelShuffleHead(embed_dim, 5) |
| self.register_buffer("scatter_idx", torch.tensor([*range(79), 83, 79, 80, 81, 82, 84]), persistent=False) |
| self.reset_parameters() |
|
|
| @classmethod |
| def smoke( |
| cls, |
| static_fields: torch.Tensor | None = None, |
| channel_mask: torch.Tensor | None = None, |
| ) -> "FuXi21": |
| static_fields = torch.zeros(6, 13, 12) if static_fields is None else static_fields |
| channel_mask = torch.ones(85, 13, 12) if channel_mask is None else channel_mask |
| return cls( |
| static_fields, |
| channel_mask, |
| grid_size=(13, 12), |
| embed_dim=32, |
| depth=2, |
| num_heads=4, |
| mlp_dim=64, |
| window_size=2, |
| ) |
|
|
| def reset_parameters(self) -> None: |
| for module in self.modules(): |
| if isinstance(module, nn.Linear): |
| nn.init.trunc_normal_(module.weight, std=0.02) |
| if module.bias is not None: |
| nn.init.zeros_(module.bias) |
| elif isinstance(module, (nn.Conv2d, nn.ConvTranspose2d)): |
| nn.init.xavier_uniform_(module.weight) |
| if module.bias is not None: |
| nn.init.zeros_(module.bias) |
| for block in self.layers: |
| nn.init.zeros_(block.adaln[-1].weight) |
| nn.init.zeros_(block.adaln[-1].bias) |
| nn.init.zeros_(self.norm_layer.scale_shift[-1].weight) |
| nn.init.zeros_(self.norm_layer.scale_shift[-1].bias) |
| for head in (self.pressure_head, self.surface_head.conv2, self.derived_head.conv2): |
| nn.init.trunc_normal_(head.weight, std=1e-3) |
|
|
| @staticmethod |
| def _time_embedding(value: torch.Tensor, periodic: bool) -> torch.Tensor: |
| frequency = torch.arange(64, device=value.device, dtype=value.dtype) |
| if periodic: |
| angles = 2 * math.pi * value.reshape(-1, 1) * frequency |
| else: |
| angles = value.reshape(-1, 1) / (10000 ** (frequency / 64)) |
| return torch.cat((angles.sin(), angles.cos()), dim=-1) |
|
|
| def forward(self, state: torch.Tensor, step: torch.Tensor, hour: torch.Tensor, doy: torch.Tensor) -> torch.Tensor: |
| expected = (2, 85, *self.grid_size) |
| if tuple(state.shape[1:]) != expected: |
| raise ValueError(f"state must have shape (B, {expected}), got {tuple(state.shape)}") |
| state = torch.nan_to_num(state) |
| state = state.clone() |
| state[:, :, DIAGNOSTIC_INDICES] = 0 |
| state = state * self.channel_mask |
| previous = state[:, -1] |
| batch = state.shape[0] |
|
|
| x = self.patch_embed(state.reshape(batch, 170, *self.grid_size)).flatten(2).transpose(1, 2) |
| x = self.patch_norm(x) |
| const = self.const_embed(self.static_fields[None].expand(batch, -1, -1, -1)).flatten(2).transpose(1, 2) |
| x = x + self.const_norm(const) |
| time_features = torch.cat( |
| (self._time_embedding(step, False), self._time_embedding(hour, True), self._time_embedding(doy, True)), dim=-1 |
| ) |
| condition = self.joint_embed_layer(time_features) |
| for layer in self.layers: |
| if self.training and self.activation_checkpointing: |
| x = activation_checkpoint(layer, x, condition, use_reentrant=False) |
| else: |
| x = layer(x, condition) |
| x = self.norm_layer(x, condition).transpose(1, 2).reshape(batch, -1, *self.token_grid) |
|
|
| pressure = self.pressure_head(x)[..., : self.grid_size[0], : self.grid_size[1]] |
| surface = self.surface_head(x, self.grid_size) |
| derived = self.derived_head(x, self.grid_size) |
| grouped = torch.cat((pressure, surface, derived), dim=1) |
| prediction = torch.empty_like(grouped) |
| prediction[:, self.scatter_idx] = grouped |
| return torch.stack((previous, prediction), dim=1) |
|
|
| @property |
| def trainable(self) -> bool: |
| return True |
|
|