"""Create diagnostic figures and metrics for virtual-data predictions.""" import argparse import csv import json from pathlib import Path import matplotlib import numpy as np matplotlib.use("Agg") import matplotlib.pyplot as plt def load_config(path: str) -> dict: import yaml with open(path, encoding="utf-8") as file: return yaml.safe_load(file) def load_variables(metadata_path: Path, channels: int) -> list[str]: if metadata_path.exists(): variables = json.loads(metadata_path.read_text(encoding="utf-8")).get("variables", []) if len(variables) == channels: return variables return [f"channel_{index}" for index in range(channels)] def compute_metrics(prediction: np.ndarray, target: np.ndarray, variables: list[str]) -> tuple[list[dict], dict]: channel_metrics = [] total_squared_error = total_absolute_error = total_error = 0.0 total_count = 0 sum_target = sum_prediction = sum_target_squared = sum_prediction_squared = sum_product = 0.0 for index, name in enumerate(variables): channel_target = target[:, index].astype(np.float64) channel_prediction = prediction[:, index].astype(np.float64) error = channel_prediction - channel_target count = error.size squared_error = float(np.sum(error**2)) absolute_error = float(np.sum(np.abs(error))) error_sum = float(np.sum(error)) target_sum = float(np.sum(channel_target)) prediction_sum = float(np.sum(channel_prediction)) target_squared = float(np.sum(channel_target**2)) prediction_squared = float(np.sum(channel_prediction**2)) product_sum = float(np.sum(channel_target * channel_prediction)) covariance = product_sum - target_sum * prediction_sum / count variance_target = target_squared - target_sum**2 / count variance_prediction = prediction_squared - prediction_sum**2 / count correlation = covariance / max(np.sqrt(variance_target * variance_prediction), 1e-12) channel_metrics.append({ "channel": index, "variable": name, "rmse": float(np.sqrt(squared_error / count)), "mae": absolute_error / count, "bias": error_sum / count, "correlation": float(correlation), }) total_squared_error += squared_error total_absolute_error += absolute_error total_error += error_sum total_count += count sum_target += target_sum sum_prediction += prediction_sum sum_target_squared += target_squared sum_prediction_squared += prediction_squared sum_product += product_sum covariance = sum_product - sum_target * sum_prediction / total_count variance_target = sum_target_squared - sum_target**2 / total_count variance_prediction = sum_prediction_squared - sum_prediction**2 / total_count overall = { "rmse": float(np.sqrt(total_squared_error / total_count)), "mae": total_absolute_error / total_count, "bias": total_error / total_count, "correlation": float(covariance / max(np.sqrt(variance_target * variance_prediction), 1e-12)), } return channel_metrics, overall def plot_diagnostics( prediction: np.ndarray, target: np.ndarray, variables: list[str], channel_metrics: list[dict], config: dict, output_dir: Path, ) -> Path: spec = config["visualization"] sample = int(spec["prediction_index"]) channel = int(spec["channel"]) if sample >= prediction.shape[0] or channel >= prediction.shape[1]: raise IndexError(f"Requested sample={sample}, channel={channel}, but prediction shape is {prediction.shape}") selected_target = target[sample, channel] selected_prediction = prediction[sample, channel] selected_error = selected_prediction - selected_target field_min = min(float(selected_target.min()), float(selected_prediction.min())) field_max = max(float(selected_target.max()), float(selected_prediction.max())) error_limit = max(float(np.abs(selected_error).max()), 1e-8) latitude = np.linspace(-89.5, 89.5, prediction.shape[2]) longitude = np.linspace(0.0, 360.0, prediction.shape[3], endpoint=False) figure = plt.figure(figsize=(16, 10), constrained_layout=True) grid = figure.add_gridspec(2, 3) extent = [longitude[0], longitude[-1], latitude[0], latitude[-1]] for axis, field, title in zip( [figure.add_subplot(grid[0, 0]), figure.add_subplot(grid[0, 1])], [selected_target, selected_prediction], ["Target field", "Predicted field"], ): image = axis.imshow(field, origin="lower", extent=extent, aspect="auto", cmap="viridis", vmin=field_min, vmax=field_max) axis.set_title(title) axis.set_xlabel("Longitude (degrees)") axis.set_ylabel("Latitude (degrees)") figure.colorbar(image, ax=axis, shrink=0.82) error_axis = figure.add_subplot(grid[0, 2]) image = error_axis.imshow(selected_error, origin="lower", extent=extent, aspect="auto", cmap="RdBu_r", vmin=-error_limit, vmax=error_limit) error_axis.set_title("Prediction error (prediction - target)") error_axis.set_xlabel("Longitude (degrees)") error_axis.set_ylabel("Latitude (degrees)") figure.colorbar(image, ax=error_axis, shrink=0.82) zonal_axis = figure.add_subplot(grid[1, 0]) zonal_axis.plot(selected_target.mean(axis=1), latitude, label="Target", linewidth=2) zonal_axis.plot(selected_prediction.mean(axis=1), latitude, label="Prediction", linewidth=2) zonal_axis.set_title("Zonal-mean profile") zonal_axis.set_xlabel("Zonal mean") zonal_axis.set_ylabel("Latitude (degrees)") zonal_axis.grid(alpha=0.25) zonal_axis.legend() scatter_axis = figure.add_subplot(grid[1, 1]) stride = max(1, selected_target.size // int(spec["scatter_points"])) x = selected_target.ravel()[::stride] y = selected_prediction.ravel()[::stride] scatter_axis.hexbin(x, y, gridsize=45, mincnt=1, cmap="magma") diagonal_min = min(float(x.min()), float(y.min())) diagonal_max = max(float(x.max()), float(y.max())) scatter_axis.plot([diagonal_min, diagonal_max], [diagonal_min, diagonal_max], "--", color="white", linewidth=1.5) scatter_axis.set_title("Pointwise agreement") scatter_axis.set_xlabel("Target") scatter_axis.set_ylabel("Prediction") sample_axis = figure.add_subplot(grid[1, 2]) sample_rmse = np.array([ np.sqrt(np.mean((prediction[index].astype(np.float64) - target[index]) ** 2)) for index in range(prediction.shape[0]) ]) sample_axis.bar(np.arange(len(sample_rmse)), sample_rmse, color="#2a6f97") sample_axis.axhline(sample_rmse.mean(), color="#d1495b", linestyle="--", label=f"Mean {sample_rmse.mean():.3f}") sample_axis.set_title("RMSE by sample") sample_axis.set_xlabel("Sample index") sample_axis.set_ylabel("RMSE") sample_axis.legend() metric = channel_metrics[channel] figure.suptitle( f"Virtual FV3GFS diagnostic | {variables[channel]} | sample {sample}\n" f"RMSE={metric['rmse']:.4f} MAE={metric['mae']:.4f} Bias={metric['bias']:.4f} Corr={metric['correlation']:.4f}", fontsize=15, ) path = output_dir / "diagnostic_dashboard.png" figure.savefig(path, dpi=int(spec["dpi"])) plt.close(figure) return path def plot_channel_metrics(channel_metrics: list[dict], output_dir: Path, dpi: int) -> Path: labels = [item["variable"] for item in channel_metrics] rmse = [item["rmse"] for item in channel_metrics] correlation = [item["correlation"] for item in channel_metrics] positions = np.arange(len(labels)) figure, axes = plt.subplots(1, 2, figsize=(16, 10), constrained_layout=True) axes[0].barh(positions, rmse, color="#457b9d") axes[0].set_title("RMSE by variable") axes[0].set_xlabel("RMSE") axes[1].barh(positions, correlation, color="#2a9d8f") axes[1].set_title("Correlation by variable") axes[1].set_xlabel("Pearson correlation") axes[1].set_xlim(-1, 1) for axis in axes: axis.set_yticks(positions, labels, fontsize=8) axis.invert_yaxis() axis.grid(axis="x", alpha=0.25) figure.suptitle("Virtual-data forecast skill by variable", fontsize=15) path = output_dir / "variable_metrics.png" figure.savefig(path, dpi=dpi) plt.close(figure) return path def create_report(config: dict) -> list[Path]: prediction_path = Path(config["inference"]["output_dir"]) / "prediction.npz" metadata_path = Path(config["synthetic_data"]["output_dir"]) / "metadata.json" if not prediction_path.exists(): raise FileNotFoundError(f"Prediction file not found: {prediction_path}. Run scripts/inference.py first.") with np.load(prediction_path) as arrays: prediction = arrays["prediction"] target = arrays["target"] if prediction.shape != target.shape or prediction.ndim != 4: raise ValueError(f"Expected matching [sample, channel, latitude, longitude] arrays, got {prediction.shape} and {target.shape}") variables = load_variables(metadata_path, prediction.shape[1]) channel_metrics, overall = compute_metrics(prediction, target, variables) output_dir = Path(config["visualization"]["output_dir"]) metrics_dir = Path(config["paths"]["metrics"]) output_dir.mkdir(parents=True, exist_ok=True) metrics_dir.mkdir(parents=True, exist_ok=True) summary_path = metrics_dir / "result_summary.json" summary_path.write_text(json.dumps({ "evaluation_scope": "virtual_data_only", "prediction_shape": list(prediction.shape), "overall": overall, "channels": channel_metrics, "note": "These metrics evaluate the synthetic task and are not paper reproduction metrics.", }, indent=2) + "\n", encoding="utf-8") csv_path = metrics_dir / "channel_metrics.csv" with csv_path.open("w", newline="", encoding="utf-8") as file: writer = csv.DictWriter(file, fieldnames=channel_metrics[0].keys()) writer.writeheader() writer.writerows(channel_metrics) dashboard = plot_diagnostics(prediction, target, variables, channel_metrics, config, output_dir) metric_plot = plot_channel_metrics(channel_metrics, output_dir, int(config["visualization"]["dpi"])) return [dashboard, metric_plot, summary_path, csv_path] def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--config", default="conf/config.yaml") args = parser.parse_args() for path in create_report(load_config(args.config)): print(f"result: {path}") if __name__ == "__main__": main()