code update
This commit is contained in:
@@ -43,7 +43,7 @@ run_configuration = {
|
|||||||
"take_every_nth": take_every_nth,
|
"take_every_nth": take_every_nth,
|
||||||
"shift_in_hours": shift_in_hours,
|
"shift_in_hours": shift_in_hours,
|
||||||
"output_window_offset": output_window_offset,
|
"output_window_offset": output_window_offset,
|
||||||
"batch_size": 256,
|
"batch_size": 128,
|
||||||
"max_lr": 1e-5,
|
"max_lr": 1e-5,
|
||||||
"num_epochs": 10,
|
"num_epochs": 10,
|
||||||
"patience": 3,
|
"patience": 3,
|
||||||
|
|||||||
@@ -55,6 +55,11 @@ def simple_x_y_collate(batch):
|
|||||||
collated_x.append(np.concatenate(current_x, axis=1))
|
collated_x.append(np.concatenate(current_x, axis=1))
|
||||||
collated_y.append(np.concatenate(current_y, axis=1))
|
collated_y.append(np.concatenate(current_y, axis=1))
|
||||||
|
|
||||||
|
# convert to numpy arrays
|
||||||
|
collated_x = np.array(collated_x)
|
||||||
|
collated_y = np.array(collated_y)
|
||||||
|
|
||||||
|
# convert to torch tensors
|
||||||
return torch.tensor(collated_x, dtype=torch.float32), torch.tensor(collated_y, dtype=torch.float32)
|
return torch.tensor(collated_x, dtype=torch.float32), torch.tensor(collated_y, dtype=torch.float32)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Executable
+58
@@ -0,0 +1,58 @@
|
|||||||
|
#!/bin/bash
|
||||||
|
|
||||||
|
# Parameter check
|
||||||
|
if [[ $# -ne 5 ]]; then
|
||||||
|
echo "Usage: $0 <run_config_module> <partition> <gpu_type> <num_gpus> <num_cpus_per_gpu>"
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
|
|
||||||
|
# Assign parameters
|
||||||
|
run_config=$1
|
||||||
|
partition=$2
|
||||||
|
gpu_type=$3
|
||||||
|
num_gpus=$4
|
||||||
|
num_cpus_per_gpu=$5
|
||||||
|
|
||||||
|
# Derive run name
|
||||||
|
run_name=$(basename "$run_config")
|
||||||
|
run_name="${run_name%.*}"
|
||||||
|
|
||||||
|
# Confirm inputs
|
||||||
|
echo "Starting run with configuration: $run_config"
|
||||||
|
echo "Partition: $partition"
|
||||||
|
echo "GPU type: $gpu_type"
|
||||||
|
echo "Number of GPUs: $num_gpus"
|
||||||
|
echo "Number of CPUs per GPU: $num_cpus_per_gpu"
|
||||||
|
|
||||||
|
# Set paths
|
||||||
|
current_dir_path=$(dirname "$(realpath "$0")")
|
||||||
|
venv_path=$(realpath "$current_dir_path/../../venv/bin/activate")
|
||||||
|
|
||||||
|
# Create log directory
|
||||||
|
log_dir="/work/rr41qemu-MA/logs/${run_name}_$(date +%Y-%m-%d_%H-%M-%S)"
|
||||||
|
mkdir -p "$log_dir"
|
||||||
|
echo "Logging to $log_dir"
|
||||||
|
|
||||||
|
# Submit job
|
||||||
|
sbatch <<EOF
|
||||||
|
#!/bin/bash
|
||||||
|
#SBATCH --job-name=$run_name
|
||||||
|
#SBATCH --output=$log_dir/log.out
|
||||||
|
#SBATCH --error=$log_dir/log.err
|
||||||
|
#SBATCH --time=48:00:00
|
||||||
|
#SBATCH --ntasks=1
|
||||||
|
#SBATCH --cpus-per-task=$(($num_gpus * $num_cpus_per_gpu))
|
||||||
|
#SBATCH --mem=32G
|
||||||
|
#SBATCH --partition=$partition
|
||||||
|
#SBATCH --gpus=$gpu_type:$num_gpus
|
||||||
|
|
||||||
|
echo "Loading python virtual environment..."
|
||||||
|
source $venv_path
|
||||||
|
|
||||||
|
echo "Loading python 3.10..."
|
||||||
|
module load Python/3.10.4-GCCcore-11.3.0
|
||||||
|
|
||||||
|
cd $current_dir_path
|
||||||
|
|
||||||
|
torchrun --nproc_per_node=$num_gpus training_wrapper.py $run_config --item_limit -1
|
||||||
|
EOF
|
||||||
+57
-25
@@ -1,8 +1,10 @@
|
|||||||
import json
|
import json
|
||||||
|
import socket
|
||||||
import sys
|
import sys
|
||||||
import os
|
import os
|
||||||
import argparse
|
import argparse
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
import lmdb
|
import lmdb
|
||||||
@@ -19,15 +21,12 @@ from utils.data_utils import LMDBIterableDataset
|
|||||||
from utils.utils import get_variable_from_module, get_logger, convert_for_json
|
from utils.utils import get_variable_from_module, get_logger, convert_for_json
|
||||||
from utils.training import train_model
|
from utils.training import train_model
|
||||||
|
|
||||||
|
print(f"PID: {os.getpid()} on host: {socket.gethostname()}")
|
||||||
|
|
||||||
dotenv.load_dotenv()
|
dotenv.load_dotenv()
|
||||||
|
|
||||||
logger = None
|
logger = None
|
||||||
|
|
||||||
# set up distributed training
|
|
||||||
dist.init_process_group(backend="nccl", init_method="env://")
|
|
||||||
local_rank = torch.distributed.get_rank()
|
|
||||||
torch.cuda.set_device(local_rank)
|
|
||||||
|
|
||||||
|
|
||||||
def prepare_run(results_dir: str,
|
def prepare_run(results_dir: str,
|
||||||
lmdb_root_dir: str,
|
lmdb_root_dir: str,
|
||||||
@@ -59,7 +58,8 @@ def train(
|
|||||||
train_ids: list,
|
train_ids: list,
|
||||||
val_ids: list,
|
val_ids: list,
|
||||||
dataset_dir: str,
|
dataset_dir: str,
|
||||||
log_dir: str) -> None:
|
log_dir: str,
|
||||||
|
num_dataloader_workers: int = 1) -> None:
|
||||||
# create training and validation loaders
|
# create training and validation loaders
|
||||||
logger.info("Creating training loaders")
|
logger.info("Creating training loaders")
|
||||||
train_loader = LMDBIterableDataset(dataset_dir,
|
train_loader = LMDBIterableDataset(dataset_dir,
|
||||||
@@ -88,14 +88,14 @@ def train(
|
|||||||
train_dataset=train_loader,
|
train_dataset=train_loader,
|
||||||
val_dataset=val_loader,
|
val_dataset=val_loader,
|
||||||
log_dir=log_dir,
|
log_dir=log_dir,
|
||||||
logger=logger
|
logger=logger,
|
||||||
|
num_dataloader_workers=num_dataloader_workers,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def evaluate(model_configuration: dict,
|
def evaluate(model_configuration: dict,
|
||||||
training_configuration: dict,
|
training_configuration: dict,
|
||||||
test_ids: list,
|
test_ids: list) -> dict:
|
||||||
device) -> dict:
|
|
||||||
logger.info(f"Evaluating model {model_configuration['id']}")
|
logger.info(f"Evaluating model {model_configuration['id']}")
|
||||||
eval_functions = get_eval_functions(model_configuration)
|
eval_functions = get_eval_functions(model_configuration)
|
||||||
results = evaluate_model(model_configuration,
|
results = evaluate_model(model_configuration,
|
||||||
@@ -114,6 +114,18 @@ def evaluate(model_configuration: dict,
|
|||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
# set up distributed training
|
||||||
|
dist.init_process_group(backend="nccl", init_method="env://")
|
||||||
|
|
||||||
|
# sleeping for a bit to allow all processes to initialize
|
||||||
|
time.sleep(5)
|
||||||
|
|
||||||
|
local_rank = torch.distributed.get_rank()
|
||||||
|
world_size = torch.distributed.get_world_size()
|
||||||
|
torch.cuda.set_device(local_rank)
|
||||||
|
|
||||||
|
print(f"Rank {local_rank} initialized with world size {world_size}")
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(description="Training wrapper for model training")
|
parser = argparse.ArgumentParser(description="Training wrapper for model training")
|
||||||
parser.add_argument("run_configuration_module",
|
parser.add_argument("run_configuration_module",
|
||||||
type=str,
|
type=str,
|
||||||
@@ -144,11 +156,11 @@ if __name__ == "__main__":
|
|||||||
required=False,
|
required=False,
|
||||||
default=None,
|
default=None,
|
||||||
help="Limit the number of items to process, default is None (no limit)")
|
help="Limit the number of items to process, default is None (no limit)")
|
||||||
parser.add_argument("--device",
|
parser.add_argument("--num_dataloader_workers",
|
||||||
type=str,
|
type=int,
|
||||||
required=False,
|
required=False,
|
||||||
default="cuda",
|
default=1,
|
||||||
help="Device to use for training, default is cuda")
|
help="Number of workers for the dataloader, default is 1")
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
# load run configuration from module
|
# load run configuration from module
|
||||||
@@ -201,8 +213,11 @@ if __name__ == "__main__":
|
|||||||
if not os.path.exists(log_dir):
|
if not os.path.exists(log_dir):
|
||||||
os.makedirs(log_dir)
|
os.makedirs(log_dir)
|
||||||
|
|
||||||
item_limit = args.item_limit if args.item_limit else run_configuration["item_limit"]
|
try:
|
||||||
if item_limit == -1:
|
item_limit = args.item_limit if args.item_limit else run_configuration["item_limit"]
|
||||||
|
if item_limit == -1:
|
||||||
|
item_limit = None
|
||||||
|
except:
|
||||||
item_limit = None
|
item_limit = None
|
||||||
|
|
||||||
run_name = run_configuration["name"]
|
run_name = run_configuration["name"]
|
||||||
@@ -214,29 +229,42 @@ if __name__ == "__main__":
|
|||||||
else:
|
else:
|
||||||
run_id = None
|
run_id = None
|
||||||
|
|
||||||
|
if torch.distributed.is_initialized():
|
||||||
|
torch.distributed.barrier()
|
||||||
|
|
||||||
# broadcast run_id to all processes
|
# broadcast run_id to all processes
|
||||||
if dist.is_initialized():
|
if dist.is_initialized():
|
||||||
|
print(f"Rank {local_rank}: Broadcasting run_id {run_id}")
|
||||||
run_id_list = [run_id]
|
run_id_list = [run_id]
|
||||||
torch.distributed.broadcast_object_list(run_id_list, src=0)
|
torch.distributed.broadcast_object_list(run_id_list, src=0)
|
||||||
run_id = run_id_list[0]
|
run_id = run_id_list[0]
|
||||||
|
|
||||||
# make sure run_id is a string
|
# make sure run_id is a string
|
||||||
run_id = str(run_id)
|
run_id = str(run_id)
|
||||||
|
print(f"Rank {local_rank}: fetched run_id {run_id}")
|
||||||
|
|
||||||
# append run_id to results_dir and log_dir
|
# append run_id to results_dir and log_dir
|
||||||
results_dir = os.path.join(results_dir, run_id)
|
results_dir = os.path.join(results_dir, run_id)
|
||||||
log_dir = os.path.join(log_dir, run_id)
|
log_dir = os.path.join(log_dir, run_id)
|
||||||
|
|
||||||
|
print(f"Rank {local_rank}: Results directory: {results_dir}")
|
||||||
|
print(f"Rank {local_rank}: Log directory: {log_dir}")
|
||||||
|
|
||||||
# create directories if they do not exist, only on the main process
|
# create directories if they do not exist, only on the main process
|
||||||
if torch.distributed.get_rank() == 0 or not dist.is_initialized():
|
if torch.distributed.get_rank() == 0 or not dist.is_initialized():
|
||||||
|
print(f"Rank {local_rank}: Creating directories for run {run_name}")
|
||||||
if not os.path.exists(results_dir):
|
if not os.path.exists(results_dir):
|
||||||
os.makedirs(results_dir)
|
os.makedirs(results_dir)
|
||||||
if not os.path.exists(log_dir):
|
if not os.path.exists(log_dir):
|
||||||
os.makedirs(log_dir)
|
os.makedirs(log_dir)
|
||||||
|
|
||||||
# set up logger
|
# sync all processes to make sure the directories are created
|
||||||
logger = get_logger(module_name=run_name, filename=os.path.join(log_dir, "main.log"))
|
if dist.is_initialized():
|
||||||
|
print(f"Rank {local_rank}: Waiting for all processes to create directories")
|
||||||
|
dist.barrier()
|
||||||
|
|
||||||
|
# set up logger for all processes
|
||||||
|
logger = get_logger(module_name=run_name, filename=os.path.join(log_dir, f"main_{local_rank}.log"))
|
||||||
|
|
||||||
logger.info(f"Rank {local_rank}: Starting run {run_name}")
|
logger.info(f"Rank {local_rank}: Starting run {run_name}")
|
||||||
logger.info(f"Rank {local_rank}: Run ID: {run_id}")
|
logger.info(f"Rank {local_rank}: Run ID: {run_id}")
|
||||||
@@ -270,14 +298,18 @@ if __name__ == "__main__":
|
|||||||
base_training_configuration=run_training_configuration
|
base_training_configuration=run_training_configuration
|
||||||
)
|
)
|
||||||
logger.info(f"Rank {local_rank}: Finished preparing run {run_step_name}")
|
logger.info(f"Rank {local_rank}: Finished preparing run {run_step_name}")
|
||||||
# sync after preparing run
|
|
||||||
if dist.is_initialized():
|
# all ranks sync here
|
||||||
dist.barrier()
|
if dist.is_initialized():
|
||||||
else:
|
if local_rank != 0:
|
||||||
# wait for rank 0 to finish preparing run
|
|
||||||
if dist.is_initialized():
|
|
||||||
dist.barrier()
|
|
||||||
logger.info(f"Rank {local_rank}: Waiting for rank 0 to finish preparing run {run_step_name}")
|
logger.info(f"Rank {local_rank}: Waiting for rank 0 to finish preparing run {run_step_name}")
|
||||||
|
dist.barrier()
|
||||||
|
if local_rank != 0:
|
||||||
|
logger.info(
|
||||||
|
f"Rank {local_rank}: Finished waiting for rank 0 to finish preparing run {run_step_name}")
|
||||||
|
|
||||||
|
if local_rank != 0:
|
||||||
|
# after the barrier, rank 0 will have prepared the run
|
||||||
run_model_configuration, run_training_configuration, train_ids, val_ids, test_ids, dataset_dir = prepare_run(
|
run_model_configuration, run_training_configuration, train_ids, val_ids, test_ids, dataset_dir = prepare_run(
|
||||||
results_dir=results_dir,
|
results_dir=results_dir,
|
||||||
lmdb_root_dir=lmdb_root_dir,
|
lmdb_root_dir=lmdb_root_dir,
|
||||||
|
|||||||
@@ -479,6 +479,10 @@ class LMDBIterableDataset(IterableDataset):
|
|||||||
self.key_subset = key_subset
|
self.key_subset = key_subset
|
||||||
# reset length
|
# reset length
|
||||||
self.len = None
|
self.len = None
|
||||||
|
# reset lmdb env
|
||||||
|
if self.lmdb_env is not None:
|
||||||
|
self.lmdb_env.close()
|
||||||
|
self.lmdb_env = None
|
||||||
|
|
||||||
def get_length_of_data_subset(self, key_set: list[str]):
|
def get_length_of_data_subset(self, key_set: list[str]):
|
||||||
"""
|
"""
|
||||||
@@ -522,12 +526,13 @@ class LMDBIterableDataset(IterableDataset):
|
|||||||
return num_steps
|
return num_steps
|
||||||
|
|
||||||
def __iter__(self):
|
def __iter__(self):
|
||||||
self.init_lmdb_env()
|
|
||||||
|
|
||||||
# if no key subset is set, use all keys
|
# if no key subset is set, use all keys
|
||||||
if self.key_subset is None:
|
if self.key_subset is None:
|
||||||
self.key_subset = self.lmdb_keys
|
self.key_subset = self.lmdb_keys
|
||||||
|
|
||||||
|
self.init_lmdb_env()
|
||||||
|
|
||||||
random.shuffle(self.key_subset)
|
random.shuffle(self.key_subset)
|
||||||
|
|
||||||
batch = list()
|
batch = list()
|
||||||
@@ -563,12 +568,12 @@ class LMDBIterableDataset(IterableDataset):
|
|||||||
meminit=False)
|
meminit=False)
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
self.init_lmdb_env()
|
|
||||||
|
|
||||||
# if no key subset is set, use all keys
|
# if no key subset is set, use all keys
|
||||||
if self.key_subset is None:
|
if self.key_subset is None:
|
||||||
self.key_subset = self.lmdb_keys
|
self.key_subset = self.lmdb_keys
|
||||||
|
|
||||||
|
self.init_lmdb_env()
|
||||||
|
|
||||||
if self.len is None:
|
if self.len is None:
|
||||||
num_steps = self.get_length_of_data_subset(self.key_subset)
|
num_steps = self.get_length_of_data_subset(self.key_subset)
|
||||||
self.len = num_steps
|
self.len = num_steps
|
||||||
|
|||||||
@@ -130,6 +130,8 @@ def evaluate_model(model_configuration: dict,
|
|||||||
if eval_fn is not None:
|
if eval_fn is not None:
|
||||||
eval_fn_name = eval_fn["name"]
|
eval_fn_name = eval_fn["name"]
|
||||||
accumulation_fn = eval_fn["accumulation_fn"]
|
accumulation_fn = eval_fn["accumulation_fn"]
|
||||||
|
if eval_fn_name not in errors:
|
||||||
|
continue
|
||||||
for key in errors[eval_fn_name]:
|
for key in errors[eval_fn_name]:
|
||||||
if len(errors[eval_fn_name][key]) == 0:
|
if len(errors[eval_fn_name][key]) == 0:
|
||||||
errors[eval_fn_name][key] = np.nan
|
errors[eval_fn_name][key] = np.nan
|
||||||
|
|||||||
+107
-22
@@ -104,7 +104,8 @@ def train_model(model: nn.Module,
|
|||||||
train_dataset: LMDBIterableDataset,
|
train_dataset: LMDBIterableDataset,
|
||||||
val_dataset: LMDBIterableDataset,
|
val_dataset: LMDBIterableDataset,
|
||||||
log_dir: str = "./logs",
|
log_dir: str = "./logs",
|
||||||
logger=None) -> torch.nn.Module:
|
num_dataloader_workers: int = 4,
|
||||||
|
logger=None) -> None:
|
||||||
if logger is None:
|
if logger is None:
|
||||||
logger = get_logger(__name__, f"{log_dir}/{model_configuration['id']}_{training_configuration['id']}.log")
|
logger = get_logger(__name__, f"{log_dir}/{model_configuration['id']}_{training_configuration['id']}.log")
|
||||||
|
|
||||||
@@ -117,7 +118,6 @@ def train_model(model: nn.Module,
|
|||||||
# get computation rank
|
# get computation rank
|
||||||
if torch.distributed.is_initialized():
|
if torch.distributed.is_initialized():
|
||||||
local_rank = torch.distributed.get_rank()
|
local_rank = torch.distributed.get_rank()
|
||||||
torch.cuda.set_device(local_rank)
|
|
||||||
device = torch.device(f"cuda:{local_rank}")
|
device = torch.device(f"cuda:{local_rank}")
|
||||||
world_size = torch.distributed.get_world_size()
|
world_size = torch.distributed.get_world_size()
|
||||||
else:
|
else:
|
||||||
@@ -132,10 +132,46 @@ def train_model(model: nn.Module,
|
|||||||
all_train_ids = train_dataset.lmdb_keys
|
all_train_ids = train_dataset.lmdb_keys
|
||||||
train_subsets = [get_ranked_ids(all_train_ids, i, local_rank, world_size) for i in range(num_epochs)]
|
train_subsets = [get_ranked_ids(all_train_ids, i, local_rank, world_size) for i in range(num_epochs)]
|
||||||
# calc total number of steps for gpu, as it is dependent on subsets
|
# calc total number of steps for gpu, as it is dependent on subsets
|
||||||
total_train_steps = sum([train_dataset.get_length_of_data_subset(subset) for subset in train_subsets])
|
local_steps = [train_dataset.get_length_of_data_subset(subset) for subset in train_subsets]
|
||||||
|
total_train_steps = sum(local_steps)
|
||||||
|
|
||||||
all_val_ids = val_dataset.lmdb_keys
|
all_val_ids = val_dataset.lmdb_keys
|
||||||
val_subsets = [get_ranked_ids(all_val_ids, i, local_rank, world_size) for i in range(num_epochs)]
|
val_subsets = [get_ranked_ids(all_val_ids, i, local_rank, world_size) for i in range(num_epochs)]
|
||||||
|
val_subset_lengths = [len(subset) for subset in val_subsets]
|
||||||
|
|
||||||
|
def get_synced_values(local_values):
|
||||||
|
"""
|
||||||
|
Sync values across all processes.
|
||||||
|
Args:
|
||||||
|
local_values: list of local values
|
||||||
|
Returns:
|
||||||
|
list of synced values
|
||||||
|
"""
|
||||||
|
if not isinstance(local_values, torch.Tensor):
|
||||||
|
local_values = torch.tensor(local_values, device=device, dtype=torch.float32)
|
||||||
|
else:
|
||||||
|
local_values = local_values.to(device)
|
||||||
|
gathered_values = [torch.zeros_like(local_values) for _ in range(world_size)]
|
||||||
|
torch.distributed.all_gather(gathered_values, local_values)
|
||||||
|
return gathered_values
|
||||||
|
|
||||||
|
# sync lengths of subsets and adjust to minimum length for equal sized training lengths
|
||||||
|
global_lengths = get_synced_values(local_steps)
|
||||||
|
global_lengths = [x.cpu().numpy() for x in global_lengths]
|
||||||
|
train_epoch_lengths = list()
|
||||||
|
for i in range(len(global_lengths[0])):
|
||||||
|
current_epoch_lengths = [x[i] for x in global_lengths]
|
||||||
|
train_epoch_lengths.append(int(min(current_epoch_lengths)))
|
||||||
|
print(f"Rank {local_rank}: Global lengths: {global_lengths} cut to {train_epoch_lengths}")
|
||||||
|
|
||||||
|
# also sync the val subsets
|
||||||
|
global_val_lengths = get_synced_values(val_subset_lengths)
|
||||||
|
global_val_lengths = [x.cpu().numpy() for x in global_val_lengths]
|
||||||
|
val_epoch_lengths = list()
|
||||||
|
for i in range(len(global_val_lengths[0])):
|
||||||
|
current_epoch_lengths = [x[i] for x in global_val_lengths]
|
||||||
|
val_epoch_lengths.append(int(min(current_epoch_lengths)))
|
||||||
|
print(f"Rank {local_rank}: Global val lengths: {global_val_lengths} cut to {val_epoch_lengths}")
|
||||||
else:
|
else:
|
||||||
train_subsets = [train_dataset.lmdb_keys] * num_epochs
|
train_subsets = [train_dataset.lmdb_keys] * num_epochs
|
||||||
val_subsets = [val_dataset.lmdb_keys] * num_epochs
|
val_subsets = [val_dataset.lmdb_keys] * num_epochs
|
||||||
@@ -156,12 +192,12 @@ def train_model(model: nn.Module,
|
|||||||
train_dataloader = DataLoader(
|
train_dataloader = DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
batch_size=None,
|
batch_size=None,
|
||||||
num_workers=4,
|
num_workers=num_dataloader_workers,
|
||||||
)
|
)
|
||||||
val_dataloader = DataLoader(
|
val_dataloader = DataLoader(
|
||||||
val_dataset,
|
val_dataset,
|
||||||
batch_size=None,
|
batch_size=None,
|
||||||
num_workers=4,
|
num_workers=num_dataloader_workers,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(f"Rank {local_rank}: Training {training_id} with {num_epochs} epochs")
|
logger.info(f"Rank {local_rank}: Training {training_id} with {num_epochs} epochs")
|
||||||
@@ -171,6 +207,12 @@ def train_model(model: nn.Module,
|
|||||||
|
|
||||||
model.to(device)
|
model.to(device)
|
||||||
|
|
||||||
|
# wrap model in DDP if distributed training
|
||||||
|
if torch.distributed.is_initialized():
|
||||||
|
model = torch.nn.parallel.DistributedDataParallel(model,
|
||||||
|
device_ids=[local_rank],
|
||||||
|
output_device=local_rank)
|
||||||
|
|
||||||
# load training state from training configuration, if available
|
# load training state from training configuration, if available
|
||||||
optimizer = AdamW(model.parameters(), lr=learning_parameters["learning_rate"])
|
optimizer = AdamW(model.parameters(), lr=learning_parameters["learning_rate"])
|
||||||
scheduler = OneCycleLR(optimizer,
|
scheduler = OneCycleLR(optimizer,
|
||||||
@@ -204,21 +246,37 @@ def train_model(model: nn.Module,
|
|||||||
train_dataloader = DataLoader(
|
train_dataloader = DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
batch_size=None,
|
batch_size=None,
|
||||||
num_workers=4,
|
num_workers=num_dataloader_workers,
|
||||||
|
# set multiprocessing start method to spawn
|
||||||
|
# multiprocessing_context="forkserver",
|
||||||
)
|
)
|
||||||
|
train_length = train_epoch_lengths[epoch - 1]
|
||||||
val_dataloader = DataLoader(
|
val_dataloader = DataLoader(
|
||||||
val_dataset,
|
val_dataset,
|
||||||
batch_size=None,
|
batch_size=None,
|
||||||
num_workers=4,
|
num_workers=num_dataloader_workers,
|
||||||
|
# set multiprocessing start method to spawn
|
||||||
|
# multiprocessing_context="forkserver",
|
||||||
)
|
)
|
||||||
|
val_length = val_epoch_lengths[epoch - 1]
|
||||||
|
|
||||||
iterator = iter(train_dataloader)
|
iterator = iter(train_dataloader)
|
||||||
for step in tqdm(range(len(train_dataloader)), total=len(train_dataloader)):
|
for step in tqdm(range(train_length)):
|
||||||
loss = batch_loss_fn(model,
|
try:
|
||||||
iterator,
|
loss = batch_loss_fn(model,
|
||||||
loss_functions,
|
iterator,
|
||||||
device,
|
loss_functions,
|
||||||
model_configuration)
|
device,
|
||||||
|
model_configuration)
|
||||||
|
except StopIteration:
|
||||||
|
# if the iterator is exhausted, reset it
|
||||||
|
logger.info(f"Rank {local_rank}: Iterator exhausted, resetting it.")
|
||||||
|
iterator = iter(train_dataloader)
|
||||||
|
loss = batch_loss_fn(model,
|
||||||
|
iterator,
|
||||||
|
loss_functions,
|
||||||
|
device,
|
||||||
|
model_configuration)
|
||||||
|
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
loss.backward()
|
loss.backward()
|
||||||
@@ -237,20 +295,35 @@ def train_model(model: nn.Module,
|
|||||||
writer.flush()
|
writer.flush()
|
||||||
|
|
||||||
avg_train_loss = total_train_loss / len(train_dataloader)
|
avg_train_loss = total_train_loss / len(train_dataloader)
|
||||||
logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs} done. Train loss: {avg_train_loss:.4f}")
|
logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs} done. Local Train loss: {avg_train_loss:.4f}")
|
||||||
|
|
||||||
|
# sync before validation
|
||||||
|
if torch.distributed.is_initialized():
|
||||||
|
torch.distributed.barrier()
|
||||||
|
|
||||||
# Validation
|
# Validation
|
||||||
|
# if local_rank == 0 or not torch.distributed.is_initialized():
|
||||||
model.eval()
|
model.eval()
|
||||||
total_val_loss = 0
|
total_val_loss = 0
|
||||||
logger.info(f"Rank {local_rank}: Validation")
|
logger.info(f"Rank {local_rank}: Validation")
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
val_iter = iter(val_dataloader)
|
val_iter = iter(val_dataloader)
|
||||||
for step in tqdm(range(len(val_dataloader))):
|
for step in tqdm(range(val_length)):
|
||||||
loss = batch_loss_fn(model,
|
try:
|
||||||
val_iter,
|
loss = batch_loss_fn(model,
|
||||||
loss_functions,
|
val_iter,
|
||||||
device,
|
loss_functions,
|
||||||
model_configuration)
|
device,
|
||||||
|
model_configuration)
|
||||||
|
except StopIteration:
|
||||||
|
# if the iterator is exhausted, reset it
|
||||||
|
logger.info(f"Rank {local_rank}: Iterator exhausted, resetting it.")
|
||||||
|
val_iter = iter(val_dataloader)
|
||||||
|
loss = batch_loss_fn(model,
|
||||||
|
val_iter,
|
||||||
|
loss_functions,
|
||||||
|
device,
|
||||||
|
model_configuration)
|
||||||
|
|
||||||
total_val_loss += loss.item()
|
total_val_loss += loss.item()
|
||||||
|
|
||||||
@@ -259,11 +332,20 @@ def train_model(model: nn.Module,
|
|||||||
avg_val_loss = None
|
avg_val_loss = None
|
||||||
else:
|
else:
|
||||||
avg_val_loss = total_val_loss / len(val_dataloader)
|
avg_val_loss = total_val_loss / len(val_dataloader)
|
||||||
|
# else:
|
||||||
|
# avg_val_loss = None
|
||||||
|
# logger.info(f"Rank {local_rank}: Validation skipped, using 0 as validation loss.")
|
||||||
|
|
||||||
|
# sync before logging
|
||||||
|
if torch.distributed.is_initialized():
|
||||||
|
torch.distributed.barrier()
|
||||||
|
|
||||||
# publish validation loss and wait for other gpus
|
# publish validation loss and wait for other gpus
|
||||||
if torch.distributed.is_initialized():
|
if torch.distributed.is_initialized():
|
||||||
if avg_val_loss is not None:
|
if avg_val_loss is not None:
|
||||||
avg_val_loss_global = torch.tensor(avg_val_loss, device=device, dtype=torch.float32)
|
avg_val_loss_global = torch.tensor(avg_val_loss, device=device, dtype=torch.float32)
|
||||||
|
else:
|
||||||
|
avg_val_loss_global = torch.tensor(0.0, device=device, dtype=torch.float32)
|
||||||
torch.distributed.all_reduce(avg_val_loss_global)
|
torch.distributed.all_reduce(avg_val_loss_global)
|
||||||
avg_val_loss_global /= torch.distributed.get_world_size()
|
avg_val_loss_global /= torch.distributed.get_world_size()
|
||||||
|
|
||||||
@@ -276,7 +358,7 @@ def train_model(model: nn.Module,
|
|||||||
|
|
||||||
# only rank 0 checks for early stopping
|
# only rank 0 checks for early stopping
|
||||||
if local_rank == 0:
|
if local_rank == 0:
|
||||||
logger.info(f"Rank {local_rank}: Epoch {epoch}/{num_epochs} done. Val loss: {avg_val_loss:.4f}")
|
logger.info(f"Rank {local_rank}: Overall val loss: {avg_val_loss_global:.4f}")
|
||||||
writer.add_scalar("Loss/Train_Epoch", avg_train_loss_global, epoch)
|
writer.add_scalar("Loss/Train_Epoch", avg_train_loss_global, epoch)
|
||||||
writer.add_scalar("Loss/Val_Epoch", avg_val_loss_global, epoch)
|
writer.add_scalar("Loss/Val_Epoch", avg_val_loss_global, epoch)
|
||||||
writer.flush()
|
writer.flush()
|
||||||
@@ -290,7 +372,10 @@ def train_model(model: nn.Module,
|
|||||||
epochs_no_improve = 0
|
epochs_no_improve = 0
|
||||||
# torch.save(model.state_dict(), os.path.join(model_configuration["id"], "model.pt"))
|
# torch.save(model.state_dict(), os.path.join(model_configuration["id"], "model.pt"))
|
||||||
save_fn = model_configuration["model_save_fn"]
|
save_fn = model_configuration["model_save_fn"]
|
||||||
save_fn(model, training_configuration)
|
if torch.distributed.is_initialized():
|
||||||
|
save_fn(model.module, training_configuration)
|
||||||
|
else:
|
||||||
|
save_fn(model, training_configuration)
|
||||||
else:
|
else:
|
||||||
epochs_no_improve += 1
|
epochs_no_improve += 1
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
Reference in New Issue
Block a user