Ai2_Climate_Emulator / model /preprocess.py
yzt15806542928's picture
Upload folder using huggingface_hub
5c365c5 verified
Raw
History Blame Contribute Delete
2.55 kB
"""FV3GFS preprocessing interfaces used by the ACE model package."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable
import numpy as np
@dataclass(frozen=True)
class VerticalCoordinate:
a_pa: np.ndarray
b: np.ndarray
source_indices: np.ndarray
@classmethod
def ace_default(cls) -> "VerticalCoordinate":
return cls(
np.asarray([64.247, 5167.14603, 12905.42546, 13982.4677, 12165.28766, 8910.07678, 4955.72632, 2155.78385, 0.0]),
np.asarray([0.0, 0.0, 0.01755, 0.11746, 0.2896, 0.49806, 0.72625, 0.88192, 1.0]),
np.asarray([0, 18, 26, 31, 36, 41, 47, 53, 63]),
)
def interface_pressure(self, surface_pressure: np.ndarray) -> np.ndarray:
ps = np.asarray(surface_pressure)
if ps.ndim < 2:
raise ValueError("surface_pressure must end in [lat, lon]")
return np.einsum("k,...ij->...kij", self.b, ps) + self.a_pa.reshape((1,) * (ps.ndim - 2) + (9, 1, 1))
def conservative_vertical_coarsen(values: np.ndarray, layer_thickness: np.ndarray, coordinate: VerticalCoordinate | None = None, layer_axis: int = -3) -> tuple[np.ndarray, np.ndarray]:
coord = coordinate or VerticalCoordinate.ace_default()
x = np.moveaxis(np.asarray(values, dtype=np.float64), layer_axis, -3)
dp = np.moveaxis(np.asarray(layer_thickness, dtype=np.float64), layer_axis, -3)
if x.shape[-3] != 63 or dp.shape != x.shape:
raise ValueError(f"expected values/dp shape [...,63,H,W], got {x.shape} and {dp.shape}")
means, thickness = [], []
for start, end in zip(coord.source_indices[:-1], coord.source_indices[1:]):
chunk_dp, chunk_x = dp[..., start:end, :, :], x[..., start:end, :, :]
dpk = chunk_dp.sum(axis=-3)
means.append((chunk_x * chunk_dp).sum(axis=-3) / np.maximum(dpk, 1e-12))
thickness.append(dpk)
return np.moveaxis(np.stack(means, axis=-3), -3, layer_axis), np.moveaxis(np.stack(thickness, axis=-3), -3, layer_axis)
def conservative_regrid(fields: np.ndarray, regrid_fn: Callable[[np.ndarray], np.ndarray] | None = None) -> np.ndarray:
if regrid_fn is None:
raise NotImplementedError("MISSING: NOAA fregrid or another conservative cubed-sphere adapter")
return np.asarray(regrid_fn(np.asarray(fields)))
def spherical_harmonic_roundtrip(fields: np.ndarray, filter_fn: Callable[[np.ndarray], np.ndarray] | None = None) -> np.ndarray:
return np.asarray(fields if filter_fn is None else filter_fn(np.asarray(fields))).copy()