yzt15806542928's picture
Upload folder using huggingface_hub
6f3c6ef verified
Raw
History Blame Contribute Delete
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()