OneScience-Group/OneForecast
07
1"""Create quick field images from OneForecast prediction files."""2 3from __future__ import annotations4 5from pathlib import Path6import argparse7import numpy as np8import yaml9 10 11def main() -> None:12 parser = argparse.ArgumentParser()13 parser.add_argument("--config", type=Path, default=Path("conf/config.yaml"))14 args = parser.parse_args()15 with args.config.open("r", encoding="utf-8") as handle:16 config = yaml.safe_load(handle)17 root = args.config.resolve().parent.parent18 input_dir = Path(config["visualization"]["input_dir"])19 output_dir = Path(config["visualization"]["output_dir"])20 if not input_dir.is_absolute():21 input_dir = root / input_dir22 if not output_dir.is_absolute():23 output_dir = root / output_dir24 output_dir.mkdir(parents=True, exist_ok=True)25 files = sorted(input_dir.glob("prediction_*.npy"))26 if not files:27 raise SystemExit(f"No prediction files found in {input_dir}")28 import matplotlib.pyplot as plt29 30 channels = config["visualization"].get("channels", [0])31 for source in files:32 prediction = np.load(source)33 if prediction.shape != (1, 69, 120, 240):34 raise ValueError(f"Expected official prediction shape [1, 69, 120, 240], got {prediction.shape}")35 field = prediction[0]36 for channel in channels:37 if channel < 0 or channel >= field.shape[0]:38 raise ValueError(f"Channel {channel} is outside prediction shape {field.shape}")39 figure, axis = plt.subplots(figsize=(8, 3.5))40 image = axis.imshow(field[channel], cmap="coolwarm", aspect="auto")41 axis.set_title(f"{source.stem}, channel {channel}")42 axis.set_xlabel("longitude index")43 axis.set_ylabel("latitude index")44 figure.colorbar(image, ax=axis, shrink=0.8)45 figure.tight_layout()46 figure.savefig(output_dir / f"{source.stem}_ch{channel}.png", dpi=160)47 plt.close(figure)48 print({"input_dir": str(input_dir), "output_dir": str(output_dir), "files": len(files)})49 50 51if __name__ == "__main__":52 main()53 