| from pathlib import Path |
| import sys |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| if str(ROOT) not in sys.path: |
| sys.path.insert(0, str(ROOT)) |
|
|
| from model.stablenn_phys import StableNNPhys, rollout |
|
|
|
|
| CHECKPOINT_FORMAT_VERSION = "stablenn_phys_checkpoint_v1" |
| REQUIRED_CHECKPOINT_FIELDS = {"model", "model_config", "format_version"} |
|
|
|
|
| def main(): |
| cfg = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| raw = np.load(ROOT / cfg["data"]["file"]) |
| checkpoint = torch.load(ROOT / cfg["paths"]["checkpoint"], map_location="cpu", weights_only=False) |
| missing = REQUIRED_CHECKPOINT_FIELDS.difference(checkpoint) |
| if missing: |
| raise ValueError(f"Checkpoint is missing required fields: {', '.join(sorted(missing))}") |
| if checkpoint["format_version"] != CHECKPOINT_FORMAT_VERSION: |
| raise ValueError( |
| f"Unsupported checkpoint format_version {checkpoint['format_version']!r}; " |
| f"expected {CHECKPOINT_FORMAT_VERSION!r}" |
| ) |
| model = StableNNPhys(checkpoint["model_config"]["hidden_size"]) |
| model.load_state_dict(checkpoint["model"]) |
| model.eval() |
| steps = cfg["runtime"]["rollout_steps"] |
| if steps != 64: |
| raise ValueError(f"runtime.rollout_steps must be 64, got {steps}") |
| initial = torch.from_numpy(raw["long_state"][:, 0].astype(np.float32)) |
| surface = torch.from_numpy(raw["long_surface"][:, :steps].astype(np.float32)) |
| advection = torch.from_numpy(raw["long_advection"][:, :steps + 1].astype(np.float32)) |
| norm = checkpoint["normalization"] |
| with torch.no_grad(): |
| prediction, tendency = rollout(model, initial, surface, advection, norm["state_mean"], norm["state_std"], |
| norm["tendency_mean"], norm["tendency_std"], checkpoint["dt_seconds"]) |
| target = raw["long_state"][:, :steps + 1] |
| out = ROOT / cfg["paths"]["inference"] |
| out.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed(out, prediction=prediction.numpy(), target=target, tendency=tendency.numpy(), |
| surface=surface.numpy(), advection=advection.numpy(), layer_mass=raw["layer_mass"], |
| source=np.arange(initial.shape[0]), |
| lead_hours=np.arange(steps + 1) * cfg["data"]["step_hours"], |
| dt_seconds=np.float32(checkpoint["dt_seconds"])) |
| print(f"saved {out}: {steps} steps, {steps * 3 / 24:.1f} days") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|