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
149 changes: 149 additions & 0 deletions scripts/prediction.py
Original file line number Diff line number Diff line change
@@ -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)
39 changes: 39 additions & 0 deletions scripts/prediction.slurm
Original file line number Diff line number Diff line change
@@ -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.************"