Download model/ace.py from OneScience-Group/Ai2_Climate_Emulator: direct link, hf CLI and curl.
- Browser
- Download file 4.16 kB
-
https://huggingface.co/OneScience-Group/Ai2_Climate_Emulator/resolve/main/model/ace.py
- Command line
-
hf download hf://OneScience-Group/Ai2_Climate_Emulator/model/ace.py
-
curl -L -o ace.py https://huggingface.co/OneScience-Group/Ai2_Climate_Emulator/resolve/main/model/ace.py
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 | |
| 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) | |
| 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) | |