File size: 8,352 Bytes
80cf062 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 | from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Any
import numpy as np
PROJECT_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_RESULTS_DIR = PROJECT_ROOT / "outputs" / "inference"
def _load_array(path: Path, name: str) -> np.ndarray:
value = np.load(path)[name]
if not np.all(np.isfinite(value)):
raise ValueError(f"{path.name}:{name} contains NaN or Inf values.")
return value
def validate_sample(path: Path, expected_channels: int = 20) -> dict[str, Any]:
with np.load(path) as result:
required = {"input", "reconstruction", "error", "mask", "step_idx", "time_index"}
missing = required - set(result.files)
if missing:
raise ValueError(f"{path} is missing fields: {sorted(missing)}")
inputs = _load_array(path, "input")
reconstruction = _load_array(path, "reconstruction")
error = _load_array(path, "error")
mask = _load_array(path, "mask")
expected_shape = (expected_channels, 720, 1440)
if inputs.shape != expected_shape:
raise ValueError(f"input shape must be {expected_shape}, got {inputs.shape}.")
if reconstruction.shape != inputs.shape or error.shape != inputs.shape:
raise ValueError("reconstruction and error must have the same shape as input.")
if mask.shape != (90, 180):
raise ValueError(f"mask must be a patch grid with shape (90, 180), got {mask.shape}.")
if not np.array_equal(mask, mask.astype(bool)):
raise ValueError("mask must contain only binary values.")
if not np.allclose(error, reconstruction - inputs):
raise ValueError("error does not equal reconstruction - input.")
spatial_mask = np.repeat(np.repeat(mask.astype(bool), 8, axis=0), 8, axis=1)
visible_error = np.abs(error[:, ~spatial_mask])
masked_error = np.abs(error[:, spatial_mask])
mask_fraction = float(mask.mean())
return {
"file": str(path),
"shape": list(inputs.shape),
"mask_fraction": mask_fraction,
"masked_pixels": int(spatial_mask.sum()),
"visible_mae": np.mean(visible_error, axis=1).tolist() if visible_error.size else [0.0] * expected_channels,
"masked_mae": np.mean(masked_error, axis=1).tolist() if masked_error.size else [0.0] * expected_channels,
"mae": np.mean(np.abs(error), axis=(1, 2)).tolist(),
"rmse": np.sqrt(np.mean(error**2, axis=(1, 2))).tolist(),
}
def save_diagnostic_png(path: Path, output_path: Path, channel: int) -> None:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
with np.load(path) as result:
input_field = result["input"][channel]
reconstruction = result["reconstruction"][channel]
error = result["error"][channel]
mask = result["mask"].astype(bool)
time_index = result["time_index"].item()
input_low, input_high = np.nanpercentile(input_field, [1, 99])
reconstruction_low, reconstruction_high = np.nanpercentile(reconstruction, [1, 99])
field_low = min(input_low, reconstruction_low)
field_high = max(input_high, reconstruction_high)
if not np.isfinite(field_low) or not np.isfinite(field_high) or field_high <= field_low:
field_low, field_high = float(np.nanmin(input_field)), float(np.nanmax(input_field))
error_limit = float(np.nanpercentile(np.abs(error), 99))
if not np.isfinite(error_limit) or error_limit <= 0:
error_limit = max(float(np.nanmax(np.abs(error))), 1.0)
latitude = np.linspace(90, -90, input_field.shape[0])
longitude = np.linspace(0, 360, input_field.shape[1], endpoint=False)
extent = [longitude[0], longitude[-1], latitude[-1], latitude[0]]
masked_pixels = np.repeat(np.repeat(mask, 8, axis=0), 8, axis=1)
masked_overlay = np.ma.masked_where(~masked_pixels, masked_pixels)
figure, axes = plt.subplots(2, 3, figsize=(18, 8), constrained_layout=True)
field_kwargs = {"cmap": "viridis", "vmin": field_low, "vmax": field_high, "extent": extent, "aspect": "auto"}
error_kwargs = {"cmap": "RdBu_r", "vmin": -error_limit, "vmax": error_limit, "extent": extent, "aspect": "auto"}
panels = [
(axes[0, 0], input_field, "Input (normalized)", field_kwargs),
(axes[0, 1], reconstruction, "Reconstruction (normalized)", field_kwargs),
(axes[0, 2], error, "Signed error", error_kwargs),
(axes[1, 0], np.abs(error), "Absolute error", {**error_kwargs, "cmap": "magma", "vmin": 0}),
]
for axis, field, title, kwargs in panels:
image = axis.imshow(field, **kwargs)
axis.set_title(title)
axis.set_xlabel("Longitude (degrees east)")
axis.set_ylabel("Latitude (degrees)")
axis.contour(masked_pixels, levels=[0.5], colors="white", linewidths=0.35, extent=extent)
figure.colorbar(image, ax=axis, shrink=0.8)
mask_image = axes[1, 1].imshow(mask.astype(float), cmap="gray_r", vmin=0, vmax=1, extent=extent, aspect="auto")
axes[1, 1].set_title("Patch mask (white = masked)")
axes[1, 1].set_xlabel("Longitude (degrees east)")
axes[1, 1].set_ylabel("Latitude (degrees)")
figure.colorbar(mask_image, ax=axes[1, 1], ticks=[0, 1], shrink=0.8)
axes[1, 2].hist(error.ravel(), bins=80, color="#315f8c", alpha=0.85)
axes[1, 2].axvline(0, color="black", linewidth=0.8)
axes[1, 2].set_title(f"Error distribution\nMAE={np.mean(np.abs(error)):.4f}, RMSE={np.sqrt(np.mean(error**2)):.4f}")
axes[1, 2].set_xlabel("Signed error (normalized)")
axes[1, 2].set_ylabel("Pixel count")
figure.suptitle(f"W-MAE reconstruction diagnostic | channel={channel} | time={time_index}")
output_path.parent.mkdir(parents=True, exist_ok=True)
figure.savefig(output_path, dpi=150)
plt.close(figure)
def main() -> None:
parser = argparse.ArgumentParser(description="Validate W-MAE inference outputs and make diagnostics.")
parser.add_argument(
"results_dir",
type=Path,
nargs="?",
default=DEFAULT_RESULTS_DIR,
help="Inference output directory (default: outputs/inference).",
)
parser.add_argument("--expected-channels", type=int, default=20)
parser.add_argument("--mask-ratio", type=float, default=None)
parser.add_argument(
"--channel",
type=int,
default=0,
help="Channel for the diagnostic image (default: 0). Use --channel -1 to disable it.",
)
parser.add_argument("--diagnostic-dir", type=Path, default=None)
parser.add_argument(
"--visualization-samples",
type=int,
default=1,
help="Number of samples to visualize; validation still covers all samples (default: 1).",
)
args = parser.parse_args()
if args.visualization_samples <= 0:
raise ValueError("--visualization-samples must be positive.")
files = sorted(args.results_dir.glob("sample_*.npz"))
if not files:
raise FileNotFoundError(f"No sample_*.npz files found in {args.results_dir}.")
reports = [validate_sample(path, args.expected_channels) for path in files]
if args.mask_ratio is not None:
for report in reports:
if not np.isclose(report["mask_fraction"], args.mask_ratio):
raise ValueError(
f"{report['file']} mask fraction {report['mask_fraction']} "
f"does not match expected {args.mask_ratio}."
)
if args.channel is not None and args.channel >= 0:
if not 0 <= args.channel < args.expected_channels:
raise ValueError("--channel is outside the configured channel range.")
diagnostic_dir = args.diagnostic_dir or args.results_dir / "diagnostics"
for path in files[:args.visualization_samples]:
save_diagnostic_png(path, diagnostic_dir / f"{path.stem}_channel_{args.channel:02d}.png", args.channel)
summary = {
"samples": len(reports),
"mean_mask_fraction": float(np.mean([report["mask_fraction"] for report in reports])),
"mean_mae": np.mean([report["mae"] for report in reports], axis=0).tolist(),
"mean_rmse": np.mean([report["rmse"] for report in reports], axis=0).tolist(),
"reports": reports,
}
print(json.dumps(summary, indent=2), flush=True)
if __name__ == "__main__":
main()
|