Publish Trackformer 1.2 announcement assets and matched DeepMind benchmark protocol
Browse files- RELEASE_NOTES_TRACKFORMER_1_2.md +29 -19
- models/trackformer_1_2_field/README.md +35 -2
- models/trackformer_1_2_field/plot_pressure.py +91 -0
- models/trackformer_1_2_field/predict.py +67 -1
- release_tools/sync_public_model_cards.py +4 -0
- release_tools/test_pressure_field_export.py +111 -0
- release_tools/verify_pressure_field_export.py +120 -0
RELEASE_NOTES_TRACKFORMER_1_2.md
CHANGED
|
@@ -1,31 +1,41 @@
|
|
| 1 |
-
# Trackformer 1.2 — research
|
| 2 |
|
| 3 |
-
**[Download weights + code](https://github.com/yu314-coder/typhoon-predict/releases/download/trackformer-1.2/
|
| 4 |
|
| 5 |
-
|
| 6 |
|
| 7 |
-
|
| 8 |
|
| 9 |
-
The
|
| 10 |
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
| 12 |
|
| 13 |
-
|
|
|
|
|
|
|
|
|
|
| 14 |
|
| 15 |
-
|
| 16 |
|
| 17 |
-
|
| 18 |
|
| 19 |
-
|
| 20 |
|
| 21 |
-
|
| 22 |
-
| --- | ---: | ---: |
|
| 23 |
-
| Mean track error | 798.4 km | 471.2 km |
|
| 24 |
-
| +120 h track error | 1,646.4 km | 1,031.6 km |
|
| 25 |
-
| Direction error | 51.58° | 34.96° |
|
| 26 |
-
| Centred shape similarity | 0.7544 | 0.8837 |
|
| 27 |
-
| Geographic path similarity | 0.5345 | 0.6560 |
|
| 28 |
|
| 29 |
-
|
| 30 |
|
| 31 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Trackformer 1.2 — pressure-field research model
|
| 2 |
|
| 3 |
+
**[Download weights + complete inference code](https://github.com/yu314-coder/typhoon-predict/releases/download/trackformer-1.2/trackformer_1_2_field_pressure_export_v2.tar.gz)** · **[Hugging Face](https://huggingface.co/euler314/typhoon-predict)** · **[Technical paper](https://github.com/yu314-coder/typhoon-predict/blob/main/paper/trackformer.pdf)**
|
| 4 |
|
| 5 |
+
Trackformer 1.2 forecasts Western Pacific storm tracks, central pressure and evolving sea-level-pressure fields through +120 hours. A multiscale environmental-attention network conditions the moving pressure core; track and central pressure are read from that evolving field.
|
| 6 |
|
| 7 |
+
## Direct detailed pressure-map output
|
| 8 |
|
| 9 |
+
The current package exports all twenty six-hour leads with:
|
| 10 |
|
| 11 |
+
- Whole-WP basin pressure, the original fixed regional composite, and the actual **65×65 moving-core pressure field in physical hPa**.
|
| 12 |
+
- Basin, regional and per-lead core latitude/longitude grids, coverage masks, issue time and exact valid times.
|
| 13 |
+
- Original track/central-pressure outputs and an explicit **one-member** count.
|
| 14 |
+
- Frozen weight/checkpoint identity and an optional PNG renderer: blue low pressure, red high pressure, labelled isobars.
|
| 15 |
|
| 16 |
+
```bash
|
| 17 |
+
python models/trackformer_1_2_field/predict.py causal_issue_packet.npz forecast.npz --device cpu --pressure-map pressure_120h.png --map-lead 120
|
| 18 |
+
python models/trackformer_1_2_field/plot_pressure.py forecast.npz pressure_24h.png --lead 24 --interval 2
|
| 19 |
+
```
|
| 20 |
|
| 21 |
+
NumPy and PyTorch are required for inference; Matplotlib is optional for images. Use `--device mps` or `--device cuda` on a compatible system. Prepare the nine-analysis causal input packet using the [published schema](https://github.com/yu314-coder/typhoon-predict/blob/main/models/trackformer_1_2_field/README.md); this package does not fetch live weather automatically.
|
| 22 |
|
| 23 |
+
The core's 20-km spacing is a learned computational reconstruction, not new native observations. Invalid coverage remains masked. Fields are not shifted onto a route, and central pressure is not inserted as a display vortex. Ensemble members require geographic registration before physical-field averaging; this command is not the separate 50-member benchmark policy.
|
| 24 |
|
| 25 |
+
**The learned modules, inference weights and forecast equations are unchanged.** The added export captures fields already produced by the released model. The original September 29 package remains available separately; use the pressure-export-v2 archive for the complete default output.
|
| 26 |
|
| 27 |
+
Inference weight SHA-256: `db49f36e85a3766defc4c172746897a1f783705d1ce8e6f9dfb8e87ae1d902cb`.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 28 |
|
| 29 |
+
## Matched development results
|
| 30 |
|
| 31 |
+
| Measure | 1.1 | 1.2 mean of 50 | Matched coverage |
|
| 32 |
+
| --- | ---: | ---: | --- |
|
| 33 |
+
| Mean track error | 798.4 km | 471.2 km | 1,473 daily starts / 270 storms |
|
| 34 |
+
| Direction error | 51.58° | 34.96° | Same daily starts |
|
| 35 |
+
| Central-pressure MAE against JMA | 13.53 hPa | 12.84 hPa | 134 common starts / 40 storms |
|
| 36 |
+
|
| 37 |
+
Storms receive equal weight after valid leads and daily starts are averaged. The mean pressure reduction is small and its paired whole-storm uncertainty includes no improvement. These repeatedly inspected results are development evidence, not a certified untouched holdout. [Verified common-support metrics](https://github.com/yu314-coder/typhoon-predict/blob/main/evaluation/released_daily/released_daily_benchmark.json).
|
| 38 |
+
|
| 39 |
+
The [model announcement](https://github.com/yu314-coder/typhoon-predict) includes the selected Mangkhut pressure-map animation and architecture. Selected examples are not representative skill. Auxiliary wind and pressure-derived radius diagnostics remain unvalidated; there is no native wind-radius forecast head in 1.2.
|
| 40 |
+
|
| 41 |
+
This is a research model, not an operational warning service or a safety-critical forecast. Trackformer 1.1 remains a separate release. No historical archive, media, benchmark predictions or training run is modified by this exporter update.
|
models/trackformer_1_2_field/README.md
CHANGED
|
@@ -26,7 +26,40 @@ The `.npz` packet must contain exactly these arrays, with no object/pickle conte
|
|
| 26 |
| `issue_time_ns` | scalar | Forecast issue timestamp in Unix nanoseconds. |
|
| 27 |
| `history_time_ns` | `(9,)` | Nine consecutive six-hour timestamps ending at issue time. |
|
| 28 |
|
| 29 |
-
Output `.npz` contains
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
|
| 31 |
It also exposes the existing `maximum_wind_auxiliary_kt` scalar and experimental
|
| 32 |
pressure-derived `pressure_wind_estimate_kt`, `rmw_estimate_km` and
|
|
@@ -47,4 +80,4 @@ The published 50-member mean is a separate evaluation policy: 50 deterministic s
|
|
| 47 |
|
| 48 |
**Matched development evaluation:** the released 1.1 and 1.2 route comparison uses the same frozen **1,473 daily starts / 270 equal-weight storms**, +6 to +120 h. The completed native-pressure comparison has **134 common starts / 40 storms**; the other 1,339 starts lack valid frozen 1.1 intensity inputs and are not zero-scored. The Site defaults to central-pressure MAE: **13.53 / 12.84 hPa** against JMA and **12.84 / 12.55 hPa** against USA (1.1 / 1.2), slight mean improvements with paired whole-storm uncertainty including zero. The optional curve-similarity diagnostic is **(1 + centred cosine) / 2 at exact common valid times**, without time shifting or warping. JMA similarity is **0.7074 / 0.7118**; USA is **0.7237 / 0.6637** on 131 eligible non-flat curves. Removing level/amplitude makes this a shape diagnostic, not proof of correct pressure levels. Actual unshifted hPa timelines, agency masks and uncertainty remain separate. See the [shared snapshot](../../evaluation/released_daily/released_daily_benchmark.json), [verification receipt](../../evaluation/released_daily/released_daily_verification.json), [pressure protocol](../../docs/intensity_benchmark.md) and [public benchmark API](https://trackformer-weatherlab.rudin-euler-8253.chatgpt.site/api/benchmarks/released). A new matched WeatherNext Cyclones Mini run is in progress on these exact daily starts; results remain pending. See the [frozen protocol](../../docs/deepmind_daily_benchmark.md); no old scores are transferred.
|
| 49 |
|
| 50 |
-
**Pressure-display coverage:** `regional_mslp_hpa`
|
|
|
|
| 26 |
| `issue_time_ns` | scalar | Forecast issue timestamp in Unix nanoseconds. |
|
| 27 |
| `history_time_ns` | `(9,)` | Nine consecutive six-hour timestamps ending at issue time. |
|
| 28 |
|
| 29 |
+
Output `.npz` contains twenty exact +6, +12, …, +120-hour forecasts. Pressure arrays are physical hPa. The moving core is exported by default, not just used internally for a central-pressure number:
|
| 30 |
+
|
| 31 |
+
| Output | Shape | Meaning |
|
| 32 |
+
| --- | --- | --- |
|
| 33 |
+
| `basin_mslp_hpa` | `(20,25,33)` | Unchanged coarse Western Pacific pressure fields. |
|
| 34 |
+
| `basin_latitude_deg`, `basin_longitude_deg` | `(25,33)` each | Basin geography from the actual input grid. |
|
| 35 |
+
| `regional_mslp_hpa` | `(20,121,121)` | Unchanged fixed issue-relative pressure composites. |
|
| 36 |
+
| `regional_latitude_deg`, `regional_longitude_deg` | `(121,121)` each | Fixed regional geography. |
|
| 37 |
+
| `regional_valid` | `(20,121,121)` | Original model domain-support masks; not a guarantee of moving-core detail throughout this fixed patch. |
|
| 38 |
+
| `core_mslp_hpa` | `(20,65,65)` | Actual learned moving-core fields, converted from model normalization to hPa. |
|
| 39 |
+
| `core_latitude_deg`, `core_longitude_deg` | `(20,65,65)` each | Original moving-frame geographic coordinates at every forecast lead. |
|
| 40 |
+
| `core_valid` | `(20,65,65)` | Original moving-core coverage masks. Mask false cells even when their stored pressure is finite. |
|
| 41 |
+
| `track_lat_lon`, `central_pressure_hpa`, `track_valid` | `(20,2)`, `(20,)`, `(20,)` | Unchanged associated track and bilinear core-pressure readout, with domain support. |
|
| 42 |
+
| `issue_time_ns`, `valid_time_ns` | scalar, `(20,)` | UTC Unix nanoseconds for the issue and each exact forecast lead. |
|
| 43 |
+
| `issue_center_lat_lon`, `member_count` | `(2,)`, scalar | Initialization location and the honest count of one clean member. |
|
| 44 |
+
| `pressure_export_json` | string scalar | Export schema, units, frozen checkpoint/weight hashes, reconstruction and coverage policy. |
|
| 45 |
+
|
| 46 |
+
No field is shifted onto the forecast or observed route, and no central-pressure scalar is inserted into a display. The moving-core information spacing is **20 km computational reconstruction**, not new native high-resolution observations. Coordinates and masks must travel with the fields; an invalid cell is not a zero-pressure forecast. `central_pressure_hpa` is sampled from this same moving field at the associated centre, so it need not equal a discrete cell minimum or a labelled contour level.
|
| 47 |
+
|
| 48 |
+
### Render a pressure map directly
|
| 49 |
+
|
| 50 |
+
Install Matplotlib in your inference environment only if you want PNG rendering. NumPy and PyTorch suffice for the numerical export.
|
| 51 |
+
|
| 52 |
+
```bash
|
| 53 |
+
python models/trackformer_1_2_field/predict.py causal_issue_packet.npz forecast.npz --device cpu --pressure-map pressure_120h.png --map-lead 120 --isobar-interval 4
|
| 54 |
+
```
|
| 55 |
+
|
| 56 |
+
The optional image shows the whole Western Pacific basin and the actual moving core side by side, on one pressure colour scale: **blue low / red high**. It keeps original geography and masks unsupported core cells. To draw another saved lead without repeating inference:
|
| 57 |
+
|
| 58 |
+
```bash
|
| 59 |
+
python models/trackformer_1_2_field/plot_pressure.py forecast.npz pressure_24h.png --lead 24 --interval 2
|
| 60 |
+
```
|
| 61 |
+
|
| 62 |
+
All original track, central-pressure, basin/regional and auxiliary-wind outputs retain their existing definitions. The exporter verifies the downloaded weight SHA-256 against `manifest.json`; the learned model modules, weights and forward equations are unchanged.
|
| 63 |
|
| 64 |
It also exposes the existing `maximum_wind_auxiliary_kt` scalar and experimental
|
| 65 |
pressure-derived `pressure_wind_estimate_kt`, `rmw_estimate_km` and
|
|
|
|
| 80 |
|
| 81 |
**Matched development evaluation:** the released 1.1 and 1.2 route comparison uses the same frozen **1,473 daily starts / 270 equal-weight storms**, +6 to +120 h. The completed native-pressure comparison has **134 common starts / 40 storms**; the other 1,339 starts lack valid frozen 1.1 intensity inputs and are not zero-scored. The Site defaults to central-pressure MAE: **13.53 / 12.84 hPa** against JMA and **12.84 / 12.55 hPa** against USA (1.1 / 1.2), slight mean improvements with paired whole-storm uncertainty including zero. The optional curve-similarity diagnostic is **(1 + centred cosine) / 2 at exact common valid times**, without time shifting or warping. JMA similarity is **0.7074 / 0.7118**; USA is **0.7237 / 0.6637** on 131 eligible non-flat curves. Removing level/amplitude makes this a shape diagnostic, not proof of correct pressure levels. Actual unshifted hPa timelines, agency masks and uncertainty remain separate. See the [shared snapshot](../../evaluation/released_daily/released_daily_benchmark.json), [verification receipt](../../evaluation/released_daily/released_daily_verification.json), [pressure protocol](../../docs/intensity_benchmark.md) and [public benchmark API](https://trackformer-weatherlab.rudin-euler-8253.chatgpt.site/api/benchmarks/released). A new matched WeatherNext Cyclones Mini run is in progress on these exact daily starts; results remain pending. See the [frozen protocol](../../docs/deepmind_daily_benchmark.md); no old scores are transferred.
|
| 82 |
|
| 83 |
+
**Pressure-display coverage:** `regional_mslp_hpa` remains the original fixed issue-relative composite. Use the separately exported `core_mslp_hpa` and its per-lead coordinates/mask when the storm moves away from that fixed patch. Do not extend it by inventing a vortex or moving it onto a route. If the moving core leaves supported basin geography, the corresponding masked cells remain unavailable. For an ensemble, register each physical member field onto a common geographic grid **before** averaging; do not average moving-frame array indices. The Mangkhut film uses that separate 50-member policy. Older fixed-patch films retain their original data in [the showcase archive](../../docs/showcase_archive.md). The automatic History archive and its separate core recovery expose geography and coverage through [the public API](https://trackformer-weatherlab.rudin-euler-8253.chatgpt.site/history-api).
|
models/trackformer_1_2_field/plot_pressure.py
ADDED
|
@@ -0,0 +1,91 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Plot the released model's saved geographic pressure fields without inference.
|
| 2 |
+
|
| 3 |
+
The basin remains at 2.5 degrees. The moving core is the model's learned 20-km
|
| 4 |
+
reconstruction, not extra native observations. Invalid coverage stays masked.
|
| 5 |
+
"""
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import argparse
|
| 9 |
+
from datetime import datetime, timezone
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def render_pressure_map(data, output, lead=120, interval=4):
|
| 16 |
+
if not np.isfinite(interval) or interval <= 0:
|
| 17 |
+
raise ValueError("Isobar interval must be positive and finite")
|
| 18 |
+
indices = np.flatnonzero(np.asarray(data["lead_hours"]) == lead)
|
| 19 |
+
if len(indices) != 1:
|
| 20 |
+
raise ValueError("Choose a saved +6 to +120-hour forecast lead")
|
| 21 |
+
required = {"core_mslp_hpa", "core_latitude_deg", "core_longitude_deg", "core_valid",
|
| 22 |
+
"basin_latitude_deg", "basin_longitude_deg", "track_valid", "valid_time_ns", "member_count"}
|
| 23 |
+
if not required.issubset(data):
|
| 24 |
+
raise ValueError("This renderer requires the moving-core pressure export v2")
|
| 25 |
+
if int(np.asarray(data["member_count"])) != 1:
|
| 26 |
+
raise ValueError("This renderer accepts the clean one-member export, not unregistered ensemble cores")
|
| 27 |
+
import matplotlib
|
| 28 |
+
matplotlib.use("Agg")
|
| 29 |
+
import matplotlib.pyplot as plt
|
| 30 |
+
|
| 31 |
+
i = int(indices[0])
|
| 32 |
+
basin = np.ma.masked_invalid(data["basin_mslp_hpa"][i])
|
| 33 |
+
core = np.ma.masked_where(~np.asarray(data["core_valid"][i], bool), data["core_mslp_hpa"][i])
|
| 34 |
+
core = np.ma.masked_invalid(core)
|
| 35 |
+
finite = np.concatenate((basin.compressed(), core.compressed()))
|
| 36 |
+
if not finite.size:
|
| 37 |
+
raise ValueError("No supported pressure field for this lead")
|
| 38 |
+
low = np.floor(finite.min() / interval) * interval
|
| 39 |
+
high = max(low + interval, np.ceil(finite.max() / interval) * interval)
|
| 40 |
+
levels = np.arange(low, high + interval * .5, interval)
|
| 41 |
+
fig, axes = plt.subplots(1, 2, figsize=(12, 5), layout="constrained")
|
| 42 |
+
plots = (
|
| 43 |
+
(basin, data["basin_longitude_deg"], data["basin_latitude_deg"], "Western Pacific · 2.5° basin field"),
|
| 44 |
+
(core, data["core_longitude_deg"][i], data["core_latitude_deg"][i], "Moving core · learned 20-km reconstruction"),
|
| 45 |
+
)
|
| 46 |
+
route = np.asarray(data["track_lat_lon"][:i + 1])
|
| 47 |
+
valid_route = np.asarray(data["track_valid"][:i + 1], bool)
|
| 48 |
+
for ax, (pressure, lon, lat, title) in zip(axes, plots):
|
| 49 |
+
ax.set_title(title, fontsize=11)
|
| 50 |
+
ax.set_xlabel("Longitude °E")
|
| 51 |
+
ax.set_ylabel("Latitude °N")
|
| 52 |
+
if pressure.count():
|
| 53 |
+
image = ax.pcolormesh(lon, lat, pressure, cmap="RdBu_r", vmin=low, vmax=high,
|
| 54 |
+
shading="auto", rasterized=True)
|
| 55 |
+
if pressure.max() > pressure.min():
|
| 56 |
+
lines = ax.contour(lon, lat, pressure, levels=levels, colors="#35424d", linewidths=.65)
|
| 57 |
+
ax.clabel(lines, levels[::2], fmt="%d", fontsize=7)
|
| 58 |
+
else:
|
| 59 |
+
ax.text(.5, .5, "Moving core outside supported basin coverage", ha="center", va="center", transform=ax.transAxes)
|
| 60 |
+
ax.plot(np.where(valid_route, route[:, 1], np.nan), np.where(valid_route, route[:, 0], np.nan),
|
| 61 |
+
color="#ad287b", linewidth=1.6)
|
| 62 |
+
if valid_route[i] and (ax is axes[0] or bool(np.asarray(data["core_valid"][i]).any())):
|
| 63 |
+
ax.scatter(route[i, 1], route[i, 0], s=25, c="#ad287b", edgecolors="white", zorder=4)
|
| 64 |
+
ax.grid(alpha=.2)
|
| 65 |
+
axes[0].set_xlim(100, 180)
|
| 66 |
+
axes[0].set_ylim(0, 60)
|
| 67 |
+
axes[1].set_xlim(float(np.min(plots[1][1])), float(np.max(plots[1][1])))
|
| 68 |
+
axes[1].set_ylim(float(np.min(plots[1][2])), float(np.max(plots[1][2])))
|
| 69 |
+
valid_time = datetime.fromtimestamp(int(np.asarray(data["valid_time_ns"])[i]) / 1e9, timezone.utc)
|
| 70 |
+
value = (f"{float(data['central_pressure_hpa'][i]):.1f} hPa" if valid_route[i] else "unavailable outside domain")
|
| 71 |
+
fig.suptitle(f"Trackformer 1.2 · +{lead:03d} h · {valid_time:%Y-%m-%d %H:%M UTC}\n"
|
| 72 |
+
f"1 clean member · central pressure {value}", fontsize=12)
|
| 73 |
+
fig.colorbar(image, ax=axes, label="Model MSLP (hPa) · blue low / red high", shrink=.85)
|
| 74 |
+
fig.supxlabel(f"{interval:g} hPa isobars · original model geography · unsupported core cells masked; no pressure or route shifting", fontsize=9)
|
| 75 |
+
fig.savefig(Path(output), dpi=160)
|
| 76 |
+
plt.close(fig)
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def main():
|
| 80 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 81 |
+
parser.add_argument("forecast", type=Path)
|
| 82 |
+
parser.add_argument("output", type=Path)
|
| 83 |
+
parser.add_argument("--lead", type=int, default=120, choices=range(6, 121, 6))
|
| 84 |
+
parser.add_argument("--interval", type=float, default=4)
|
| 85 |
+
args = parser.parse_args()
|
| 86 |
+
with np.load(args.forecast, allow_pickle=False) as saved:
|
| 87 |
+
render_pressure_map({key: saved[key] for key in saved.files}, args.output, args.lead, args.interval)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
if __name__ == "__main__":
|
| 91 |
+
main()
|
models/trackformer_1_2_field/predict.py
CHANGED
|
@@ -6,6 +6,7 @@ must supply past/issue-time analyses on the exact grids in manifest.json.
|
|
| 6 |
from __future__ import annotations
|
| 7 |
|
| 8 |
import argparse
|
|
|
|
| 9 |
import json
|
| 10 |
from pathlib import Path
|
| 11 |
|
|
@@ -27,6 +28,38 @@ INPUT_SHAPES = {
|
|
| 27 |
"issue_mask": (2,),
|
| 28 |
}
|
| 29 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
|
| 31 |
def load_packet(path: Path) -> dict[str, np.ndarray]:
|
| 32 |
with np.load(path, allow_pickle=False) as packet:
|
|
@@ -54,7 +87,12 @@ def forecast(packet: Path, model_dir: Path, device: str = "cpu") -> dict[str, np
|
|
| 54 |
raise ValueError("Unexpected public model version")
|
| 55 |
contract = metadata["data_contract"]
|
| 56 |
model = CoreForecaster(contract).to(device).eval()
|
| 57 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 58 |
model.load_state_dict(state, strict=True)
|
| 59 |
inputs = {key: torch.from_numpy(value).to(device) for key, value in load_packet(packet).items()}
|
| 60 |
with torch.inference_mode():
|
|
@@ -73,6 +111,9 @@ def forecast(packet: Path, model_dir: Path, device: str = "cpu") -> dict[str, np
|
|
| 73 |
"regional_mslp_hpa": torch.stack([o["regional"][0, 0] for o in outputs]).cpu().numpy() * scale + offset,
|
| 74 |
"maximum_wind_auxiliary_kt": torch.stack([o["vmax"][0] for o in outputs]).cpu().numpy(),
|
| 75 |
}
|
|
|
|
|
|
|
|
|
|
| 76 |
diagnostics = [summarize_members(diagnose_outputs(o, contract)) for o in outputs]
|
| 77 |
result['maximum_wind_auxiliary_kt_valid'] = np.asarray(
|
| 78 |
[d['estimates']['maximum_wind_auxiliary_kt']['mean'] is not None for d in diagnostics], dtype=bool)
|
|
@@ -84,6 +125,24 @@ def forecast(packet: Path, model_dir: Path, device: str = "cpu") -> dict[str, np
|
|
| 84 |
if not all(np.isfinite(value).all() for value in result.values()):
|
| 85 |
raise ValueError("Non-finite forecast")
|
| 86 |
result['wind_estimation_json'] = np.asarray(json.dumps(diagnostics, allow_nan=False))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 87 |
return result
|
| 88 |
|
| 89 |
|
|
@@ -92,10 +151,17 @@ def main() -> None:
|
|
| 92 |
parser.add_argument("packet", type=Path, help="Causal normalized .npz issue packet")
|
| 93 |
parser.add_argument("output", type=Path, help="Output .npz path")
|
| 94 |
parser.add_argument("--device", default="cpu", choices=("cpu", "mps", "cuda"))
|
|
|
|
|
|
|
|
|
|
| 95 |
args = parser.parse_args()
|
| 96 |
result = forecast(args.packet, Path(__file__).resolve().parent, args.device)
|
| 97 |
np.savez_compressed(args.output, **result)
|
| 98 |
print(f"Saved 20 six-hour forecasts to {args.output}")
|
|
|
|
|
|
|
|
|
|
|
|
|
| 99 |
|
| 100 |
|
| 101 |
if __name__ == "__main__":
|
|
|
|
| 6 |
from __future__ import annotations
|
| 7 |
|
| 8 |
import argparse
|
| 9 |
+
import hashlib
|
| 10 |
import json
|
| 11 |
from pathlib import Path
|
| 12 |
|
|
|
|
| 28 |
"issue_mask": (2,),
|
| 29 |
}
|
| 30 |
|
| 31 |
+
PRESSURE_EXPORT_SCHEMA = "trackformer-1.2-pressure-fields-v2"
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def export_pressure_fields(outputs, inputs, contract, issue_time_ns):
|
| 35 |
+
"""Capture existing forward outputs; never relocate or fit a display vortex."""
|
| 36 |
+
scale = float(contract["normalization"]["std"][0])
|
| 37 |
+
offset = float(contract["normalization"]["mean"][0])
|
| 38 |
+
global_static = inputs["global_static"][0].detach().cpu().numpy()
|
| 39 |
+
regional_static = inputs["regional_static"][0].detach().cpu().numpy()
|
| 40 |
+
leads = np.arange(6, 121, 6, dtype=np.int16)
|
| 41 |
+
if len(outputs) != len(leads):
|
| 42 |
+
raise ValueError("Detailed pressure export requires all twenty forecast leads")
|
| 43 |
+
return {
|
| 44 |
+
"core_mslp_hpa": torch.stack([
|
| 45 |
+
o["core"][0, 0] * scale + offset for o in outputs]).cpu().numpy(),
|
| 46 |
+
"core_latitude_deg": torch.stack([o["core_lat"][0] for o in outputs]).cpu().numpy(),
|
| 47 |
+
"core_longitude_deg": torch.stack([o["core_lon"][0] for o in outputs]).cpu().numpy(),
|
| 48 |
+
"core_valid": torch.stack([o["core_valid"][0, 0] for o in outputs]).cpu().numpy().astype(bool),
|
| 49 |
+
"basin_latitude_deg": global_static[0] * 90,
|
| 50 |
+
"basin_longitude_deg": (global_static[1] + 1) * 180,
|
| 51 |
+
"regional_latitude_deg": regional_static[0] * 90,
|
| 52 |
+
"regional_longitude_deg": (regional_static[1] + 1) * 180,
|
| 53 |
+
"regional_valid": torch.stack([
|
| 54 |
+
o["regional_valid"][0, 0] for o in outputs]).cpu().numpy().astype(bool),
|
| 55 |
+
"track_valid": torch.stack([o["track_valid"][0] for o in outputs]).cpu().numpy().astype(bool),
|
| 56 |
+
"issue_center_lat_lon": inputs["center"][0].detach().cpu().numpy(),
|
| 57 |
+
"issue_time_ns": np.asarray(issue_time_ns, dtype=np.int64),
|
| 58 |
+
"valid_time_ns": issue_time_ns + leads.astype(np.int64) * (3600 * 10**9),
|
| 59 |
+
"member_count": np.asarray(1, dtype=np.int16),
|
| 60 |
+
"core_information_spacing_km": np.asarray(20, dtype=np.float32),
|
| 61 |
+
}
|
| 62 |
+
|
| 63 |
|
| 64 |
def load_packet(path: Path) -> dict[str, np.ndarray]:
|
| 65 |
with np.load(path, allow_pickle=False) as packet:
|
|
|
|
| 87 |
raise ValueError("Unexpected public model version")
|
| 88 |
contract = metadata["data_contract"]
|
| 89 |
model = CoreForecaster(contract).to(device).eval()
|
| 90 |
+
weights = model_dir / "weights.pt"
|
| 91 |
+
with weights.open("rb") as stream:
|
| 92 |
+
weight_sha256 = hashlib.file_digest(stream, "sha256").hexdigest()
|
| 93 |
+
if weight_sha256 != metadata["inference_weights_sha256"]:
|
| 94 |
+
raise ValueError("Weights do not match the released 1.2 manifest")
|
| 95 |
+
state = torch.load(weights, map_location="cpu", weights_only=True)
|
| 96 |
model.load_state_dict(state, strict=True)
|
| 97 |
inputs = {key: torch.from_numpy(value).to(device) for key, value in load_packet(packet).items()}
|
| 98 |
with torch.inference_mode():
|
|
|
|
| 111 |
"regional_mslp_hpa": torch.stack([o["regional"][0, 0] for o in outputs]).cpu().numpy() * scale + offset,
|
| 112 |
"maximum_wind_auxiliary_kt": torch.stack([o["vmax"][0] for o in outputs]).cpu().numpy(),
|
| 113 |
}
|
| 114 |
+
with np.load(packet, allow_pickle=False) as source:
|
| 115 |
+
issue_time_ns = int(source["issue_time_ns"])
|
| 116 |
+
result.update(export_pressure_fields(outputs, inputs, contract, issue_time_ns))
|
| 117 |
diagnostics = [summarize_members(diagnose_outputs(o, contract)) for o in outputs]
|
| 118 |
result['maximum_wind_auxiliary_kt_valid'] = np.asarray(
|
| 119 |
[d['estimates']['maximum_wind_auxiliary_kt']['mean'] is not None for d in diagnostics], dtype=bool)
|
|
|
|
| 125 |
if not all(np.isfinite(value).all() for value in result.values()):
|
| 126 |
raise ValueError("Non-finite forecast")
|
| 127 |
result['wind_estimation_json'] = np.asarray(json.dumps(diagnostics, allow_nan=False))
|
| 128 |
+
result['pressure_export_json'] = np.asarray(json.dumps({
|
| 129 |
+
"schema": PRESSURE_EXPORT_SCHEMA,
|
| 130 |
+
"public_version": "1.2",
|
| 131 |
+
"architecture": metadata["architecture"],
|
| 132 |
+
"inference_weights_sha256": weight_sha256,
|
| 133 |
+
"source_checkpoint_sha256": metadata["source_checkpoint_sha256"],
|
| 134 |
+
"members": 1,
|
| 135 |
+
"units": {"pressure": "hPa", "latitude": "degrees_north", "longitude": "degrees_east"},
|
| 136 |
+
"core_information_spacing_km": 20,
|
| 137 |
+
"core_method": "unchanged learned moving pressure field; physical hPa with original geographic coordinates",
|
| 138 |
+
"native_high_resolution_observations": False,
|
| 139 |
+
"native_detail_history_available": bool(inputs["detail_available"][0, 0].item()),
|
| 140 |
+
"regional_grid": "fixed issue-relative composite; outside moving-core coverage only basin information remains",
|
| 141 |
+
"coverage_policy": "apply core_valid, regional_valid and track_valid; invalid finite storage is not a supported forecast",
|
| 142 |
+
"central_pressure_policy": "existing bilinear moving-core readout at the associated forecast centre; not an independently inserted scalar",
|
| 143 |
+
"ensemble_policy": "one clean member; register physical fields on common geographic coordinates before any ensemble average",
|
| 144 |
+
"forecast_equations_changed": False,
|
| 145 |
+
}, allow_nan=False))
|
| 146 |
return result
|
| 147 |
|
| 148 |
|
|
|
|
| 151 |
parser.add_argument("packet", type=Path, help="Causal normalized .npz issue packet")
|
| 152 |
parser.add_argument("output", type=Path, help="Output .npz path")
|
| 153 |
parser.add_argument("--device", default="cpu", choices=("cpu", "mps", "cuda"))
|
| 154 |
+
parser.add_argument("--pressure-map", type=Path, help="Optional PNG of basin and actual moving-core pressure (requires Matplotlib)")
|
| 155 |
+
parser.add_argument("--map-lead", type=int, default=120, choices=range(6, 121, 6))
|
| 156 |
+
parser.add_argument("--isobar-interval", type=float, default=4, help="Pressure contour interval in hPa")
|
| 157 |
args = parser.parse_args()
|
| 158 |
result = forecast(args.packet, Path(__file__).resolve().parent, args.device)
|
| 159 |
np.savez_compressed(args.output, **result)
|
| 160 |
print(f"Saved 20 six-hour forecasts to {args.output}")
|
| 161 |
+
if args.pressure_map:
|
| 162 |
+
from plot_pressure import render_pressure_map
|
| 163 |
+
render_pressure_map(result, args.pressure_map, args.map_lead, args.isobar_interval)
|
| 164 |
+
print(f"Saved pressure map to {args.pressure_map}")
|
| 165 |
|
| 166 |
|
| 167 |
if __name__ == "__main__":
|
release_tools/sync_public_model_cards.py
CHANGED
|
@@ -29,9 +29,13 @@ SYNC_FILES = (
|
|
| 29 |
'models/trackformer_1_2_field/README.md',
|
| 30 |
'models/trackformer_1_2_field/WIND_ESTIMATION.md',
|
| 31 |
'models/trackformer_1_2_field/predict.py',
|
|
|
|
| 32 |
'models/trackformer_1_2_field/wind_estimation.py',
|
| 33 |
'release_tools/sync_public_model_cards.py',
|
| 34 |
'release_tools/test_public_model_cards.py',
|
|
|
|
|
|
|
|
|
|
| 35 |
'release_tools/build_release_benchmark.py',
|
| 36 |
'release_tools/plot_release_pressure_benchmark.py',
|
| 37 |
'release_tools/plot_daily_storm_final.py',
|
|
|
|
| 29 |
'models/trackformer_1_2_field/README.md',
|
| 30 |
'models/trackformer_1_2_field/WIND_ESTIMATION.md',
|
| 31 |
'models/trackformer_1_2_field/predict.py',
|
| 32 |
+
'models/trackformer_1_2_field/plot_pressure.py',
|
| 33 |
'models/trackformer_1_2_field/wind_estimation.py',
|
| 34 |
'release_tools/sync_public_model_cards.py',
|
| 35 |
'release_tools/test_public_model_cards.py',
|
| 36 |
+
'release_tools/test_pressure_field_export.py',
|
| 37 |
+
'release_tools/verify_pressure_field_export.py',
|
| 38 |
+
'RELEASE_NOTES_TRACKFORMER_1_2.md',
|
| 39 |
'release_tools/build_release_benchmark.py',
|
| 40 |
'release_tools/plot_release_pressure_benchmark.py',
|
| 41 |
'release_tools/plot_daily_storm_final.py',
|
release_tools/test_pressure_field_export.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Regression checks for the capture-only released 1.2 pressure export."""
|
| 2 |
+
import hashlib
|
| 3 |
+
import json
|
| 4 |
+
import sys
|
| 5 |
+
import tempfile
|
| 6 |
+
import unittest
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
|
| 12 |
+
ROOT = Path(__file__).resolve().parents[1]
|
| 13 |
+
sys.path.insert(0, str(ROOT / "models/trackformer_1_2_field"))
|
| 14 |
+
from predict import export_pressure_fields
|
| 15 |
+
from plot_pressure import render_pressure_map
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class PressureFieldExportTest(unittest.TestCase):
|
| 19 |
+
def setUp(self):
|
| 20 |
+
self.contract = {"normalization": {"std": [40], "mean": [1000]}}
|
| 21 |
+
lat, lon = torch.meshgrid(torch.linspace(60, 0, 25), torch.linspace(100, 180, 33), indexing="ij")
|
| 22 |
+
rlat, rlon = torch.meshgrid(torch.linspace(30, 10, 121), torch.linspace(120, 140, 121), indexing="ij")
|
| 23 |
+
static = lambda y, x: torch.stack((y/90, x/180-1, torch.zeros_like(y), torch.zeros_like(y)))[None]
|
| 24 |
+
self.inputs = {"global_static": static(lat, lon), "regional_static": static(rlat, rlon),
|
| 25 |
+
"center": torch.tensor([[20., 130.]])}
|
| 26 |
+
self.outputs = []
|
| 27 |
+
for i in range(20):
|
| 28 |
+
y, x = torch.meshgrid(torch.linspace(25, 15, 65), torch.linspace(125+i*.5, 135+i*.5, 65), indexing="ij")
|
| 29 |
+
core = -torch.exp(-((x-(130+i*.5))**2+(y-20)**2)/3)[None, None]
|
| 30 |
+
mask = torch.ones_like(core, dtype=torch.bool)
|
| 31 |
+
mask[:, :, :, 0] = False
|
| 32 |
+
self.outputs.append({"core": core, "core_lat": y[None], "core_lon": x[None], "core_valid": mask,
|
| 33 |
+
"regional_valid": torch.ones((1,1,121,121),dtype=torch.bool), "track_valid": torch.tensor([True]),
|
| 34 |
+
"pressure": torch.tensor([800.])})
|
| 35 |
+
self.issue = 1536624000 * 10**9
|
| 36 |
+
|
| 37 |
+
def test_core_units_geography_and_every_valid_time(self):
|
| 38 |
+
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
|
| 39 |
+
self.assertEqual(result["core_mslp_hpa"].shape, (20,65,65))
|
| 40 |
+
self.assertEqual(result["core_latitude_deg"].shape, (20,65,65))
|
| 41 |
+
self.assertEqual(result["core_longitude_deg"].shape, (20,65,65))
|
| 42 |
+
self.assertEqual(result["core_valid"].dtype, np.dtype(bool))
|
| 43 |
+
self.assertEqual(result["regional_valid"].shape, (20,121,121))
|
| 44 |
+
np.testing.assert_array_equal(result["valid_time_ns"], self.issue+np.arange(6,121,6,dtype=np.int64)*3600*10**9)
|
| 45 |
+
self.assertEqual(int(result["member_count"]), 1)
|
| 46 |
+
self.assertEqual(float(result["core_information_spacing_km"]), 20)
|
| 47 |
+
self.assertAlmostEqual(float(result["core_mslp_hpa"][0,32,32]),960)
|
| 48 |
+
np.testing.assert_array_equal(result["core_longitude_deg"][19], self.outputs[19]["core_lon"][0].numpy())
|
| 49 |
+
self.assertFalse(result["core_valid"][:,:,0].any())
|
| 50 |
+
|
| 51 |
+
def test_capture_preserves_tensors_and_never_inserts_scalar_pressure(self):
|
| 52 |
+
before = [{k:v.clone() for k,v in o.items()} for o in self.outputs]
|
| 53 |
+
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
|
| 54 |
+
self.assertGreater(float(result["core_mslp_hpa"].min()), float(self.outputs[0]["pressure"][0]))
|
| 55 |
+
for old, new in zip(before, self.outputs):
|
| 56 |
+
for key in old:
|
| 57 |
+
self.assertTrue(torch.equal(old[key], new[key]), key)
|
| 58 |
+
|
| 59 |
+
def test_a_moving_core_can_leave_the_original_fixed_patch(self):
|
| 60 |
+
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
|
| 61 |
+
fixed_east = float(result["regional_longitude_deg"].max())
|
| 62 |
+
self.assertGreater(float(result["core_longitude_deg"][-1].max()), fixed_east)
|
| 63 |
+
np.testing.assert_array_equal(result["regional_longitude_deg"], (self.inputs["regional_static"][0,1].numpy()+1)*180)
|
| 64 |
+
|
| 65 |
+
def test_incomplete_rollout_is_rejected(self):
|
| 66 |
+
with self.assertRaisesRegex(ValueError, "twenty"):
|
| 67 |
+
export_pressure_fields(self.outputs[:-1], self.inputs, self.contract, self.issue)
|
| 68 |
+
|
| 69 |
+
def plot_data(self):
|
| 70 |
+
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
|
| 71 |
+
result.update(lead_hours=np.arange(6,121,6), basin_mslp_hpa=np.stack([
|
| 72 |
+
1005+self.inputs["global_static"][0,0].numpy()*10+i*.1 for i in range(20)]),
|
| 73 |
+
central_pressure_hpa=np.full(20,960),
|
| 74 |
+
track_lat_lon=np.column_stack((np.full(20,20),130+np.arange(20)*.5)))
|
| 75 |
+
return result
|
| 76 |
+
|
| 77 |
+
def test_renderer_accepts_real_geography_and_distinct_leads(self):
|
| 78 |
+
with tempfile.TemporaryDirectory(prefix="pressure-export-test-") as folder:
|
| 79 |
+
first, last = Path(folder)/"first.png", Path(folder)/"last.png"
|
| 80 |
+
data = self.plot_data()
|
| 81 |
+
render_pressure_map(data, first, 6, 4)
|
| 82 |
+
render_pressure_map(data, last, 120, 4)
|
| 83 |
+
self.assertGreater(first.stat().st_size, 10000)
|
| 84 |
+
self.assertNotEqual(hashlib.sha256(first.read_bytes()).digest(), hashlib.sha256(last.read_bytes()).digest())
|
| 85 |
+
|
| 86 |
+
def test_missing_core_is_masked_not_filled_from_scalar(self):
|
| 87 |
+
data = self.plot_data()
|
| 88 |
+
data["core_valid"][:] = False
|
| 89 |
+
original = data["core_mslp_hpa"].copy()
|
| 90 |
+
with tempfile.TemporaryDirectory(prefix="pressure-export-test-") as folder:
|
| 91 |
+
render_pressure_map(data, Path(folder)/"masked.png", 120)
|
| 92 |
+
np.testing.assert_array_equal(data["core_mslp_hpa"], original)
|
| 93 |
+
|
| 94 |
+
def test_invalid_contour_interval_and_unregistered_ensemble_are_rejected(self):
|
| 95 |
+
for interval in (0,-1,np.nan):
|
| 96 |
+
with self.assertRaisesRegex(ValueError, "interval"):
|
| 97 |
+
render_pressure_map(self.plot_data(), "unused.png", interval=interval)
|
| 98 |
+
data = self.plot_data()
|
| 99 |
+
data["member_count"] = np.array(50)
|
| 100 |
+
with self.assertRaisesRegex(ValueError, "one-member"):
|
| 101 |
+
render_pressure_map(data, "unused.png")
|
| 102 |
+
|
| 103 |
+
def test_released_neural_source_hashes_are_unchanged(self):
|
| 104 |
+
directory = ROOT/"models/trackformer_1_2_field"
|
| 105 |
+
manifest = json.loads((directory/"manifest.json").read_text())
|
| 106 |
+
for name, expected in manifest["source_module_sha256"].items():
|
| 107 |
+
self.assertEqual(hashlib.sha256((directory/name).read_bytes()).hexdigest(), expected, name)
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
if __name__ == "__main__":
|
| 111 |
+
unittest.main()
|
release_tools/verify_pressure_field_export.py
ADDED
|
@@ -0,0 +1,120 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""CPU-only export regression using an existing causal batched issue packet.
|
| 2 |
+
|
| 3 |
+
This checks packaging/field consistency, not forecast skill. It never calls MPS,
|
| 4 |
+
changes neural weights, or writes into historical archive/training directories.
|
| 5 |
+
"""
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import argparse
|
| 9 |
+
import hashlib
|
| 10 |
+
import json
|
| 11 |
+
import shutil
|
| 12 |
+
import sys
|
| 13 |
+
from pathlib import Path
|
| 14 |
+
from unittest.mock import patch
|
| 15 |
+
|
| 16 |
+
import numpy as np
|
| 17 |
+
import torch
|
| 18 |
+
|
| 19 |
+
ROOT = Path(__file__).resolve().parents[1]
|
| 20 |
+
MODEL = ROOT / "models/trackformer_1_2_field"
|
| 21 |
+
sys.path.insert(0, str(MODEL))
|
| 22 |
+
import predict
|
| 23 |
+
from model import CoreForecaster
|
| 24 |
+
from baseline_model import base
|
| 25 |
+
from plot_pressure import render_pressure_map
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def sha(path):
|
| 29 |
+
with Path(path).open("rb") as stream:
|
| 30 |
+
return hashlib.file_digest(stream, "sha256").hexdigest()
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def verify(batched_packet, weights, work):
|
| 34 |
+
work = work.resolve()
|
| 35 |
+
if not work.is_relative_to(Path("/Volumes/D")):
|
| 36 |
+
raise ValueError("Diagnostic artifacts must remain on D")
|
| 37 |
+
work.mkdir(parents=True, exist_ok=True)
|
| 38 |
+
directory = work / "inference"
|
| 39 |
+
directory.mkdir(exist_ok=True)
|
| 40 |
+
for source in MODEL.iterdir():
|
| 41 |
+
if source.is_file() and source.suffix in (".py", ".md", ".json"):
|
| 42 |
+
shutil.copy2(source, directory / source.name)
|
| 43 |
+
shutil.copy2(weights, directory / "weights.pt")
|
| 44 |
+
metadata = json.loads((directory / "manifest.json").read_text())
|
| 45 |
+
for name, expected in metadata["source_module_sha256"].items():
|
| 46 |
+
if sha(directory / name) != expected:
|
| 47 |
+
raise ValueError("Frozen neural module changed: " + name)
|
| 48 |
+
with np.load(batched_packet, allow_pickle=False) as source:
|
| 49 |
+
arrays = {key: np.asarray(source[key])[0] for key in predict.INPUT_SHAPES}
|
| 50 |
+
arrays["history_time_ns"] = np.asarray(source["history_time_ns"], dtype=np.int64)
|
| 51 |
+
arrays["issue_time_ns"] = np.asarray(arrays["history_time_ns"][-1], dtype=np.int64)
|
| 52 |
+
packet = work / "causal_issue_packet.npz"
|
| 53 |
+
np.savez_compressed(packet, **arrays)
|
| 54 |
+
captured, origins = [], []
|
| 55 |
+
|
| 56 |
+
class CapturingModel(CoreForecaster):
|
| 57 |
+
def step(self, state):
|
| 58 |
+
next_state, out = super().step(state)
|
| 59 |
+
captured.append(out)
|
| 60 |
+
origins.append(next_state["origin"])
|
| 61 |
+
return next_state, out
|
| 62 |
+
|
| 63 |
+
torch.set_num_threads(2)
|
| 64 |
+
torch.set_num_interop_threads(1)
|
| 65 |
+
with patch.object(predict, "CoreForecaster", CapturingModel):
|
| 66 |
+
result = predict.forecast(packet, directory, "cpu")
|
| 67 |
+
contract = metadata["data_contract"]
|
| 68 |
+
scale, offset = contract["normalization"]["std"][0], contract["normalization"]["mean"][0]
|
| 69 |
+
reference = {
|
| 70 |
+
"track_lat_lon": torch.stack([o["center"][0] for o in captured]).numpy(),
|
| 71 |
+
"central_pressure_hpa": torch.stack([o["pressure"][0] for o in captured]).numpy(),
|
| 72 |
+
"basin_mslp_hpa": torch.stack([o["global"][0,0] for o in captured]).numpy()*scale+offset,
|
| 73 |
+
"regional_mslp_hpa": torch.stack([o["regional"][0,0] for o in captured]).numpy()*scale+offset,
|
| 74 |
+
"maximum_wind_auxiliary_kt": torch.stack([o["vmax"][0] for o in captured]).numpy(),
|
| 75 |
+
}
|
| 76 |
+
for key, expected in reference.items():
|
| 77 |
+
np.testing.assert_array_equal(result[key], expected, err_msg="Original readout changed: " + key)
|
| 78 |
+
alignment = []
|
| 79 |
+
for i, (out, origin) in enumerate(zip(captured, origins)):
|
| 80 |
+
# The original model's own geographic sampler, not a centre/minimum fit.
|
| 81 |
+
grid = CoreForecaster.core_grid(None, out["center"][:,0,None,None], out["center"][:,1,None,None], origin)
|
| 82 |
+
sampled = base.sample_field(torch.from_numpy(result["core_mslp_hpa"][i])[None,None], grid)[0,0,0,0]
|
| 83 |
+
alignment.append(abs(float(sampled) - float(result["central_pressure_hpa"][i])))
|
| 84 |
+
np.testing.assert_array_equal(result["core_latitude_deg"][i], out["core_lat"][0].numpy())
|
| 85 |
+
np.testing.assert_array_equal(result["core_longitude_deg"][i], out["core_lon"][0].numpy())
|
| 86 |
+
np.testing.assert_array_equal(result["core_valid"][i], out["core_valid"][0,0].numpy())
|
| 87 |
+
if max(alignment) > 1e-3:
|
| 88 |
+
raise ValueError("Exported core does not reproduce its original central-pressure readout")
|
| 89 |
+
np.savez_compressed(work / "forecast.npz", **result)
|
| 90 |
+
for lead in (6, 120):
|
| 91 |
+
render_pressure_map(result, work / f"pressure_{lead:03d}h.png", lead, 4)
|
| 92 |
+
receipt = {
|
| 93 |
+
"state": "verified_cpu_export_regression",
|
| 94 |
+
"scientific_benchmark": False,
|
| 95 |
+
"model": "Trackformer 1.2", "members": 1,
|
| 96 |
+
"device": "cpu", "forecast_steps": 20,
|
| 97 |
+
"issue_time_ns": int(result["issue_time_ns"]),
|
| 98 |
+
"core_shape": list(result["core_mslp_hpa"].shape),
|
| 99 |
+
"core_min_hpa": float(result["core_mslp_hpa"][result["core_valid"]].min()),
|
| 100 |
+
"core_max_hpa": float(result["core_mslp_hpa"][result["core_valid"]].max()),
|
| 101 |
+
"central_pressure_alignment_max_error_hpa": max(alignment),
|
| 102 |
+
"original_outputs_bit_identical": list(reference),
|
| 103 |
+
"coordinates_and_masks_match_forward_outputs": True,
|
| 104 |
+
"inference_weights_sha256": sha(weights),
|
| 105 |
+
"source_packet_sha256": sha(batched_packet),
|
| 106 |
+
"forecast_sha256": sha(work / "forecast.npz"),
|
| 107 |
+
"native_detail_history_available": bool(arrays["detail_available"][0]),
|
| 108 |
+
"core_is_learned_reconstruction_not_native_observations": True,
|
| 109 |
+
}
|
| 110 |
+
(work / "verification.json").write_text(json.dumps(receipt, indent=2) + "\n")
|
| 111 |
+
print(json.dumps(receipt, indent=2))
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
if __name__ == "__main__":
|
| 115 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 116 |
+
parser.add_argument("--batched-packet", type=Path, required=True)
|
| 117 |
+
parser.add_argument("--weights", type=Path, required=True)
|
| 118 |
+
parser.add_argument("--work", type=Path, required=True)
|
| 119 |
+
args = parser.parse_args()
|
| 120 |
+
verify(args.batched_packet, args.weights, args.work)
|