diff --git a/climanet/train.py b/climanet/train.py index 682133f..c5ed08d 100644 --- a/climanet/train.py +++ b/climanet/train.py @@ -25,6 +25,9 @@ class TrainConfig: patience: int = 10 accumulation_steps: int = 1 optimizer_lr: float = 1e-3 + optimizer_weight_decay: float = 1e-2 + scheduler_lr_factor: float = 0.1 + scheduler_min_lr: float = 1e-5 device: str = "cpu" verbose: bool = False verbose_epoch_interval: int = 20 @@ -104,7 +107,9 @@ def train_monthly_model( # Set the optimizer optimizer = torch.optim.AdamW( - model.parameters(), lr=training_config.optimizer_lr, weight_decay=1e-2 + model.parameters(), + lr=training_config.optimizer_lr, + weight_decay=training_config.optimizer_weight_decay, ) best_loss = float("inf") @@ -120,9 +125,9 @@ def train_monthly_model( scheduler = ReduceLROnPlateau( optimizer, mode="min", - factor=0.5, + factor=training_config.scheduler_lr_factor, patience=training_config.patience // 2, # Reduce LR before early stop triggers - min_lr=1e-7, + min_lr=training_config.scheduler_min_lr, ) model.train() @@ -215,6 +220,8 @@ def train_monthly_model( if counter >= training_config.patience and current_lr <= scheduler.min_lrs[0]: if training_config.store_logs: writer.add_text("Training", f"Early stop at epoch {epoch}", epoch) + if training_config.verbose: + print(f"Early stopping triggered at epoch {epoch}. Best loss: {best_loss:.6f}") break # Restore best model diff --git a/scripts/README.md b/scripts/README.md new file mode 100644 index 0000000..c9569cb --- /dev/null +++ b/scripts/README.md @@ -0,0 +1,43 @@ +# Scripts + +## Structure + +- `data_preparation.*`: Scripts for preparing the data for training, tuning and evaluation. Mainly converting large netCDF files to Zarr storage with specific chunking strategies. This allows executing training for larger-than-memory datasets. +- `example_training.*`: example training script +- `tuning.*`: Scripts for hyperparameter tuning. +- `run_best_tuned_model.*`: Scripts for running the best tuned model on the test set. +- `training.*`: Scripts for training the model with the best hyperparameters found in the tuning experiments. + +## Experiments + +### Tuning experiments + +- datasplit: train set = 2020, validation set = 2021, test set = 2022 +- path of tuning results: `/work//eso4clima/tune/`. +- test loss: 0.036662004509047774 +- hyperparameters of the best model: + + ``` + {'patch_size': 8, + 'overlap': 1, + 'embed_dim': 64, + 'dropout': 0.2, + 'hidden': 32, + 'spatial_depth': 3, + 'spatial_heads': 2, + 'optimizer_lr': 0.001787422899066508, + 'batch_config': {'accumulation_steps': 2}} + ``` + +### Training experiments + +Use the best hyperparameters found in the tuning experiments to train the model +on the training set; three years 2018-2020 for training, 2021 for validation in +the training loop. Because three years of hourly data is too large to fit into +memory, we use `load_lazy` option in the dataset. This makes the training +process slower, but allows us to train on larger-than-memory datasets. The +training is done for 100 epochs, and the best model is saved based on the +validation loss. Note that the training is done only on one GPU node, including +4 GPUs. + +The results are stored at `/work//eso4clima/train/sst_01/`. diff --git a/scripts/run_best_tuned_model.py b/scripts/run_best_tuned_model.py new file mode 100644 index 0000000..7c2964d --- /dev/null +++ b/scripts/run_best_tuned_model.py @@ -0,0 +1,184 @@ +import argparse +from pathlib import Path + +import xarray as xr +from ray import tune + +from climanet.dataset import DataLoaderConfig, STDataset +from climanet.predict import PredictionConfig, predict_monthly_var +from climanet.utils import data_preparation, read_st_data + + +def build_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser( + description=( + "Load the best Ray Tune checkpoint, prepare the test data, and evaluate the " + "trained model on the 2023 test period." + ) + ) + parser.add_argument( + "--experiment-path", + type=Path, + required=True, + help="Path to the Ray Tune experiment directory containing the checkpoint.", + ) + parser.add_argument( + "--test-data-dir", + type=Path, + required=True, + help="Directory containing the test NetCDF files.", + ) + parser.add_argument( + "--lsm-file-path", + type=Path, + required=True, + help="Path to the land-sea mask NetCDF file.", + ) + parser.add_argument( + "--run-dir", + type=Path, + default=Path("./run_dir_tune_test").resolve(), + help="Directory used for the evaluation run and saved logs.", + ) + parser.add_argument( + "--var-name", + type=str, + default="tos", + help="Variable name to evaluate in the NetCDF files.", + ) + parser.add_argument( + "--year", + type=str, + default="2022", + help="Year pattern to include in the test files (e.g. 2022).", + ) + return parser + + +def main() -> None: + args = build_parser().parse_args() + experiment_path = args.experiment_path.resolve() + test_data_dir = args.test_data_dir.resolve() + lsm_file_path = args.lsm_file_path.resolve() + run_dir = args.run_dir.resolve() + run_dir.mkdir(parents=True, exist_ok=True) + + if not experiment_path.exists(): + raise FileNotFoundError( + f"Experiment directory does not exist: {experiment_path}" + ) + if not test_data_dir.exists(): + raise FileNotFoundError(f"Test data directory does not exist: {test_data_dir}") + if not lsm_file_path.exists(): + raise FileNotFoundError(f"LSM file does not exist: {lsm_file_path}") + + daily_files = list( + test_data_dir.glob(f"{args.year}*_hr_ERA5dc_masked_{args.var_name}*.nc") + ) + monthly_files = list( + test_data_dir.glob(f"{args.year}*_mon_ERA5dc_masked_{args.var_name}*.nc") + ) + + if not daily_files: + raise FileNotFoundError( + f"No daily test files found for year '{args.year}' in '{test_data_dir}'" + ) + if not monthly_files: + raise FileNotFoundError( + f"No monthly test files found for year '{args.year}' in '{test_data_dir}'" + ) + + print(f"Using daily files ({len(daily_files)}): {daily_files[:3]} ...") + print(f"Using monthly files ({len(monthly_files)}): {monthly_files[:3]} ...") + + daily_data_test = xr.open_mfdataset( + daily_files, combine="by_coords", parallel=False + ) + monthly_data_test = xr.open_mfdataset( + monthly_files, combine="by_coords", parallel=False + ) + + test_data_zarr_dir = run_dir / "test_data_zarr" + test_data_zarr_dir.mkdir(parents=True, exist_ok=True) + + _ = data_preparation( + daily_data_test[args.var_name], + monthly_data_test[args.var_name], + calculate_residuals=True, + is_hourly=True, + save_to_zarr=True, + run_dir=test_data_zarr_dir, + ) + + input_da, input_da_nan_mask, monthly_da, padded_days_mask, time_features = ( + read_st_data( + data_path=test_data_zarr_dir, + var_name=args.var_name, + ) + ) + + lsm_mask = xr.open_dataset(lsm_file_path) + + num_patches = (10, 10) + patch_size = (1, 4, 4) + spatial_patch_size = ( + patch_size[1] * num_patches[0], + patch_size[2] * num_patches[1], + ) + stride = (spatial_patch_size[0] // 5, spatial_patch_size[1] // 5) + + 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["lsm"], + patch_size=(1, *spatial_patch_size), + stride=stride, + sh_embed_dim=96, + sh_order_L=10, + verbose=True, + load_lazy=False, + ) + print(f"Created test dataset with {len(dataset_test)} patches.") + + analysis = tune.ExperimentAnalysis(str(experiment_path)) + best_result = analysis.get_best_trial("loss", "min") + best_checkpoint = best_result.checkpoint + model_path = Path(best_checkpoint.path) / "checkpoint.pt" + print(f"Best checkpoint path: {model_path}") + + prediction_config = PredictionConfig( + calculate_residuals=True, + return_numpy=True, + save_predictions=False, + return_loss=True, + device="cpu", + verbose=False, + ) + + dataloader_config = DataLoaderConfig( + batch_size=10, + shuffle=True, + num_workers=0, + pin_memory=False, + persistent_workers=False, + device="cpu", + multiprocessing_context=None, + ) + + test_loss = predict_monthly_var( + model=model_path, + dataset=dataset_test, + dataloader_config=dataloader_config, + prediction_config=prediction_config, + run_dir=run_dir, + ) + + print("Test loss:") + print(test_loss) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_best_tuned_model.slurm b/scripts/run_best_tuned_model.slurm new file mode 100644 index 0000000..aa3f3b4 --- /dev/null +++ b/scripts/run_best_tuned_model.slurm @@ -0,0 +1,24 @@ +#!/bin/bash +#SBATCH --job-name=climanet_eval +#SBATCH --nodes=1 +#SBATCH --ntasks-per-node=1 +#SBATCH --cpus-per-task=128 +#SBATCH --time=02:00:00 +#SBATCH --account=bd0854 +#SBATCH --partition=compute +#SBATCH --output=climanet_eval_%j.out +#SBATCH --error=climanet_eval_%j.err + +set -euo pipefail + +source /home/b/b383704/eso4clima/ClimaNet/.venv/bin/activate + +python -u /home/b/b383704/eso4clima/run_best_tuned_model/run_best_tuned_model.py \ + --experiment-path /work/bd0854/eso4clima/tune/sst_01 \ + --test-data-dir /work/bd0854/b380103/eso4clima/output/sst/concatenated/ \ + --lsm-file-path /home/b/b383704/eso4clima/data/era5_lsm_bool.nc \ + --run-dir /home/b/b383704/eso4clima/run_best_tuned_model/run_dir \ + --var-name tos \ + --year 2022 + +printf "\nFinished evaluation run.\n" diff --git a/scripts/training.py b/scripts/training.py new file mode 100644 index 0000000..dc4c98a --- /dev/null +++ b/scripts/training.py @@ -0,0 +1,175 @@ +import argparse +from pathlib import Path + +import ray +import xarray as xr + +from climanet.dataset import DataLoaderConfig, STDataset +from climanet.st_encoder_decoder import SpatioTemporalModel +from climanet.train import TrainConfig, train_monthly_model +from climanet.utils import configure_compute_resources, read_st_data, set_seed + + +def _build_dataset( + prepared_data_dir: Path, + years: list[int], + var_name: str, + land_mask: xr.DataArray, + patch_size: tuple[int, int, int], + stride: tuple[int, int], +) -> STDataset: + + data = [read_st_data(data_path=f"{prepared_data_dir}/{year}", var_name=var_name) for year in years] + input_das, input_da_nan_masks, monthly_das, padded_days_masks, time_features_list = zip(*data) + + input_da = xr.concat(input_das, dim="M") + input_da_nan_mask = xr.concat(input_da_nan_masks, dim="M") + monthly_da = xr.concat(monthly_das, dim="M") + padded_days_mask = xr.concat(padded_days_masks, dim="M") + time_features = xr.concat(time_features_list, dim="M") + + return 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=land_mask, + patch_size=patch_size, + stride=stride, + sh_embed_dim=96, + sh_order_L=10, + verbose=False, + load_lazy=True, + ) + + +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( + "--tune-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() + tune_dir = Path(args.tune_dir).resolve() + run_dir = Path(args.run_dir).resolve() + + # Load the best hyperparameters from tuning + analysis = ray.tune.ExperimentAnalysis(str(tune_dir)) + best_result = analysis.get_best_trial("loss", "min") + best_config = best_result.config + + # set the random seed for reproducibility + set_seed() + + # 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 + + dataset_patch_size = (1, 40, 40) + dataset_stride = (20, 20) + + train_years = [2018, 2019, 2020] + dataset_train = _build_dataset( + prepared_data_dir=prepared_data_dir, + years=train_years, + var_name=var_name, + land_mask=lsm_mask, + patch_size=dataset_patch_size, + stride=dataset_stride, + ) + + validation_year = [2021] + dataset_validation = _build_dataset( + prepared_data_dir=prepared_data_dir, + years=validation_year, + var_name=var_name, + land_mask=lsm_mask, + patch_size=dataset_patch_size, + stride=dataset_stride, + ) + + # Build the dataloader config + dataloader_num_workers = 32 # adjust if needed + use_cuda = device == "cuda" + dataloader_config = DataLoaderConfig( + batch_size=100, # adjust if OOM issue + shuffle=True, + num_workers=dataloader_num_workers, + pin_memory=use_cuda, + persistent_workers=True, + device=device, + multiprocessing_context="spawn", + ) + + # Build the model with the best hyperparameters from tuning + patch_size = (1, best_config["patch_size"], best_config["patch_size"]) + overlap = best_config["overlap"] + embed_dim = best_config["embed_dim"] + dropout = best_config["dropout"] + hidden = best_config["hidden"] + spatial_depth = best_config["spatial_depth"] + spatial_heads = best_config["spatial_heads"] + + model = SpatioTemporalModel( + patch_size=patch_size, + overlap=overlap, + embed_dim=embed_dim, + dropout=dropout, + hidden=hidden, + spatial_depth=spatial_depth, + spatial_heads=spatial_heads, + ) + + # 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 + training_config = TrainConfig( + calculate_residuals=True, + num_epoch=101, + patience=10, + accumulation_steps=2, + optimizer_lr=best_config["optimizer_lr"], + device=device, + verbose=True, + verbose_epoch_interval=20, + tune_checkpoint=False, + store_model=True, + store_logs=True, + ) + + trained_model = train_monthly_model( + model=model, + dataset_train=dataset_train, + dataloader_config=dataloader_config, + training_config=training_config, + dataset_validation=dataset_validation, + run_dir=run_dir, + ) diff --git a/scripts/training.slurm b/scripts/training.slurm new file mode 100644 index 0000000..0e8d7ce --- /dev/null +++ b/scripts/training.slurm @@ -0,0 +1,37 @@ +#!/bin/bash +#SBATCH --job-name=training +#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=training_%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/train" + +# data directory (adjust this path to your data location) +PREPARED_DATA_DIR="/scratch/b/$USER/data" +TUNE_DIR="/scratch/b/$USER/tune" +LSM_DIR="/scratch/b/$USER/data" + +echo "Starting training.py script..." +python -u $HOME/ClimaNet/scripts/training.py \ + --run-dir "$RUN_DIR" \ + --prepared-data-dir "$PREPARED_DATA_DIR" \ + --tune-dir "$TUNE_DIR" \ + --lsm-dir "$LSM_DIR" + +echo "**********Training script completed.************"