Download scripts/result.py from OneScience-Group/Spherical_DYffusion: direct link, hf CLI and curl.
- Browser
- Download file 10.7 kB
-
https://huggingface.co/OneScience-Group/Spherical_DYffusion/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/Spherical_DYffusion/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/Spherical_DYffusion/resolve/main/scripts/result.py
10.7 kB
| """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() | |