Изграждане на персонализирано пакетно ансамблово прогнозиране на времето с NVIDIA Earth2Studio
В този урок изграждаме работен процес за ансамблово прогнозиране на времето с NVIDIA Earth2Studio. Инсталираме необходимите компоненти на Earth2Studio, като запазваме съществуващата CUDA-съвместима среда на PyTorch в Colab, зареждаме прогностичния модел FCN и извличаме начални атмосферни условия от GFS. След това реализираме персонализирана диагностика на вятърната енергия, която преобразува компонентите на вятъра на височина 10 метра в коефициенти на използване на мощността на турбината, заедно със система за мащабирани по променливи смущения, прилагаща физически подходящи амплитуди на шума към различни атмосферни променливи, като същевременно запазва непертурбиран контролен член. С помощта на нискоуровневия итератор, съпоставянето на координати, пакетната обработка и Zarr API на Earth2Studio изграждаме собствен конвейер за изпълнение на ансамбъла, записваме прогнозните и диагностичните полета в хранилище с координатна информация и проверяваме прогнозите спрямо анализите на GFS чрез RMSE с тегла по географска ширина, справедлив CRPS, ансамблов разброс и съотношения разброс-умение. Накрая визуализираме неопределеността на ансамбъла чрез пространствени карти, контури тип „спагети“ на геопотенциалната височина, точкови ветрилообразни диаграми, прогнози на коефициента на използване на вятърната мощност и криви на умението спрямо времето за прогнозиране.
import importlib.util, os, subprocess, sys
if importlib.util.find_spec("earth2studio") is None:
import numpy as _np, torch as _torch
cfile = os.path.join(os.getcwd(), "e2s_constraints.txt")
with open(cfile, "w") as f:
f.write(f"torch=={_torch.__version__.split('+')[0]}\n")
f.write(f"numpy=={_np.__version__}\n")
env = {**os.environ, "PIP_CONSTRAINT": cfile}
subprocess.check_call(
[sys.executable, "-m", "pip", "install", "-q",
"earth2studio[fcn,data,perturbation,statistics]"], env=env)
print("\n>>> Install done. If the imports below fail: Runtime > Restart session, re-run.\n")
os.environ.setdefault("EARTH2STUDIO_CACHE", "/content/e2s_cache")
os.makedirs("outputs", exist_ok=True)
from collections import OrderedDict
from datetime import datetime, timedelta, timezone
from tqdm.auto import tqdm
from earth2studio.data import GFS, fetch_data
from earth2studio.io import ZarrBackend
from earth2studio.models.batch import batch_coords, batch_func
from earth2studio.models.px import FCN
from earth2studio.statistics import rmse
from earth2studio.utils import handshake_coords, handshake_dim
from earth2studio.utils.coords import map_coords
from earth2studio.utils.time import to_time_array
from earth2studio.utils.type import CoordSystem
if DEVICE.type == "cpu":
print("!! No GPU detected — this will be very slow. Runtime > Change runtime type > T4 GPU")
NENSEMBLE = 8
BATCH_SIZE = 2
NSTEPS = 8
SAVE_VARS = ["t2m", "z500", "u10m", "v10m", "tcwv"]
VERIFY_VARS = ["t2m", "z500", "u10m"]
INIT = (datetime.now(timezone.utc) - timedelta(days=7)).replace()
INIT_STR = INIT.strftime("%Y-%m-%dT%H:%M:%S")
POI = ("New Delhi", 28.61, 77.21)
print(f"Initialization: {INIT_STR} | device: {DEVICE}")
Инсталираме Earth2Studio, като запазваме съществуващата CUDA-съвместима среда на PyTorch и NumPy в Colab чрез ограничения на пакетите. Конфигурираме кеша на модела, импортираме помощните програми за прогнозиране, данни, статистика, визуализация и управление на координати и откриваме наличното изчислително устройство. Също така задаваме размера на ансамбъла, размера на пакета, продължителността на прогнозата, запазваните променливи, променливите за проверка, началния момент и точката на интерес в Ню Делхи.
class WindPowerCF(torch.nn.Module):
"""Turbine capacity factor [0,1] from 10 m winds via power-law shear + power curve."""
def __init__(self, lat, lon, hub=100.0, alpha=0.143,
cut_in=3.0, rated=12.0, cut_out=25.0):
super().__init__()
self.lat, self.lon = lat, lon
self.hub, self.alpha = hub, alpha
self.cut_in, self.rated, self.cut_out = cut_in, rated, cut_out
def input_coords(self) -> CoordSystem:
return OrderedDict({
"batch": np.empty(0),
"variable": np.array(["u10m", "v10m"]),
"lat": self.lat,
"lon": self.lon,
})
@batch_coords()
def output_coords(self, input_coords: CoordSystem) -> CoordSystem:
target = self.input_coords()
for i, (key, _) in enumerate(target.items()):
if key != "batch":
handshake_dim(input_coords, key, i)
handshake_coords(input_coords, target, key)
oc = OrderedDict({
"batch": np.empty(0),
"variable": np.array(["wind_cf"]),
"lat": self.lat,
"lon": self.lon,
})
oc["batch"] = input_coords["batch"]
return oc
@batch_func()
def __call__(self, x: torch.Tensor, coords: CoordSystem):
oc = self.output_coords(coords)
u, v = x[..., 0:1, :, :], x[..., 1:2, :, :]
ws10 = torch.sqrt(u * u + v * v)
ws = ws10 * (self.hub / 10.0) ** self.alpha
ramp = (ws ** 3 - self.cut_in ** 3) / (self.rated ** 3 - self.cut_in ** 3)
cf = torch.zeros_like(ws)
cf = torch.where((ws >= self.cut_in) & (ws < self.rated), ramp.clamp(0, 1), cf)
cf = torch.where((ws >= self.rated) & (ws <= self.cut_out), torch.ones_like(cf), cf)
return cf, oc
class VariableScaledNoise:
"""Spatially correlated noise with per-variable amplitudes + control member."""
def __init__(self, amplitudes: dict, default: float = 0.0, control_member: bool = True):
self.amplitudes, self.default, self.control = amplitudes, default, control_member
try:
from earth2studio.perturbation import SphericalGaussian
self.sampler, self.kind = SphericalGaussian(noise_amplitude=1.0), "SphericalGaussian"
except Exception:
from earth2studio.perturbation import Brown
self.sampler, self.kind = Brown(noise_amplitude=1.0), "Brown"
def __call__(self, x: torch.Tensor, coords: CoordSystem):
noise, _ = self.sampler(torch.zeros_like(x), coords)
vax = list(coords).index("variable")
amps = torch.tensor([self.amplitudes.get(str(v), self.default)
for v in coords["variable"]], device=x.device, dtype=x.dtype)
shape = [1] * x.ndim; shape[vax] = amps.numel()
pert = noise * amps.reshape(shape)
if self.control and "ensemble" in coords:
eax = list(coords).index("ensemble")
mask = torch.tensor((np.asarray(coords["ensemble"]) != 0).astype(np.float32),
device=x.device, dtype=x.dtype)
mshape = [1] * x.ndim; mshape[eax] = mask.numel()
pert = pert * mask.reshape(mshape)
return x + pert, coords
Създаваме персонализиран диагностичен модел, който преобразува компонентите на вятъра на височина 10 метра във вятърна скорост на височината на главината и коефициент на използване на мощността на турбината. Проверяваме съвместимостта на координатите чрез помощните функции за съгласуване на Earth2Studio и поддържаме пакетирани входни данни с предоставените декоратори. Реализираме и специфични за променливите пространствени смущения, които запазват член нула като непертурбирана контролна прогноза.
def write_vars(io, x, coords, names):
"""Write selected channels of a (…, variable, lat, lon) tensor to the IO backend."""
vax = list(coords).index("variable")
sub = OrderedDict((k, v) for k, v in coords.items() if k != "variable")
for name in names:
hit = np.where(np.asarray(coords["variable"]) == name)[0]
if hit.size:
io.write(x.select(vax, int(hit[0])).cpu(), sub, name)
def run_ensemble(time, nsteps, nensemble, batch_size, prognostic, diagnostic,
perturbation, data, io, save_vars, device):
time = to_time_array(time)
ic = prognostic.input_coords()
x0, c0 = fetch_data(source=data, time=time, lead_time=ic["lead_time"],
variable=ic["variable"], device=device)
print(f"Initial condition tensor: {tuple(x0.shape)} dims={list(c0)}")
oc = prognostic.output_coords(ic)
dt = oc["lead_time"]
prog_vars = [v for v in save_vars if v in set(map(str, oc["variable"]))]
total = OrderedDict({
"ensemble": np.arange(nensemble),
"time": time,
"lead_time": np.asarray([dt * i for i in range(nsteps + 1)]).flatten(),
"lat": oc["lat"],
"lon": oc["lon"],
})
io.add_array(total, prog_vars + ["wind_cf"])
dx_target = OrderedDict((k, v) for k, v in diagnostic.input_coords().items() if k != "batch")
nbatch = int(np.ceil(nensemble / batch_size))
with torch.inference_mode():
for b in tqdm(range(nbatch), desc="ensemble batches"):
lo = b * batch_size
n = min(batch_size, nensemble - lo)
x = x0.unsqueeze(0).repeat(n, *([1] * x0.ndim))
coords = OrderedDict({"ensemble": np.arange(lo, lo + n), **c0})
x, coords = perturbation(x, coords)
x, coords = map_coords(x, coords, ic)
for step, (xs, cs) in enumerate(prognostic.create_iterator(x, coords)):
write_vars(io, xs, cs, prog_vars)
xw, cw = map_coords(xs, cs, dx_target)
xw, cw = diagnostic(xw, cw)
write_vars(io, xw, cw, ["wind_cf"])
if step >= nsteps:
break
torch.cuda.empty_cache() if device.type == "cuda" else None
return io
model = FCN.load_model(FCN.load_default_package()).to(DEVICE)
grid = model.output_coords(model.input_coords())
LAT, LON = grid["lat"], grid["lon"]
diagnostic = WindPowerCF(LAT, LON).to(DEVICE)
pert = VariableScaledNoise(
amplitudes={"t2m": 0.20, "t850": 0.20, "z500": 40.0, "z850": 25.0,
"u10m": 0.25, "v10m": 0.25, "u500": 0.40, "v500": 0.40, "tcwv": 0.30},
default=0.0, control_member=True)
print(f"Perturbation sampler: {pert.kind}")
io = ZarrBackend(file_name="outputs/e2s_ensemble.zarr",
chunks={"ensemble": 1, "time": 1, "lead_time": 1},
backend_kwargs={"overwrite": True})
io = run_ensemble([INIT_STR], NSTEPS, NENSEMBLE, BATCH_SIZE,
model, diagnostic, pert, GFS(), io, SAVE_VARS, DEVICE)
print(io.root.tree())
Дефинираме помощни функции, които избират атмосферни канали и ги записват в Zarr хранилище с координатна информация. Изграждаме персонализиран пакетен цикъл за ансамбъла, който извлича начални условия от GFS, прилага смущения към членовете на ансамбъла, подравнява координатите, итерира модела FCN и свързва диагностиката на вятърната енергия. След това зареждаме модела, инициализираме диагностичните компоненти и компонентите за смущения, изпълняваме прогнозата и проверяваме получената структура на Zarr.
leads = np.asarray(io["lead_time"][:]).astype("timedelta64[ns]")
lead_h = leads.astype("timedelta64[h]").astype(int)
valid = to_time_array([INIT_STR])[0] + leads
truth, tc = fetch_data(source=GFS(), time=valid,
lead_time=np.array([np.timedelta64(0, "h")]),
variable=np.array(VERIFY_VARS), device="cpu")
truth = truth[:, 0]
w = torch.cos(torch.deg2rad(torch.as_tensor(np.asarray(LAT), dtype=torch.float32)))
w2d = w[:, None].expand(len(LAT), len(LON)).contiguous()
mcoords = OrderedDict({"lead_time": leads, "lat": np.asarray(LAT), "lon": np.asarray(LON)})
def fair_crps(ens, obs, weights):
"""Fair (unbiased) CRPS, lat-weighted. ens: (M, lat, lon), obs: (lat, lon)."""
M = ens.shape[0]
wn = weights / weights.sum()
skill = ((ens - obs).abs() * wn).sum(dim=(-2, -1)).mean()
spread = torch.zeros((), dtype=ens.dtype)
for i in range(M):
spread = spread + ((ens[i] - ens).abs() * wn).sum(dim=(-2, -1)).sum()
return (skill - spread / (2 * M * (M - 1))).item()
scores = {}
for k, var in enumerate(VERIFY_VARS):
fc = torch.as_tensor(np.asarray(io[var][:]))[:, 0].float()
ob = truth[:, k].float()
mean = fc.mean(0)
try:
metric = rmse(reduction_dimensions=["lat", "lon"], weights=w2d)
r, _ = metric(mean, mcoords, ob, mcoords)
r = r.numpy()
except Exception as e:
print(f"(built-in rmse unavailable: {e})")
wn = (w2d / w2d.sum())
r = torch.sqrt((((mean - ob) ** 2) * wn).sum(dim=(-2, -1))).numpy()
wn = w2d / w2d.sum()
spread = torch.sqrt((fc.var(0, unbiased=True) * wn).sum(dim=(-2, -1))).numpy()
crps = np.array([fair_crps(fc[:, t], ob[t], w2d) for t in range(fc.shape[1])])
scores[var] = dict(rmse=r, spread=spread, crps=crps, fc=fc, obs=ob, mean=mean)
print(f"\n=== {var} ===")
print(f"{'lead[h]':>8}{'RMSE':>12}{'spread':>12}{'ratio':>9}{'CRPS':>12}")
for t in range(len(lead_h)):
ratio = spread[t] / r[t] if r[t] > 0 else np.nan
print(f"{lead_h[t]:>8}{r[t]:>12.3f}{spread[t]:>12.3f}{ratio:>9.2f}{crps[t]:>12.3f}")
Извличаме анализите на GFS за всеки валиден момент на прогнозата и ги използваме като референтни данни за проверка. Изчисляваме RMSE с тегла по географска ширина, ансамблов разброс, справедлив CRPS и съотношения разброс-грешка за температурата, геопотенциалната височина и вятърните променливи. Съхраняваме прогнозните полета и оценъчните показатели в структурирана речникоподобна колекция и отпечатваме обобщения на умението за различните срокове на прогнозата за всяка променлива.
lat_np, lon_np = np.asarray(LAT), np.asarray(LON)
ilat = int(np.argmin(np.abs(lat_np - POI[1])))
ilon = int(np.argmin(np.abs(lon_np - (POI[2] % 360))))
last = -1
d = scores["t2m"]
fields = [(d["mean"][last].numpy() - 273.15, "ensemble mean t2m [C]", "RdBu_r", None),
(d["fc"][:, last].std(0).numpy(), "ensemble spread [K]", "magma", None),
(d["obs"][last].numpy() - 273.15, "GFS analysis [C]", "RdBu_r", None),
((d["mean"][last] - d["obs"][last]).numpy(), "mean error [K]", "coolwarm", 5)]
fig, axs = plt.subplots(2, 2, figsize=(15, 7), constrained_layout=True)
for ax, (f, title, cmap, lim) in zip(axs.ravel(), fields):
kw = dict(vmin=-lim, vmax=lim) if lim else {}
im = ax.pcolormesh(lon_np, lat_np, f, cmap=cmap, shading="auto", **kw)
ax.set_title(f"{title} — +{lead_h[last]} h"); plt.colorbar(im, ax=ax, shrink=0.85)
plt.show()
z = scores["z500"]["fc"][:, last].numpy() / 9.81
la = (lat_np > 25) & (lat_np < 75)
lo = (lon_np > 280) | (lon_np < 40)
lon_shift = np.where(lon_np > 180, lon_np - 360, lon_np)
order = np.argsort(lon_shift[lo])
plt.figure(figsize=(11, 5))
for m in range(z.shape[0]):
sub = z[m][np.ix_(la, lo)][:, order]
plt.contour(lon_shift[lo][order], lat_np[la], sub, levels=[5520],
colors=["k" if m == 0 else "C0"], linewidths=[2.0 if m == 0 else 0.8])
zo = scores["z500"]["obs"][last].numpy() / 9.81
plt.contour(lon_shift[lo][order], lat_np[la], zo[np.ix_(la, lo)][:, order],
levels=[5520], colors="crimson", linewidths=2.5)
plt.title(f"z500 5520 m spaghetti at +{lead_h[last]} h "
f"(black=control, blue=members, red=GFS analysis)")
plt.xlabel("lon"); plt.ylabel("lat"); plt.show()
t2m_pt = scores["t2m"]["fc"][:, :, ilat, ilon].numpy() - 273.15
obs_pt = scores["t2m"]["obs"][:, ilat, ilon].numpy() - 273.15
cf_pt = np.asarray(io["wind_cf"][:])[:, 0, :, ilat, ilon]
fig, (a1, a2) = plt.subplots(1, 2, figsize=(14, 4))
a1.fill_between(lead_h, t2m_pt.min(0), t2m_pt.max(0), alpha=0.25, label="member range")
a1.plot(lead_h, t2m_pt.mean(0), "o-", label="ensemble mean")
a1.plot(lead_h, t2m_pt[0], "k--", label="control")
a1.plot(lead_h, obs_pt, "r^-", label="GFS analysis")
a1.set_title(f"2 m temperature — {POI[0]}"); a1.set_xlabel("lead [h]"); a1.set_ylabel("C")
a1.legend(); a1.grid(alpha=.3)
a2.fill_between(lead_h, cf_pt.min(0), cf_pt.max(0), alpha=0.25, color="seagreen")
a2.plot(lead_h, cf_pt.mean(0), "o-", color="seagreen")
a2.set_title(f"wind capacity factor (custom diagnostic) — {POI[0]}")
a2.set_xlabel("lead [h]"); a2.set_ylim(0, 1); a2.grid(alpha=.3)
plt.tight_layout(); plt.show()
fig, axs = plt.subplots(1, len(VERIFY_VARS), figsize=(5 * len(VERIFY_VARS), 3.6))
for ax, var in zip(np.atleast_1d(axs), VERIFY_VARS):
s = scores[var]
ax.plot(lead_h, s["rmse"], "o-", label="RMSE (ens. mean)")
ax.plot(lead_h, s["spread"], "s--", label="spread")
ax.plot(lead_h, s["crps"], "^:", label="fair CRPS")
ax.set_title(var); ax.set_xlabel("lead [h]"); ax.grid(alpha=.3); ax.legend(fontsize=8)
plt.tight_layout(); plt.show()
import xarray as xr
ds = xr.open_zarr("outputs/e2s_ensemble.zarr")
print(ds)
Визуализираме поведението на ансамбъла чрез карти на средната температура, разброса, анализа и грешката в последния срок на прогнозата. Генерираме контури тип „спагети“ на геопотенциалната височина, ветрилообразна диаграма на температурата в Ню Делхи, прогноза на коефициента на използване на вятърната мощност и криви на умението спрямо времето за прогнозиране. Накрая отваряме изхода от Zarr с Xarray, за да можем да инспектираме, анализираме или експортираме пълния набор от ансамблови данни.
В заключение изградихме гъвкав и разширяем работен процес с Earth2Studio, който надхвърля изпълнението на предварително дефинирана функция за ансамбъл. Получихме пряк контрол върху смущенията на началните условия, пакетната обработка на членовете, итерацията на модела, свързването на диагностиката, подравняването на координатите, съхраняването на данните, проверката и визуализацията в рамките на една среда Colab. Демонстрирахме също как физически мащабираните смущения и непертурбираният контролен член ни помагат да интерпретираме ансамбловия разброс. В същото време диагностиките RMSE, справедлив CRPS и разброс-умение ни позволяват да оценим точността и калибрирането на прогнозата за различните срокове. Полученият набор от данни в Zarr запазва пълната структура на ансамбъла и остава достъпен чрез Xarray за допълнителен анализ или конвертиране. Тъй като работният процес следва интерфейсите на компонентите на Earth2Studio, можем да го разширим, като заменим прогностичния модел, променим източника на атмосферни данни, добавим нови диагностики, увеличим размера на ансамбъла или възприемем асинхронно съхранение, без да преработваме целия конвейер за прогнозиране.
Вижте ПЪЛНИТЕ КОДОВЕ тук. Също така можете да ни последвате в Twitter и не забравяйте да се присъедините към нашия ML SubReddit с над 150 хил. членове и да се абонирате за нашия информационен бюлетин. Почакайте! В Telegram ли сте? Вече можете да се присъедините към нас и в Telegram.
Имате нужда от партньор за популяризиране на вашето GitHub хранилище, страница в Hugging Face, продуктова премиера, уебинар и т.н.? Свържете се с нас
Преведено автоматично от английски. Оригиналната статия е на връзката по-долу.