ML-MODIS / scripts /result.py
zhangrenchao's picture
Publish ML-MODIS engineering reproduction
73ddb67 verified
Raw
History Blame Contribute Delete
7.36 kB
#!/usr/bin/env python3
"""Evaluate OOB skill, 2014 cloud responses, importance and radiative contributions."""
from __future__ import annotations
import argparse
import json
import math
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "model"))
from ml_modis import BootstrapRandomForestRegressor, feature_names, regression_metrics
def finite(value: float):
return float(value) if math.isfinite(float(value)) else None
def weighted_mean(values: np.ndarray, latitude: np.ndarray) -> float:
valid = np.isfinite(values)
weights = np.cos(np.deg2rad(latitude[valid]))
return float(np.sum(values[valid] * weights) / np.sum(weights))
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", default=str(ROOT / "conf/config.yaml"))
parser.add_argument("--data", default=None)
parser.add_argument("--checkpoint", default=None)
parser.add_argument("--predictions", default=None)
parser.add_argument("--output", default=None)
parser.add_argument("--skip-importance", action="store_true")
args = parser.parse_args()
config = yaml.safe_load(Path(args.config).read_text())
with np.load(ROOT / (args.data or config["data"]["path"])) as archive:
data = {key: archive[key] for key in archive.files}
with np.load(ROOT / (args.predictions or config["paths"]["predictions"])) as archive:
predictions = {key: archive[key] for key in archive.files}
checkpoint = torch.load(ROOT / (args.checkpoint or config["paths"]["checkpoint"]), map_location="cpu", weights_only=False)
targets = list(checkpoint["model_config"]["targets"])
report = {"format_version": config["format_version"],
"evidence_scope": "Synthetic structured data smoke reproduction; not paper numerical results.",
"oob": {}, "all_sample_skill": {}, "response_2014": {}, "susceptibility": {},
"radiative_relative_contribution_percent": {}, "permutation_importance_top10": {}}
names = feature_names()
for month in checkpoint["model_config"]["months"]:
for target_index, target in enumerate(targets):
key = f"{month}:{target}"
model_info = checkpoint["model"][key]
report["oob"][key] = {metric: finite(value) if metric != "n" else int(value)
for metric, value in model_info["oob_metrics"].items()}
month_mask = data["month"] == month
metrics = regression_metrics(predictions["obs"][month_mask, target_index], predictions["pred"][month_mask, target_index])
report["all_sample_skill"][key] = {metric: finite(value) if metric != "n" else int(value) for metric, value in metrics.items()}
if not args.skip_importance:
train_mask = month_mask & (data["year"] != 2014)
forest = BootstrapRandomForestRegressor.from_state_dict(model_info["state"])
importance = forest.permutation_importance(data["X"][train_mask], data["Y"][train_mask, target_index], config["runtime"]["seed"] + target_index)
order = np.argsort(importance)[::-1][:config["evaluation"]["importance_top_k"]]
report["permutation_importance_top10"][key] = [
{"feature": names[index], "delta_oob_mse": float(importance[index])} for index in order
]
eruption = predictions["year"] == 2014
monthly_log_response = {target: [] for target in targets}
for month in checkpoint["model_config"]["months"]:
mask = eruption & (predictions["month"] == month)
for target_index, target in enumerate(targets):
ratio = predictions["obs_over_pred"][mask, target_index]
mean_ratio = weighted_mean(ratio, predictions["latitude"][mask])
response = mean_ratio - 1.0
report["response_2014"][f"{month}:{target}"] = {
"area_weighted_obs_over_pred": mean_ratio,
"area_weighted_relative_percent": 100.0 * response,
"samples": int(mask.sum()),
}
monthly_log_response[target].append(math.log(max(mean_ratio, 1e-8)))
nd_change = float(np.mean(monthly_log_response["Nd"]))
for target in ("reff", "LWP", "CF"):
report["susceptibility"][f"dln{target}_dlnNd"] = finite(float(np.mean(monthly_log_response[target])) / nd_change)
alpha_cloud = float(config["evaluation"]["cloud_albedo"])
alpha_clear = float(config["evaluation"]["clear_sky_ocean_albedo"])
s_lwp = report["susceptibility"]["dlnLWP_dlnNd"] or 0.0
s_cf = report["susceptibility"]["dlnCF_dlnNd"] or 0.0
terms = {
"Twomey": alpha_cloud * (1 - alpha_cloud) / 3.0,
"LWP": alpha_cloud * (1 - alpha_cloud) * (5.0 / 6.0) * s_lwp,
"CF": (alpha_cloud - alpha_clear) * s_cf,
}
denominator = sum(terms.values())
report["radiative_relative_contribution_percent"] = {
key: finite(100.0 * value / denominator) for key, value in terms.items()
}
report["radiative_assumptions"] = {
"cloud_albedo": alpha_cloud, "clear_sky_ocean_albedo": alpha_clear,
"method": "Paper equations 1-3; common SWdown, CF and dlnNd/dlnAOD factors cancel in relative terms.",
"twomey_note": "The 1/3 term follows the paper equation; observed dlnreff/dlnNd is reported separately."
}
output_dir = ROOT / config["paths"]["evaluation_dir"]
output = ROOT / args.output if args.output else output_dir / "metrics.json"
output_dir.mkdir(parents=True, exist_ok=True)
output.parent.mkdir(parents=True, exist_ok=True)
serialized = json.dumps(report, indent=2, allow_nan=False) + "\n"
output.write_text(serialized)
figure, axes = plt.subplots(1, 2, figsize=(12, 4.5))
labels = [f"{month}-{target}" for month in checkpoint["model_config"]["months"] for target in targets]
pearson = [report["all_sample_skill"][label.replace("-", ":")]["pearson"] for label in labels]
axes[0].bar(labels, pearson, color=["#275d6c", "#d98b3a", "#6b8e23", "#8b5a83"] * 2)
axes[0].set(ylabel="Pearson correlation", title="All-sample model skill")
axes[0].tick_params(axis="x", rotation=45, labelsize=8)
response_labels = [f"{month}-{target}" for month in checkpoint["model_config"]["months"] for target in targets]
responses = [report["response_2014"][label.replace("-", ":")]["area_weighted_relative_percent"] for label in response_labels]
axes[1].bar(response_labels, responses, color=["#275d6c", "#d98b3a", "#6b8e23", "#8b5a83"] * 2)
axes[1].axhline(0, color="black", linewidth=0.7)
axes[1].set(ylabel="Area-weighted response (%)", title="Observed / counterfactual in 2014")
axes[1].tick_params(axis="x", rotation=45, labelsize=8)
figure.tight_layout()
figure.savefig(output_dir / "comparison.png", dpi=int(config["evaluation"]["figure_dpi"]))
plt.close(figure)
print(json.dumps({"output": str(output), "response_2014": report["response_2014"],
"susceptibility": report["susceptibility"],
"radiative_percent": report["radiative_relative_contribution_percent"]}, indent=2))
if __name__ == "__main__":
main()