Sakhanda Wire
NVDA MSFT GOOGL META AMZN
← К новостям

Создание пользовательского пакетного ансамблевого прогнозирования погоды с NVIDIA Earth2Studio

В этом руководстве мы создаём рабочий процесс ансамблевого прогнозирования погоды с помощью NVIDIA Earth2Studio. Мы устанавливаем необходимые компоненты Earth2Studio, сохраняя существующую в Colab среду PyTorch с поддержкой CUDA, загружаем прогностическую модель FCN и получаем начальные атмосферные условия из GFS. Затем мы реализуем пользовательскую диагностику ветроэнергетики, преобразующую компоненты ветра на высоте 10 метров в коэффициенты использования мощности турбины, а также систему возмущений с масштабированием по переменным, применяющую физически корректные амплитуды шума к различным атмосферным переменным при сохранении невозмущённого контрольного члена. Используя низкоуровневый итератор Earth2Studio, API сопоставления координат, пакетной обработки и Zarr, мы создаём собственный конвейер выполнения ансамбля, записываем поля прогнозов и диагностики в хранилище с учётом координат и проверяем прогнозы по анализам GFS с помощью среднеквадратичной ошибки с весами по широте, честного 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, сохраняя существующие в Colab окружения PyTorch с поддержкой CUDA и NumPy с помощью ограничений пакетов. Мы настраиваем кэш модели, импортируем инструменты прогнозирования, работы с данными, статистики, построения графиков и управления координатами, а также определяем доступное вычислительное устройство. Кроме того, мы задаём размер ансамбля, размер пакета, продолжительность прогноза, сохраняемые переменные, переменные для проверки, время инициализации и точку интереса в Нью-Дели.

Копировать кодСкопированоИспользовать другой браузер
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 для каждого времени, соответствующего прогнозу, и используем их как эталонные данные для проверки. Мы рассчитываем среднеквадратичную ошибку с весами по широте, разброс ансамбля, честный 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 и не забудьте присоединиться к нашему сабреддиту о машинном обучении с более чем 150 тыс. участников и подписаться на нашу рассылку. Постойте! Вы есть в Telegram? теперь вы также можете присоединиться к нам в Telegram.

Хотите сотрудничать с нами в продвижении вашего репозитория GitHub, страницы Hugging Face, релиза продукта, вебинара и т. д.? Свяжитесь с нами

Переведено автоматически с английского. Оригинал статьи — по ссылке ниже.

Впервые опубликовано изданием MarkTechPost

Читать оригинал на MarkTechPost ↗

Текст и изображения принадлежат MarkTechPost и приводятся здесь с указанием авторства и ссылкой на оригинальную публикацию.

← К новостям

Ещё новости

Все последние новости