From 793b4a69f51a0a09569804331c940aae40427b26 Mon Sep 17 00:00:00 2001 From: Daewi Gao Date: Tue, 13 Jan 2026 12:14:23 -0800 Subject: [PATCH 1/2] Fix dimension mismatch bug in extract_elm_restart_point --- scripts/extract_elm_restart_point.py | 62 ++++++++++++++++++++++------ 1 file changed, 49 insertions(+), 13 deletions(-) diff --git a/scripts/extract_elm_restart_point.py b/scripts/extract_elm_restart_point.py index 3a5909a..ccc265c 100644 --- a/scripts/extract_elm_restart_point.py +++ b/scripts/extract_elm_restart_point.py @@ -10,6 +10,7 @@ import os import argparse from datetime import datetime +from netCDF4 import Dataset as NC4Dataset # ==================== Configuration Parameters (Default Values) ==================== SOURCE_NC = '/global/cfs/cdirs/m4814/daweigao/14_Code/all_dataset_1_degree/20250117_trendytest_ICB1850CNRDCTCBC_ad_spinup.elm.r.0021-01-01-00000.nc' @@ -94,11 +95,19 @@ def extract_single_point_elm(source_nc=None, output_nc=None, target_lat=None, ta # ============ Step 1: Open dataset ============ print("\n[1/5] Loading dataset...") + + # First, get ALL dimensions using netCDF4 (xarray may filter some out) + nc4_file = NC4Dataset(source_file, 'r') + all_nc_dims = {dim: len(nc4_file.dimensions[dim]) for dim in nc4_file.dimensions} + nc4_file.close() + ds = xr.open_dataset(source_file, decode_times=False, mask_and_scale=False) print(f" ✓ Dataset loaded successfully") print(f" - File size: {os.path.getsize(source_file) / (1024**3):.2f} GB") - print(f" - Dimensions: gridcell={ds.dims['gridcell']}, topounit={ds.dims['topounit']}, " + + print(f" - Dimensions (netCDF4): {len(all_nc_dims)} total") + print(f" - Dimensions (xarray): {len(ds.dims)} visible") + print(f" - Main dims: gridcell={ds.dims['gridcell']}, topounit={ds.dims['topounit']}, " + f"landunit={ds.dims['landunit']}, column={ds.dims['column']}, pft={ds.dims['pft']}") print(f" - Number of variables: {len(ds.data_vars)}") @@ -119,33 +128,33 @@ def extract_single_point_elm(source_nc=None, output_nc=None, target_lat=None, ta print(f" - All unique longitudes (original): {unique_lons}") # Handle longitude format conversion (0-360° vs -180 to 180°) - # 智能选择:如果数据集和目标都是同一格式,保持原格式;否则统一转换 + # Smart selection: if dataset and target use the same format, keep original format for efficiency dataset_uses_360 = lon_vals.max() > 200 - # 判断目标经度格式:> 180 肯定是 0-360,< 0 肯定是 -180~180,0-180 之间默认认为是 0-360 + # Determine target longitude format: > 180 must be 0-360, < 0 must be -180~180 target_uses_360 = lon > 180 or (lon >= 0 and lon <= 180 and dataset_uses_360) - # 如果数据集和目标都是 0-360 格式,保持原格式比较(更高效) + # If both dataset and target are in 0-360 format, keep original format (more efficient) if dataset_uses_360 and target_uses_360: - lon_adjusted = lon_vals.copy() # 保持 0-360 格式 - target_lon_adjusted = lon # 保持 0-360 格式 + lon_adjusted = lon_vals.copy() # Keep 0-360 format + target_lon_adjusted = lon # Keep 0-360 format print(f" Dataset uses 0-360° format, target also in 0-360° format ({lon:.6f})") print(f" → Using 0-360° format for comparison (no conversion needed)") - # 如果数据集是 0-360 但目标是 -180~180,转换数据集 + # If dataset is 0-360 but target is -180~180, convert dataset elif dataset_uses_360 and not target_uses_360: lon_adjusted = np.where(lon_vals > 180, lon_vals - 360, lon_vals) - target_lon_adjusted = lon # 目标已经是 -180~180 格式 + target_lon_adjusted = lon # Target is already in -180~180 format print(f" Dataset uses 0-360° format, target uses -180~180° format") print(f" → Converting dataset to -180~180° format for comparison") unique_lons_adjusted = np.unique(lon_adjusted) print(f" - Unique longitudes (adjusted): {len(unique_lons_adjusted)} values, range [{unique_lons_adjusted.min():.6f}, {unique_lons_adjusted.max():.6f}]") print(f" - All unique longitudes (adjusted): {unique_lons_adjusted}") - # 如果数据集是 -180~180 但目标是 0-360,转换目标 + # If dataset is -180~180 but target is 0-360, convert target elif not dataset_uses_360 and target_uses_360: - lon_adjusted = lon_vals.copy() # 数据集已经是 -180~180 格式 - target_lon_adjusted = lon - 360 if lon > 180 else lon # 转换目标到 -180~180 + lon_adjusted = lon_vals.copy() # Dataset is already in -180~180 format + target_lon_adjusted = lon - 360 if lon > 180 else lon # Convert target to -180~180 print(f" Dataset uses -180~180° format, target uses 0-360° format ({lon:.6f})") print(f" → Converting target to -180~180° format ({target_lon_adjusted:.6f}) for comparison") - # 如果两者都是 -180~180 格式,直接使用 + # If both use -180~180 format, use directly else: lon_adjusted = lon_vals.copy() target_lon_adjusted = lon @@ -212,12 +221,27 @@ def extract_single_point_elm(source_nc=None, output_nc=None, target_lat=None, ta # Filter data by level subset = ds.copy() - # Filter each dimension + # Identify global dimensions (dimensions that don't vary with gridcell) + # These are common vertical layers or other global dimensions in ELM restart files + all_dims = set(ds.dims.keys()) + spatial_dims = set(indices.keys()) + global_dims = all_dims - spatial_dims + + print(f" - Spatial dimensions to filter: {spatial_dims}") + print(f" - Global dimensions to preserve: {global_dims}") + + # Filter each spatial dimension (only filter spatial dimensions, preserve global dimensions) for dim_name, dim_indices in indices.items(): if dim_name in subset.dims and len(dim_indices) > 0: subset = subset.isel({dim_name: dim_indices}) print(f" ✓ Filtered {dim_name}: {len(dim_indices)} elements") + # Verify that global dimensions are preserved + print(f" ✓ Preserved global dimensions:") + for dim in global_dims: + if dim in subset.dims: + print(f" - {dim}: {subset.dims[dim]} elements") + # Calculate data compression ratio original_size_estimate = sum([ ds.dims['gridcell'], @@ -276,6 +300,18 @@ def extract_single_point_elm(source_nc=None, output_nc=None, target_lat=None, ta unlimited_dims=None ) + # Add missing dimensions that were in original file but filtered by xarray + missing_dims = set(all_nc_dims.keys()) - set(subset.dims.keys()) + if missing_dims: + print(f" Adding {len(missing_dims)} missing dimensions from original file...") + nc_out = NC4Dataset(output_file, 'a') + for dim_name in missing_dims: + if dim_name not in nc_out.dimensions: + dim_size = all_nc_dims[dim_name] + nc_out.createDimension(dim_name, dim_size) + print(f" + {dim_name}: {dim_size}") + nc_out.close() + # Verify output output_size = os.path.getsize(output_file) / (1024**2) # MB print(f" ✓ File saved successfully") From 3ad7374d11feb3aa6e9cdd098085cbf7d58967d7 Mon Sep 17 00:00:00 2001 From: Daewi Gao Date: Thu, 15 Jan 2026 14:25:51 -0800 Subject: [PATCH 2/2] Add IJCAI-style NEE time series plotting script --- scripts/plot_nee_timeseries_ijcai.py | 130 +++++++++++++++++++++++++++ 1 file changed, 130 insertions(+) create mode 100644 scripts/plot_nee_timeseries_ijcai.py diff --git a/scripts/plot_nee_timeseries_ijcai.py b/scripts/plot_nee_timeseries_ijcai.py new file mode 100644 index 0000000..2c1f68a --- /dev/null +++ b/scripts/plot_nee_timeseries_ijcai.py @@ -0,0 +1,130 @@ +import argparse +import re +from pathlib import Path + +import numpy as np +import xarray as xr +import matplotlib.pyplot as plt + + +DEFAULT_SIM_DIR = ( + "/global/cfs/cdirs/m4814/daweigao/14_Code/0_dataset_construction/" + "3_restarted simulation" +) +DEFAULT_GT_PATH = ( + "/global/cfs/cdirs/m4814/daweigao/14_Code/0_dataset_construction/" + "20251201_TRENDY2024_default_ICB1850CNPRDCTCBC.elm.h0.0801-01-01-00000.nc" +) + + +def _to_nan_fillvalue(arr, fill_threshold=1e35): + data = np.asarray(arr, dtype=float) + data[np.abs(data) >= fill_threshold] = np.nan + return data + + +def _apply_ijcai_style(): + plt.rcParams.update({ + "font.family": "serif", + "font.serif": ["Times New Roman", "Times", "DejaVu Serif"], + "font.size": 11, + "axes.labelsize": 12, + "axes.titlesize": 12, + "axes.linewidth": 1.0, + "xtick.direction": "in", + "ytick.direction": "in", + "xtick.major.size": 4, + "ytick.major.size": 4, + "xtick.minor.size": 2, + "ytick.minor.size": 2, + "legend.frameon": False, + }) + + +def _extract_year(path: Path): + match = re.search(r"\.h0\.(\d{4})-", path.name) + if not match: + return None + return int(match.group(1)) + + +def _sum_variable(ds, variable): + if variable not in ds: + raise KeyError(f"Variable {variable} not found in {ds.encoding.get('source', 'dataset')}.") + data = _to_nan_fillvalue(ds[variable].values) + return float(np.nansum(data)) + + +def _collect_simulation_series(sim_dir: Path, variable: str): + nc_files = sorted(sim_dir.glob("*.h0.*.nc")) + years = [] + sums = [] + for path in nc_files: + year = _extract_year(path) + if year is None: + continue + with xr.open_dataset(path) as ds: + value = _sum_variable(ds, variable) + years.append(year) + sums.append(value) + + if not years: + raise FileNotFoundError(f"No .h0.*.nc files with year found in {sim_dir}") + + order = np.argsort(years) + years = np.asarray(years)[order] + sums = np.asarray(sums)[order] + return years, sums + + +def plot_timeseries(sim_dir, gt_path, variable="NEE", output=None): + sim_dir = Path(sim_dir) + gt_path = Path(gt_path) + if output: + output = Path(output) + else: + output = Path(__file__).resolve().parent / f"{variable.lower()}_timeseries_ijcai.png" + + years, sim_sums = _collect_simulation_series(sim_dir, variable) + with xr.open_dataset(gt_path) as ds_gt: + gt_sum = _sum_variable(ds_gt, variable) + + _apply_ijcai_style() + fig, ax = plt.subplots(figsize=(6.4, 3.6)) + + ax.plot(years, sim_sums, color="#1f77b4", linewidth=2.0, label="Simulation") + ax.hlines(gt_sum, years.min(), years.max(), colors="#d62728", linestyles="--", linewidth=2.0, label="Ground Truth") + + ax.set_xlabel("Year") + ax.set_ylabel(f"{variable} Sum") + ax.grid(True, linestyle="--", linewidth=0.5, alpha=0.4) + ax.legend() + + fig.tight_layout() + output.parent.mkdir(parents=True, exist_ok=True) + fig.savefig(output, dpi=300) + pdf_output = output.with_suffix(".pdf") + fig.savefig(pdf_output) + plt.close(fig) + + print(f"Saved time series plot: {output}") + print(f"Saved time series plot: {pdf_output}") + + +def parse_args(): + parser = argparse.ArgumentParser(description="Plot IJCAI-style time series of NEE sum.") + parser.add_argument("--sim-dir", type=str, default=DEFAULT_SIM_DIR, help="Directory of simulation .h0.*.nc files.") + parser.add_argument("--gt-path", type=str, default=DEFAULT_GT_PATH, help="Ground-truth NetCDF path.") + parser.add_argument("--variable", type=str, default="NEE", help="Variable name to plot.") + parser.add_argument("--output", type=str, default=None, help="Output figure path.") + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + plot_timeseries( + sim_dir=args.sim_dir, + gt_path=args.gt_path, + variable=args.variable, + output=args.output, + )