diff --git a/scripts/prediction.py b/scripts/prediction.py new file mode 100644 index 0000000..b7a4235 --- /dev/null +++ b/scripts/prediction.py @@ -0,0 +1,149 @@ +import argparse +from pathlib import Path + +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, +) + +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(), + ) + parser.add_argument( + "--raw-data-folder", + 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() + raw_data_folder = Path(args.raw_data_folder).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 = 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=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, + 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, + ) + + # Predict residuals + predictions = predict_monthly_var( + model=model, + dataset=dataset_test, + dataloader_config=dataloader_config, + 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 new file mode 100644 index 0000000..3841a19 --- /dev/null +++ b/scripts/prediction.slurm @@ -0,0 +1,39 @@ +#!/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=1: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" +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"\ + --raw-data-folder "$RAW_DATA_FOLDER" + +echo "**********Prediction script completed.************"