typhoon-predict / release_tools /test_pressure_field_export.py
euler314's picture
Publish Trackformer 1.2 announcement assets and matched DeepMind benchmark protocol
4ebcd24 verified
Raw History Blame Contribute Delete
6.14 kB
"""Regression checks for the capture-only released 1.2 pressure export."""
import hashlib
import json
import sys
import tempfile
import unittest
from pathlib import Path
import numpy as np
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "models/trackformer_1_2_field"))
from predict import export_pressure_fields
from plot_pressure import render_pressure_map
class PressureFieldExportTest(unittest.TestCase):
def setUp(self):
self.contract = {"normalization": {"std": [40], "mean": [1000]}}
lat, lon = torch.meshgrid(torch.linspace(60, 0, 25), torch.linspace(100, 180, 33), indexing="ij")
rlat, rlon = torch.meshgrid(torch.linspace(30, 10, 121), torch.linspace(120, 140, 121), indexing="ij")
static = lambda y, x: torch.stack((y/90, x/180-1, torch.zeros_like(y), torch.zeros_like(y)))[None]
self.inputs = {"global_static": static(lat, lon), "regional_static": static(rlat, rlon),
"center": torch.tensor([[20., 130.]])}
self.outputs = []
for i in range(20):
y, x = torch.meshgrid(torch.linspace(25, 15, 65), torch.linspace(125+i*.5, 135+i*.5, 65), indexing="ij")
core = -torch.exp(-((x-(130+i*.5))**2+(y-20)**2)/3)[None, None]
mask = torch.ones_like(core, dtype=torch.bool)
mask[:, :, :, 0] = False
self.outputs.append({"core": core, "core_lat": y[None], "core_lon": x[None], "core_valid": mask,
"regional_valid": torch.ones((1,1,121,121),dtype=torch.bool), "track_valid": torch.tensor([True]),
"pressure": torch.tensor([800.])})
self.issue = 1536624000 * 10**9
def test_core_units_geography_and_every_valid_time(self):
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
self.assertEqual(result["core_mslp_hpa"].shape, (20,65,65))
self.assertEqual(result["core_latitude_deg"].shape, (20,65,65))
self.assertEqual(result["core_longitude_deg"].shape, (20,65,65))
self.assertEqual(result["core_valid"].dtype, np.dtype(bool))
self.assertEqual(result["regional_valid"].shape, (20,121,121))
np.testing.assert_array_equal(result["valid_time_ns"], self.issue+np.arange(6,121,6,dtype=np.int64)*3600*10**9)
self.assertEqual(int(result["member_count"]), 1)
self.assertEqual(float(result["core_information_spacing_km"]), 20)
self.assertAlmostEqual(float(result["core_mslp_hpa"][0,32,32]),960)
np.testing.assert_array_equal(result["core_longitude_deg"][19], self.outputs[19]["core_lon"][0].numpy())
self.assertFalse(result["core_valid"][:,:,0].any())
def test_capture_preserves_tensors_and_never_inserts_scalar_pressure(self):
before = [{k:v.clone() for k,v in o.items()} for o in self.outputs]
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
self.assertGreater(float(result["core_mslp_hpa"].min()), float(self.outputs[0]["pressure"][0]))
for old, new in zip(before, self.outputs):
for key in old:
self.assertTrue(torch.equal(old[key], new[key]), key)
def test_a_moving_core_can_leave_the_original_fixed_patch(self):
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
fixed_east = float(result["regional_longitude_deg"].max())
self.assertGreater(float(result["core_longitude_deg"][-1].max()), fixed_east)
np.testing.assert_array_equal(result["regional_longitude_deg"], (self.inputs["regional_static"][0,1].numpy()+1)*180)
def test_incomplete_rollout_is_rejected(self):
with self.assertRaisesRegex(ValueError, "twenty"):
export_pressure_fields(self.outputs[:-1], self.inputs, self.contract, self.issue)
def plot_data(self):
result = export_pressure_fields(self.outputs, self.inputs, self.contract, self.issue)
result.update(lead_hours=np.arange(6,121,6), basin_mslp_hpa=np.stack([
1005+self.inputs["global_static"][0,0].numpy()*10+i*.1 for i in range(20)]),
central_pressure_hpa=np.full(20,960),
track_lat_lon=np.column_stack((np.full(20,20),130+np.arange(20)*.5)))
return result
def test_renderer_accepts_real_geography_and_distinct_leads(self):
with tempfile.TemporaryDirectory(prefix="pressure-export-test-") as folder:
first, last = Path(folder)/"first.png", Path(folder)/"last.png"
data = self.plot_data()
render_pressure_map(data, first, 6, 4)
render_pressure_map(data, last, 120, 4)
self.assertGreater(first.stat().st_size, 10000)
self.assertNotEqual(hashlib.sha256(first.read_bytes()).digest(), hashlib.sha256(last.read_bytes()).digest())
def test_missing_core_is_masked_not_filled_from_scalar(self):
data = self.plot_data()
data["core_valid"][:] = False
original = data["core_mslp_hpa"].copy()
with tempfile.TemporaryDirectory(prefix="pressure-export-test-") as folder:
render_pressure_map(data, Path(folder)/"masked.png", 120)
np.testing.assert_array_equal(data["core_mslp_hpa"], original)
def test_invalid_contour_interval_and_unregistered_ensemble_are_rejected(self):
for interval in (0,-1,np.nan):
with self.assertRaisesRegex(ValueError, "interval"):
render_pressure_map(self.plot_data(), "unused.png", interval=interval)
data = self.plot_data()
data["member_count"] = np.array(50)
with self.assertRaisesRegex(ValueError, "one-member"):
render_pressure_map(data, "unused.png")
def test_released_neural_source_hashes_are_unchanged(self):
directory = ROOT/"models/trackformer_1_2_field"
manifest = json.loads((directory/"manifest.json").read_text())
for name, expected in manifest["source_module_sha256"].items():
self.assertEqual(hashlib.sha256((directory/name).read_bytes()).hexdigest(), expected, name)
if __name__ == "__main__":
unittest.main()