CoolFace
Modelpublic

OneScience-Group/OneForecast

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes7downloads
result.py53 linesDownload Raw Back to scripts
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