yzt15806542928's picture
Upload folder using huggingface_hub
439c523 verified
Raw
History Blame Contribute Delete
6.32 kB
#!/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()