File size: 2,553 Bytes
5c365c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
"""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()