"""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()