yzt15806542928's picture
Upload folder using huggingface_hub
5c365c5 verified
Raw
History Blame Contribute Delete
4.16 kB
"""ACE 40-channel input, 44-channel output and autoregressive rollout."""
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Callable
import torch
from torch import nn
from ACE.model.variables import (
FORCING_CHANNELS,
INPUT_CHANNELS,
OUTPUT_CHANNELS,
PROGNOSTIC_CHANNELS,
split_output,
validate_channels,
)
from ACE.model.sfno import SFNOAdapter, SFNOConfig
@dataclass
class ACEModelConfig:
nlat: int = 180
nlon: int = 360
input_channels: int = len(INPUT_CHANNELS)
output_channels: int = len(OUTPUT_CHANNELS)
prognostic_channels: int = len(PROGNOSTIC_CHANNELS)
forcing_channels: int = len(FORCING_CHANNELS)
embed_dim: int = 256
num_layers: int = 8
filter_type: str = "linear"
operator_type: str = "dhconv"
scale_factor: int = 1
spectral_layers: int = 3
grid: str = "legendre-gauss"
grid_internal: str = "legendre-gauss"
mlp_ratio: float = 2.0
fallback: bool = False
def to_dict(self) -> dict:
return asdict(self)
class ACEModel(nn.Module):
def __init__(self, config: ACEModelConfig | None = None) -> None:
super().__init__()
self.config = config or ACEModelConfig()
if self.config.input_channels != len(INPUT_CHANNELS) or self.config.output_channels != len(OUTPUT_CHANNELS):
raise ValueError("ACE channel contract must remain 40 input and 44 output channels")
sfno_config = SFNOConfig(
nlat=self.config.nlat,
nlon=self.config.nlon,
in_channels=self.config.input_channels,
out_channels=self.config.output_channels,
embed_dim=self.config.embed_dim,
num_layers=self.config.num_layers,
filter_type=self.config.filter_type,
operator_type=self.config.operator_type,
scale_factor=self.config.scale_factor,
spectral_layers=self.config.spectral_layers,
grid=self.config.grid,
grid_internal=self.config.grid_internal,
mlp_ratio=self.config.mlp_ratio,
fallback=self.config.fallback,
)
self.sfno = SFNOAdapter(sfno_config)
def forward(self, x: torch.Tensor) -> torch.Tensor:
validate_channels(x, self.config.input_channels, "ACE input")
return self.sfno(x)
def step(self, prognostic: torch.Tensor, forcing: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
validate_channels(prognostic, self.config.prognostic_channels, "prognostic state")
validate_channels(forcing, self.config.forcing_channels, "forcing")
predicted = self.forward(torch.cat([prognostic, forcing], dim=1))
return split_output(predicted)
@torch.no_grad()
def rollout(
self,
initial_prognostic: torch.Tensor,
forcings: torch.Tensor | Callable[[int, torch.Tensor], torch.Tensor],
steps: int | None = None,
) -> torch.Tensor:
"""Return predictions with shape ``[B,T,44,H,W]``.
``forcings`` is either `[B,T,6,H,W]` or a callable receiving
`(step, current_prognostic)` and returning `[B,6,H,W]`.
"""
validate_channels(initial_prognostic, self.config.prognostic_channels, "initial prognostic")
if callable(forcings):
if steps is None or steps < 1:
raise ValueError("steps is required for callable forcing")
forcing_steps = steps
else:
if forcings.ndim != 5 or forcings.shape[2] != self.config.forcing_channels:
raise ValueError("tensor forcings must have shape [B,T,6,H,W]")
forcing_steps = forcings.shape[1] if steps is None else min(steps, forcings.shape[1])
state = initial_prognostic
outputs = []
for step in range(forcing_steps):
forcing = forcings[:, step] if not callable(forcings) else forcings(step, state)
predicted_state, diagnostics = self.step(state, forcing)
outputs.append(torch.cat([predicted_state, diagnostics], dim=1))
state = predicted_state
return torch.stack(outputs, dim=1)