Earthformer / script /result.py
yzt15806542928's picture
Upload folder using huggingface_hub
8bfc737 verified
Raw
History Blame Contribute Delete
7.1 kB
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
if __package__ in (None, ""):
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from model.earthformer import Earthformer
from script.data_loader import make_loader
from script.metrics import metric_sums, metrics_from_sums
from script.utils import clean_state_dict, load_checkpoint_payload, resolve_cli_path, resolve_device
def _squeeze(frame: np.ndarray) -> np.ndarray:
return np.squeeze(frame)
def plot_lead_time(truth: np.ndarray, prediction: np.ndarray, output_path: Path, title: str) -> Path:
"""Three-panel Truth / Prediction / Difference image, fuxi/fengwu style."""
truth, prediction = _squeeze(truth), _squeeze(prediction)
difference = truth - prediction
rmse = float(np.sqrt(np.mean(difference**2)))
vmin, vmax = min(truth.min(), prediction.min()), max(truth.max(), prediction.max())
diff_abs_max = max(float(np.abs(difference).max()), 1e-8)
panels = [
{"data": truth, "title": "Truth", "cmap": "viridis", "vmin": vmin, "vmax": vmax},
{"data": prediction, "title": "Prediction", "cmap": "viridis", "vmin": vmin, "vmax": vmax},
{"data": difference, "title": f"Difference (RMSE={rmse:.2f})", "cmap": "RdBu_r", "vmin": -diff_abs_max, "vmax": diff_abs_max},
]
fig, axes = plt.subplots(1, 3, figsize=(15, 4))
for ax, panel in zip(axes, panels):
image = ax.imshow(panel["data"], cmap=panel["cmap"], vmin=panel["vmin"], vmax=panel["vmax"])
ax.set_title(panel["title"], fontsize=12, pad=4)
ax.set_xlabel("Pixel")
ax.set_ylabel("Pixel")
plt.colorbar(image, ax=ax, orientation="horizontal")
fig.suptitle(title, fontsize=14, fontweight="bold", y=0.98)
fig.savefig(output_path, dpi=300, bbox_inches="tight")
plt.close(fig)
return output_path
def visualize_predictions(
truth: np.ndarray,
prediction: np.ndarray,
output_dir: str | Path,
max_samples: int = 2,
stride: int = 2,
) -> list[str]:
"""Write per-sample/lead-time PNG images and a per-sample mosaic into output_dir."""
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
if truth.shape != prediction.shape:
raise ValueError(f"truth and prediction must have the same shape, got {truth.shape} vs {prediction.shape}")
samples = min(max_samples, int(truth.shape[0]))
lead_times = list(range(0, int(truth.shape[1]), stride))
generated: list[str] = []
for sample in range(samples):
for lead_time in lead_times:
path = output_dir / f"earthformer_sample_{sample:03d}_t{lead_time:02d}.png"
plot_lead_time(
truth[sample, lead_time],
prediction[sample, lead_time],
path,
f"Earthformer VIL sample {sample} lead time {lead_time} (+{lead_time * 5} min)",
)
generated.append(str(path))
mosaic = output_dir / f"earthformer_overview_sample_{sample:03d}.png"
_plot_mosaic(truth[sample], prediction[sample], lead_times, mosaic)
generated.append(str(mosaic))
return generated
def _plot_mosaic(truth: np.ndarray, prediction: np.ndarray, lead_times: list[int], output_path: Path) -> Path:
rows, columns = len(lead_times) * 2, len(lead_times)
fig, axes = plt.subplots(rows, columns, figsize=(columns * 2.6, rows * 2.2))
for column, lead_time in enumerate(lead_times):
truth_frame = _squeeze(truth[lead_time])
pred_frame = _squeeze(prediction[lead_time])
vmin, vmax = min(truth_frame.min(), pred_frame.min()), max(truth_frame.max(), pred_frame.max())
axes[0, column].imshow(truth_frame, cmap="viridis", vmin=vmin, vmax=vmax)
axes[0, column].set_title(f"+{lead_time * 5} min", fontsize=9)
axes[len(lead_times), column].imshow(pred_frame, cmap="viridis", vmin=vmin, vmax=vmax)
for row in range(rows):
axes[row, column].set_xticks([])
axes[row, column].set_yticks([])
axes[0, 0].set_ylabel("Truth", fontsize=10)
axes[len(lead_times), 0].set_ylabel("Prediction", fontsize=10)
fig.suptitle("Earthformer VIL sample overview", fontsize=13, fontweight="bold", y=0.99)
fig.tight_layout(rect=(0, 0, 1, 0.97))
fig.savefig(output_path, dpi=200, bbox_inches="tight")
plt.close(fig)
return output_path
def evaluate(model: torch.nn.Module, loader, device: torch.device) -> dict[str, float]:
sums = torch.zeros(22, dtype=torch.float64, device=device)
with torch.no_grad():
for inputs, targets in loader:
prediction = model(inputs.to(device)).clamp(0.0, 1.0)
sums += metric_sums(prediction, targets.to(device))
result = metrics_from_sums(sums.cpu())
result["note"] = "Lightweight metrics on configured data; mean CSI is not the official complete SEVIR evaluation"
return result
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Evaluate Earthformer and visualize predictions (fuxi/fengwu style PNG images)"
)
parser.add_argument("--checkpoint", default="data/checkpoint/earthformer.pt")
parser.add_argument("--split", choices=("train", "val", "test"), default="test")
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--predictions", default="output/predictions.npz", help="Inference arrays for visualization")
parser.add_argument("--output-dir", default="output/visualization", help="Directory for visualization PNG images")
parser.add_argument("--max-samples", type=int, default=2)
parser.add_argument("--stride", type=int, default=2, help="Lead-time stride, aligned with official plot_stride")
parser.add_argument("--skip-visualization", action="store_true", help="Only evaluate, do not render PNG images")
return parser.parse_args()
def main() -> None:
args = parse_args()
device = resolve_device(args.device)
payload = load_checkpoint_payload(resolve_cli_path(args.checkpoint), device)
config = payload["config"]
model = Earthformer(config).to(device)
model.load_state_dict(clean_state_dict(payload["model"]))
model.eval()
loader, _ = make_loader(config, args.split, shuffle=False)
result = evaluate(model, loader, device)
report = {"metrics": result}
if args.skip_visualization:
print(json.dumps(report, indent=2))
return
predictions_path = Path(resolve_cli_path(args.predictions))
with np.load(predictions_path) as payload:
truth = payload["targets"]
prediction = payload["predictions"]
output_dir = resolve_cli_path(args.output_dir)
images = visualize_predictions(truth, prediction, output_dir, args.max_samples, args.stride)
report["output_dir"] = output_dir
report["images"] = images
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()