| """Generate native packed-wedge ensemble forecasts.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import sys |
| ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) |
|
|
| from model.echocast_3d import PackedRadarDataset, load_config, DiffusionSchedule, EchoCast3D |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--config", default="conf/config.yaml") |
| parser.add_argument("--data", default="data") |
| parser.add_argument("--checkpoint", default="result/checkpoints/echocast_3d.pt") |
| parser.add_argument("--output", default="result/output/predictions.npz") |
| args = parser.parse_args() |
| config = load_config(ROOT / args.config) |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| model = EchoCast3D(config).to(device) |
| checkpoint = torch.load(ROOT / args.checkpoint, map_location=device, weights_only=False) |
| model.load_state_dict(checkpoint["model"]) |
| model.eval() |
| sample = PackedRadarDataset(ROOT / args.data, "validation")[0] |
| history = sample["values"][None, :3].float().to(device) |
| observed = sample["observed"][None, :3].bool().to(device) |
| schedule = DiffusionSchedule(device=device, **{k: config["diffusion"][k] for k in ("steps", "beta_start", "beta_end")}) |
| ensemble = [] |
| for member in range(config["diffusion"]["ensemble_size"]): |
| torch.manual_seed(config["seed"] + member) |
| ensemble.append(schedule.sample(model, history, observed)[0].cpu().numpy()) |
| output = ROOT / args.output |
| output.parent.mkdir(parents=True, exist_ok=True) |
| np.savez( |
| output, |
| ensemble=np.stack(ensemble), |
| truth=sample["truth"][3:].numpy(), |
| validity=sample["validity"][3:].numpy(), |
| observed_history=sample["observed"][:3].numpy(), |
| ) |
| print(f"saved {output}: ensemble={len(ensemble)}, steps={schedule.steps}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|