#!/usr/bin/env python3 """Render NowcastNet predictions and truth comparisons as RGB PNG files.""" from __future__ import annotations import argparse import json from pathlib import Path import struct import zlib import numpy as np import yaml PROJECT_ROOT = Path(__file__).resolve().parents[1] RAIN_THRESHOLDS = np.asarray([0.1, 1.0, 2.0, 4.0, 8.0, 16.0, 32.0, 64.0], dtype=np.float32) RAIN_COLORS = np.asarray( [ [0, 0, 0], [70, 70, 70], [0, 110, 255], [0, 205, 255], [0, 190, 80], [255, 230, 0], [255, 145, 0], [235, 35, 30], [205, 0, 180], ], dtype=np.uint8, ) ERROR_THRESHOLDS = np.asarray([0.1, 0.5, 1.0, 2.0, 4.0, 8.0, 16.0, 32.0], dtype=np.float32) ERROR_COLORS = np.asarray( [ [0, 0, 0], [40, 40, 40], [35, 80, 170], [30, 165, 215], [80, 200, 120], [245, 225, 65], [245, 145, 45], [220, 55, 40], [245, 245, 245], ], dtype=np.uint8, ) def _png_chunk(kind: bytes, payload: bytes) -> bytes: checksum = zlib.crc32(kind + payload) & 0xFFFFFFFF return struct.pack(">I", len(payload)) + kind + payload + struct.pack(">I", checksum) def write_png(path: Path, image: np.ndarray) -> None: """Write an H x W x 3 uint8 array as a standards-compliant RGB PNG.""" image = np.asarray(image, dtype=np.uint8) if image.ndim != 3 or image.shape[2] != 3: raise ValueError(f"Expected RGB image [H,W,3], got {image.shape}") raw = b"".join(b"\x00" + row.tobytes() for row in image) header = struct.pack(">IIBBBBB", image.shape[1], image.shape[0], 8, 2, 0, 0, 0) path.write_bytes( b"\x89PNG\r\n\x1a\n" + _png_chunk(b"IHDR", header) + _png_chunk(b"IDAT", zlib.compress(raw, 1)) + _png_chunk(b"IEND", b"") ) def colorize(image: np.ndarray, thresholds: np.ndarray, colors: np.ndarray) -> np.ndarray: values = np.nan_to_num(np.asarray(image, dtype=np.float32), nan=0.0, posinf=128.0, neginf=0.0) return colors[np.searchsorted(thresholds, np.maximum(values, 0.0), side="right")] def comparison_image(truth: np.ndarray, prediction: np.ndarray) -> np.ndarray: truth_rgb = colorize(truth, RAIN_THRESHOLDS, RAIN_COLORS) prediction_rgb = colorize(prediction, RAIN_THRESHOLDS, RAIN_COLORS) error_rgb = colorize(np.abs(prediction - truth), ERROR_THRESHOLDS, ERROR_COLORS) separator = np.full((truth.shape[0], 4, 3), 255, dtype=np.uint8) return np.concatenate([truth_rgb, separator, prediction_rgb, separator, error_rgb], axis=1) def main() -> None: parser = argparse.ArgumentParser(description="Render NowcastNet inference results as PNG images") parser.add_argument("--config", default=str(PROJECT_ROOT / "conf/config.yaml")) parser.add_argument("--input-dir", help="directory containing *_pred.npy and *_target.npy") parser.add_argument("--output-dir") parser.add_argument("--threshold", type=float) args = parser.parse_args() cfg = yaml.safe_load(Path(args.config).read_text()) src = Path(args.input_dir) if args.input_dir else PROJECT_ROOT / cfg["inference"]["output_dir"] out = Path(args.output_dir) if args.output_dir else PROJECT_ROOT / cfg["visualization"]["output_dir"] threshold = args.threshold if args.threshold is not None else float(cfg["inference"]["threshold"]) expected_frames = int(cfg["model"]["total_length"]) - int(cfg["model"]["input_length"]) prediction_dir = out / "predictions" comparison_dir = out / "comparison" prediction_dir.mkdir(parents=True, exist_ok=True) comparison_dir.mkdir(parents=True, exist_ok=True) summary: dict[str, dict[str, object]] = {} pred_paths = sorted(src.glob("*_pred.npy")) if not pred_paths: raise FileNotFoundError(f"No *_pred.npy inference results found under {src}") for pred_path in pred_paths: event = pred_path.name.removesuffix("_pred.npy") target_path = src / f"{event}_target.npy" if not target_path.is_file(): raise FileNotFoundError( f"Truth file not found: {target_path}. Rerun scripts/inference.py to export targets." ) prediction = np.load(pred_path) truth = np.load(target_path) if prediction.shape != truth.shape: raise ValueError(f"Prediction shape {prediction.shape} != truth shape {truth.shape} for {event}") if prediction.ndim != 3 or prediction.shape[0] != expected_frames: raise ValueError( f"Expected {expected_frames} frames [T,H,W] for {event}, got {prediction.shape}" ) absolute_error = np.abs(prediction - truth) mae_by_lead = absolute_error.mean(axis=(1, 2)) rmse_by_lead = np.sqrt(np.square(prediction - truth).mean(axis=(1, 2))) for index in range(expected_frames): filename = f"{event}_t{index + 1:02d}.png" write_png( prediction_dir / filename, colorize(prediction[index], RAIN_THRESHOLDS, RAIN_COLORS), ) write_png( comparison_dir / filename, comparison_image(truth[index], prediction[index]), ) summary[event] = { "shape": list(prediction.shape), "prediction_png_count": expected_frames, "comparison_png_count": expected_frames, "comparison_layout": ["truth", "prediction", "absolute_error"], "prediction_min": float(prediction.min()), "prediction_max": float(prediction.max()), "prediction_mean": float(prediction.mean()), "threshold": threshold, "threshold_fraction": float((prediction >= threshold).mean()), "mae": float(absolute_error.mean()), "rmse": float(np.sqrt(np.square(prediction - truth).mean())), "mae_by_lead": [float(value) for value in mae_by_lead], "rmse_by_lead": [float(value) for value in rmse_by_lead], } summary_path = out / "summary.json" summary_path.write_text(json.dumps(summary, indent=2) + "\n") print(f"prediction_png_dir={prediction_dir}") print(f"comparison_png_dir={comparison_dir}") print(f"summary={summary_path}") if __name__ == "__main__": main()