From 4f7c59841cc0b8a88e811c89b7685a4d5583b226 Mon Sep 17 00:00:00 2001 From: SarahAlidoost Date: Tue, 15 Sep 2026 11:10:44 +0200 Subject: [PATCH 1/4] add prediction scripts --- scripts/prediction.py | 116 +++++++++++++++++++++++++++++++++++++++ scripts/prediction.slurm | 37 +++++++++++++ 2 files changed, 153 insertions(+) create mode 100644 scripts/prediction.py create mode 100644 scripts/prediction.slurm diff --git a/scripts/prediction.py b/scripts/prediction.py new file mode 100644 index 0000000..d33a7b5 --- /dev/null +++ b/scripts/prediction.py @@ -0,0 +1,116 @@ +import argparse +from pathlib import Path + +import ray +import xarray as xr + +from climanet.dataset import DataLoaderConfig, STDataset +from climanet.predict import PredictionConfig, predict_monthly_var +from climanet.utils import configure_compute_resources, load_model, read_st_data, set_seed + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument( + "--run-dir", + type=str, + default=Path("./run_dir").resolve(), + ) + parser.add_argument( + "--prepared-data-dir", + type=str, + default=Path("./data").resolve(), + ) + parser.add_argument( + "--train-dir", + type=str, + default=Path("./data").resolve(), + ) + parser.add_argument( + "--lsm-dir", + type=str, + default=Path("./data").resolve(), + ) + args = parser.parse_args() + + var_name = "tos" + device = "cuda" + prepared_data_dir = Path(args.prepared_data_dir).resolve() + lsm_dir = Path(args.lsm_dir).resolve() + train_dir = Path(args.train_dir).resolve() + run_dir = Path(args.run_dir).resolve() + + # set the random seed for reproducibility + set_seed() + + # Load the trained model for prediction + model_path = train_dir / "best_model.pth" + model = load_model(model_path, device) + model_patch_size = model.config["patch_size"] + + # Build dataset for training and validation + lsm_file_path = lsm_dir / "era5_lsm_bool.nc" + lsm_mask = xr.open_dataset(lsm_file_path)["lsm"] # make sure is dask array + + predict_year = 2022 + data = read_st_data(data_path=f"{prepared_data_dir}/{predict_year}", var_name=var_name) + input_da, input_da_nan_mask, monthly_da, padded_days_mask, time_features = zip(*data) + + monthly_shape = monthly_da.shape[1:] # the whole dataset + crop_size = (1, *monthly_shape) + + dataset_test = STDataset( + input_da=input_da, + input_da_nan_mask=input_da_nan_mask, + monthly_da=monthly_da, + padded_days_mask=padded_days_mask, + time_features=time_features, + land_mask=lsm_mask, + crop_size=crop_size, + stride=None, # no overlap for prediction + model_patch_size=model_patch_size, + sh_embed_dim=96, + sh_order_L=10, + verbose=False, + load_lazy=False, # load all data into memory for prediction 1 year + ) + + # Build the dataloader config + dataloader_num_workers = 10 # adjust if needed + use_cuda = device == "cuda" + dataloader_config = DataLoaderConfig( + batch_size=100, # adjust if OOM issue + shuffle=False, # for prediction and reconstruction, it should be False + num_workers=dataloader_num_workers, + pin_memory=use_cuda, + persistent_workers=True, + device=device, + multiprocessing_context=None, # one year fits in memory + ) + + # move the model to GPU and configure compute resources + model = configure_compute_resources( + model, + device=device, + compute_threads=None, # on gpu, it is not used + dataloader_num_workers=dataloader_num_workers + ) + + # Training configuration# create prediction config + prediction_config = PredictionConfig( + calculate_residuals=True, + return_numpy=False, + save_predictions=True, + return_loss=False, + device=device, + verbose=False, + store_logs=False, + ) + + predictions = predict_monthly_var( + model=model, + dataset=dataset_test, + dataloader_config=dataloader_config, + prediction_config=prediction_config, + run_dir=run_dir, + ) diff --git a/scripts/prediction.slurm b/scripts/prediction.slurm new file mode 100644 index 0000000..8c6e3d5 --- /dev/null +++ b/scripts/prediction.slurm @@ -0,0 +1,37 @@ +#!/bin/bash +#SBATCH --job-name=prediction +#SBATCH --partition=gpu +#SBATCH --constraint=a100_80 +#SBATCH --nodes=1 +#SBATCH --ntasks-per-node=1 +#SBATCH --cpus-per-task=128 +#SBATCH --gpus-per-task=4 +#SBATCH --exclusive +#SBATCH --mem=0 +#SBATCH --time=12:00:00 +#SBATCH --account=bd0854 +#SBATCH --output=prediction_%j.out + +set -euo pipefail +ulimit -s 204800 + +# Activate uv env +UV_ENV="$HOME/climanet_py314" +source "$UV_ENV/bin/activate" + +# Set the scratch directory because they are avialable to all nodes +RUN_DIR="/scratch/b/$USER/predict" + +# data directory (adjust this path to your data location) +PREPARED_DATA_DIR="/scratch/b/$USER/data" +TRAIN_DIR="/scratch/b/$USER/train" +LSM_DIR="/scratch/b/$USER/data" + +echo "Starting prediction.py script..." +python -u $HOME/ClimaNet/scripts/prediction.py \ + --run-dir "$RUN_DIR" \ + --prepared-data-dir "$PREPARED_DATA_DIR" \ + --train-dir "$TRAIN_DIR" \ + --lsm-dir "$LSM_DIR" + +echo "**********Prediction script completed.************" From be28104520ea366f8b07946a0523f0d994da8fbe Mon Sep 17 00:00:00 2001 From: SarahAlidoost Date: Tue, 15 Sep 2026 11:33:20 +0200 Subject: [PATCH 2/4] fix prediction scripts --- scripts/prediction.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/prediction.py b/scripts/prediction.py index d33a7b5..3f01db3 100644 --- a/scripts/prediction.py +++ b/scripts/prediction.py @@ -54,7 +54,7 @@ predict_year = 2022 data = read_st_data(data_path=f"{prepared_data_dir}/{predict_year}", var_name=var_name) - input_da, input_da_nan_mask, monthly_da, padded_days_mask, time_features = zip(*data) + input_da, input_da_nan_mask, monthly_da, padded_days_mask, time_features = data monthly_shape = monthly_da.shape[1:] # the whole dataset crop_size = (1, *monthly_shape) From 7159d080bed340839f3f6a3ce125b034a19a6acd Mon Sep 17 00:00:00 2001 From: SarahAlidoost Date: Tue, 15 Sep 2026 14:13:47 +0200 Subject: [PATCH 3/4] fix scripts --- scripts/prediction.py | 32 +++++++++++++++++++++++++++++++- scripts/prediction.slurm | 6 ++++-- 2 files changed, 35 insertions(+), 3 deletions(-) diff --git a/scripts/prediction.py b/scripts/prediction.py index 3f01db3..eaf5630 100644 --- a/scripts/prediction.py +++ b/scripts/prediction.py @@ -3,6 +3,7 @@ import ray import xarray as xr +import numpy as np from climanet.dataset import DataLoaderConfig, STDataset from climanet.predict import PredictionConfig, predict_monthly_var @@ -31,6 +32,11 @@ type=str, default=Path("./data").resolve(), ) + parser.add_argument( + "--raw-data-folder", + type=str, + default=Path("./data").resolve(), + ) args = parser.parse_args() var_name = "tos" @@ -39,6 +45,7 @@ lsm_dir = Path(args.lsm_dir).resolve() train_dir = Path(args.train_dir).resolve() run_dir = Path(args.run_dir).resolve() + raw_data_folder = Path(args.raw_data_folder).resolve() # set the random seed for reproducibility set_seed() @@ -79,7 +86,7 @@ dataloader_num_workers = 10 # adjust if needed use_cuda = device == "cuda" dataloader_config = DataLoaderConfig( - batch_size=100, # adjust if OOM issue + batch_size=1, # adjust if OOM issue, monthly batch shuffle=False, # for prediction and reconstruction, it should be False num_workers=dataloader_num_workers, pin_memory=use_cuda, @@ -107,6 +114,7 @@ store_logs=False, ) + # Predict residuals predictions = predict_monthly_var( model=model, dataset=dataset_test, @@ -114,3 +122,25 @@ prediction_config=prediction_config, run_dir=run_dir, ) + + # add residuals to the averaged monthly data + # load hourly data + files = sorted(raw_data_folder.glob(f"{predict_year}*_hr_ERA5dc_masked_{var_name}.nc")) + input_data = xr.open_mfdataset(files) + + # load prediction residuals + files = sorted(run_dir.glob(f"{predict_year}*_{var_name}_prediction_residual.nc")) + predictions_res = xr.open_mfdataset(files) + + # add residuals to the averaged monthly data + input_data_averaged = input_data.resample({"time": "MS"}).mean(skipna=True) + input_data_averaged["time"] = predictions_res["time"] + adjusted_data = input_data_averaged[var_name] + predictions_res + + # save the adjusted data to a new NetCDF file, one file per month + times = adjusted_data.coords["time"].values + + for t in times: + time_str = np.datetime_as_string(t, unit="M").replace("-", "") + file_name = f"{run_dir}/{time_str}_{var_name}_prediction.nc" + adjusted_data.sel(time=[t]).to_netcdf(file_name) diff --git a/scripts/prediction.slurm b/scripts/prediction.slurm index 8c6e3d5..3841a19 100644 --- a/scripts/prediction.slurm +++ b/scripts/prediction.slurm @@ -8,7 +8,7 @@ #SBATCH --gpus-per-task=4 #SBATCH --exclusive #SBATCH --mem=0 -#SBATCH --time=12:00:00 +#SBATCH --time=1:00:00 #SBATCH --account=bd0854 #SBATCH --output=prediction_%j.out @@ -26,12 +26,14 @@ RUN_DIR="/scratch/b/$USER/predict" PREPARED_DATA_DIR="/scratch/b/$USER/data" TRAIN_DIR="/scratch/b/$USER/train" LSM_DIR="/scratch/b/$USER/data" +RAW_DATA_FOLDER="/scratch/b/$USER/output/sst/concatenated" echo "Starting prediction.py script..." python -u $HOME/ClimaNet/scripts/prediction.py \ --run-dir "$RUN_DIR" \ --prepared-data-dir "$PREPARED_DATA_DIR" \ --train-dir "$TRAIN_DIR" \ - --lsm-dir "$LSM_DIR" + --lsm-dir "$LSM_DIR"\ + --raw-data-folder "$RAW_DATA_FOLDER" echo "**********Prediction script completed.************" From dcdb963dbda431f94107e6d0402aa162690ee875 Mon Sep 17 00:00:00 2001 From: SarahAlidoost Date: Tue, 15 Sep 2026 14:14:39 +0200 Subject: [PATCH 4/4] fix linter errors --- scripts/prediction.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/scripts/prediction.py b/scripts/prediction.py index eaf5630..b7a4235 100644 --- a/scripts/prediction.py +++ b/scripts/prediction.py @@ -1,14 +1,17 @@ import argparse from pathlib import Path -import ray -import xarray as xr import numpy as np +import xarray as xr from climanet.dataset import DataLoaderConfig, STDataset from climanet.predict import PredictionConfig, predict_monthly_var -from climanet.utils import configure_compute_resources, load_model, read_st_data, set_seed - +from climanet.utils import ( + configure_compute_resources, + load_model, + read_st_data, + set_seed, +) if __name__ == "__main__": parser = argparse.ArgumentParser()