Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 49 additions & 13 deletions scripts/extract_elm_restart_point.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'
Expand Down Expand Up @@ -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)}")

Expand All @@ -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
Expand Down Expand Up @@ -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'],
Expand Down Expand Up @@ -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")
Expand Down
130 changes: 130 additions & 0 deletions scripts/plot_nee_timeseries_ijcai.py
Original file line number Diff line number Diff line change
@@ -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,
)