CoolFace
Modelpublic

OneScience-Group/ML-MODIS

sourceHugging Faceapache-2.0updated 13d agoView on Hugging Face
0likes26downloads
result.py137 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""Evaluate OOB skill, 2014 cloud responses, importance and radiative contributions."""3 4from __future__ import annotations5 6import argparse7import json8import math9import sys10from pathlib import Path11 12import numpy as np13import torch14import yaml15 16import matplotlib17matplotlib.use("Agg")18import matplotlib.pyplot as plt19 20ROOT = Path(__file__).resolve().parents[1]21sys.path.insert(0, str(ROOT / "model"))22from ml_modis import BootstrapRandomForestRegressor, feature_names, regression_metrics23 24 25def finite(value: float):26    return float(value) if math.isfinite(float(value)) else None27 28 29def weighted_mean(values: np.ndarray, latitude: np.ndarray) -> float:30    valid = np.isfinite(values)31    weights = np.cos(np.deg2rad(latitude[valid]))32    return float(np.sum(values[valid] * weights) / np.sum(weights))33 34 35def main() -> None:36    parser = argparse.ArgumentParser()37    parser.add_argument("--config", default=str(ROOT / "conf/config.yaml"))38    parser.add_argument("--data", default=None)39    parser.add_argument("--checkpoint", default=None)40    parser.add_argument("--predictions", default=None)41    parser.add_argument("--output", default=None)42    parser.add_argument("--skip-importance", action="store_true")43    args = parser.parse_args()44    config = yaml.safe_load(Path(args.config).read_text())45    with np.load(ROOT / (args.data or config["data"]["path"])) as archive:46        data = {key: archive[key] for key in archive.files}47    with np.load(ROOT / (args.predictions or config["paths"]["predictions"])) as archive:48        predictions = {key: archive[key] for key in archive.files}49    checkpoint = torch.load(ROOT / (args.checkpoint or config["paths"]["checkpoint"]), map_location="cpu", weights_only=False)50    targets = list(checkpoint["model_config"]["targets"])51    report = {"format_version": config["format_version"],52              "evidence_scope": "Synthetic structured data smoke reproduction; not paper numerical results.",53              "oob": {}, "all_sample_skill": {}, "response_2014": {}, "susceptibility": {},54              "radiative_relative_contribution_percent": {}, "permutation_importance_top10": {}}55    names = feature_names()56    for month in checkpoint["model_config"]["months"]:57        for target_index, target in enumerate(targets):58            key = f"{month}:{target}"59            model_info = checkpoint["model"][key]60            report["oob"][key] = {metric: finite(value) if metric != "n" else int(value)61                                  for metric, value in model_info["oob_metrics"].items()}62            month_mask = data["month"] == month63            metrics = regression_metrics(predictions["obs"][month_mask, target_index], predictions["pred"][month_mask, target_index])64            report["all_sample_skill"][key] = {metric: finite(value) if metric != "n" else int(value) for metric, value in metrics.items()}65            if not args.skip_importance:66                train_mask = month_mask & (data["year"] != 2014)67                forest = BootstrapRandomForestRegressor.from_state_dict(model_info["state"])68                importance = forest.permutation_importance(data["X"][train_mask], data["Y"][train_mask, target_index], config["runtime"]["seed"] + target_index)69                order = np.argsort(importance)[::-1][:config["evaluation"]["importance_top_k"]]70                report["permutation_importance_top10"][key] = [71                    {"feature": names[index], "delta_oob_mse": float(importance[index])} for index in order72                ]73    eruption = predictions["year"] == 201474    monthly_log_response = {target: [] for target in targets}75    for month in checkpoint["model_config"]["months"]:76        mask = eruption & (predictions["month"] == month)77        for target_index, target in enumerate(targets):78            ratio = predictions["obs_over_pred"][mask, target_index]79            mean_ratio = weighted_mean(ratio, predictions["latitude"][mask])80            response = mean_ratio - 1.081            report["response_2014"][f"{month}:{target}"] = {82                "area_weighted_obs_over_pred": mean_ratio,83                "area_weighted_relative_percent": 100.0 * response,84                "samples": int(mask.sum()),85            }86            monthly_log_response[target].append(math.log(max(mean_ratio, 1e-8)))87    nd_change = float(np.mean(monthly_log_response["Nd"]))88    for target in ("reff", "LWP", "CF"):89        report["susceptibility"][f"dln{target}_dlnNd"] = finite(float(np.mean(monthly_log_response[target])) / nd_change)90 91    alpha_cloud = float(config["evaluation"]["cloud_albedo"])92    alpha_clear = float(config["evaluation"]["clear_sky_ocean_albedo"])93    s_lwp = report["susceptibility"]["dlnLWP_dlnNd"] or 0.094    s_cf = report["susceptibility"]["dlnCF_dlnNd"] or 0.095    terms = {96        "Twomey": alpha_cloud * (1 - alpha_cloud) / 3.0,97        "LWP": alpha_cloud * (1 - alpha_cloud) * (5.0 / 6.0) * s_lwp,98        "CF": (alpha_cloud - alpha_clear) * s_cf,99    }100    denominator = sum(terms.values())101    report["radiative_relative_contribution_percent"] = {102        key: finite(100.0 * value / denominator) for key, value in terms.items()103    }104    report["radiative_assumptions"] = {105        "cloud_albedo": alpha_cloud, "clear_sky_ocean_albedo": alpha_clear,106        "method": "Paper equations 1-3; common SWdown, CF and dlnNd/dlnAOD factors cancel in relative terms.",107        "twomey_note": "The 1/3 term follows the paper equation; observed dlnreff/dlnNd is reported separately."108    }109    output_dir = ROOT / config["paths"]["evaluation_dir"]110    output = ROOT / args.output if args.output else output_dir / "metrics.json"111    output_dir.mkdir(parents=True, exist_ok=True)112    output.parent.mkdir(parents=True, exist_ok=True)113    serialized = json.dumps(report, indent=2, allow_nan=False) + "\n"114    output.write_text(serialized)115    figure, axes = plt.subplots(1, 2, figsize=(12, 4.5))116    labels = [f"{month}-{target}" for month in checkpoint["model_config"]["months"] for target in targets]117    pearson = [report["all_sample_skill"][label.replace("-", ":")]["pearson"] for label in labels]118    axes[0].bar(labels, pearson, color=["#275d6c", "#d98b3a", "#6b8e23", "#8b5a83"] * 2)119    axes[0].set(ylabel="Pearson correlation", title="All-sample model skill")120    axes[0].tick_params(axis="x", rotation=45, labelsize=8)121    response_labels = [f"{month}-{target}" for month in checkpoint["model_config"]["months"] for target in targets]122    responses = [report["response_2014"][label.replace("-", ":")]["area_weighted_relative_percent"] for label in response_labels]123    axes[1].bar(response_labels, responses, color=["#275d6c", "#d98b3a", "#6b8e23", "#8b5a83"] * 2)124    axes[1].axhline(0, color="black", linewidth=0.7)125    axes[1].set(ylabel="Area-weighted response (%)", title="Observed / counterfactual in 2014")126    axes[1].tick_params(axis="x", rotation=45, labelsize=8)127    figure.tight_layout()128    figure.savefig(output_dir / "comparison.png", dpi=int(config["evaluation"]["figure_dpi"]))129    plt.close(figure)130    print(json.dumps({"output": str(output), "response_2014": report["response_2014"],131                      "susceptibility": report["susceptibility"],132                      "radiative_percent": report["radiative_relative_contribution_percent"]}, indent=2))133 134 135if __name__ == "__main__":136    main()137