Download release_tools/test_pressure_field_export.py from euler314/typhoon-predict: direct link, hf CLI and curl.
- Browser
- Download file 6.14 kB
-
https://huggingface.co/euler314/typhoon-predict/resolve/main/release_tools/test_pressure_field_export.py
- Command line
-
hf download hf://euler314/typhoon-predict/release_tools/test_pressure_field_export.py
-
curl -L -o test_pressure_field_export.py https://huggingface.co/euler314/typhoon-predict/resolve/main/release_tools/test_pressure_field_export.py
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() | |