diff --git a/model/train.py b/model/train.py new file mode 100644 index 0000000..3855c19 --- /dev/null +++ b/model/train.py @@ -0,0 +1,166 @@ +import json +import os +from argparse import ArgumentParser +from datetime import datetime, timezone +from pathlib import Path + +import torch +from dotenv import load_dotenv +from ignite.engine import Events, create_supervised_evaluator, create_supervised_trainer +from ignite.handlers import global_step_from_engine +from ignite.handlers.checkpoint import Checkpoint +from ignite.handlers.param_scheduler import create_lr_scheduler_with_warmup +from ignite.handlers.tensorboard_logger import TensorboardLogger +from ignite.handlers.tqdm_logger import ProgressBar +from ignite.metrics import Average, Loss, RunningAverage +from src.utils import VCDADataset, get_class +from torch import nn +from torch.optim.lr_scheduler import ExponentialLR +from torch.utils.data import DataLoader + +from shared.data_store import connect_minio + + +def parse_arguments(): + parser = ArgumentParser() + parser.add_argument('-c', '--config', type=Path, required=True) + parser.add_argument('-d', '--device', type=str, default='cpu') + parser.add_argument('-b', '--batch-size', type=int, default=32) + parser.add_argument('--epochs', type=int, default=50) + parser.add_argument('--epoch-length', type=int, default=128) + parser.add_argument('--checkpoint', type=Path) + parser.add_argument('--lr', type=float, default=1e-5) + parser.add_argument('--lr-gamma', type=float, default=0.95) + parser.add_argument('--lr-warmup-start', type=float, default=1e-8) + parser.add_argument('--lr-warmup-duration', type=int, default=5) + parser.add_argument('--train-seed', type=int, default=1312) + parser.add_argument('--val-seed', type=int, default=1313) + parser.add_argument('--loader-workers', type=int, default=8) + return parser.parse_args() + + +args = parse_arguments() + +load_dotenv('server.env') + +# read model config +with open(args.config) as fh: + config = json.load(fh) +model_name = args.config.stem +run_name = datetime.now(timezone.utc).strftime('%Y-%m-%d-%H%M') + '_' + model_name +device = torch.device(args.device) + +# create model +model_class = get_class(config['backbone']['class']) +model = model_class().to(device) +optimizer = torch.optim.Adam(model.parameters(), lr=args.lr) +loss_fn = nn.CrossEntropyLoss() + +# create datasets and loaders +minio_client = connect_minio() +with open('model/src/dataset/train.csv', encoding='utf-8') as fh: + train_data_name_list = fh.read().split('\n') +train_dataset = VCDADataset( + minio_client=minio_client, + data_name_list=train_data_name_list, +) +train_loader = DataLoader(dataset=train_dataset, num_workers=args.loader_workers) + +with open('model/src/dataset/val.csv', encoding='utf-8') as fh: + val_data_name_list = fh.read().split('\n') +val_dataset = VCDADataset(minio_client=minio_client, data_name_list=val_data_name_list) +val_loader = DataLoader(dataset=val_dataset, num_workers=args.loader_workers) + +# create trainer and evaluator +trainer = create_supervised_trainer( + model=model, + optimizer=optimizer, + loss_fn=loss_fn, + device=device, +) +Average(output_transform=lambda x: x).attach(trainer, name='loss') +RunningAverage(output_transform=lambda x: x, alpha=0.5).attach( + trainer, + 'running_avg_loss', +) +ProgressBar(desc='Train', ncols=80).attach(trainer, ['running_avg_loss']) + +val_metrics = { + 'loss': Loss(loss_fn=loss_fn, device=device), +} +evaluator = create_supervised_evaluator( + model=model, + metrics=val_metrics, + device=device, +) +ProgressBar(desc='Val', ncols=80).attach(evaluator) + +# create lr scheduler +lr_scheduler = ExponentialLR( + optimizer=optimizer, + gamma=args.lr_gamma, +) +lr_handler = create_lr_scheduler_with_warmup( + lr_scheduler=lr_scheduler, + warmup_start_value=args.lr_warmup_start, + warmup_duration=args.lr_warmup_duration, + warmup_end_value=args.lr, +) +trainer.add_event_handler(Events.EPOCH_STARTED, lr_handler) + + +# log metrics to TensorBoard +@trainer.on(Events.EPOCH_COMPLETED) +def evaluate(): + evaluator.run(val_loader, epoch_length=args.val_epoch_length) + + +tb_logger = TensorboardLogger(log_dir=f'runs/logs/{run_name}') +tb_logger.attach_opt_params_handler( + engine=trainer, + event_name=Events.EPOCH_COMPLETED, + optimizer=optimizer, +) +for tag, engine in [('train', trainer), ('val', evaluator)]: + tb_logger.attach_output_handler( + engine, + event_name=Events.EPOCH_COMPLETED, + tag=tag, + metric_names='all', + global_step_transform=global_step_from_engine(trainer), + ) + +# Set up checkpoint saving +to_save = { + 'model': model, + 'optimizer': optimizer, + 'trainer': trainer, + 'lr_scheduler': lr_scheduler, + 'lr_handler': lr_handler, +} +checkpoint_handler = Checkpoint( + to_save, + f"runs/checkpoints/{run_name}", + n_saved=3, + filename_prefix='best', + score_function=lambda engine: -engine.state.metrics['loss'], + score_name='neg_val_loss', + global_step_transform=global_step_from_engine(trainer), +) +evaluator.add_event_handler(Events.COMPLETED, checkpoint_handler) + +if args.checkpoint: + Checkpoint.load_objects(to_load=to_save, checkpoint=str(args.checkpoint)) + +# save model config +os.makedirs('runs/configs/', exist_ok=True) +with open(f"runs/configs/{run_name}.json", 'w', encoding='utf-8') as fh: + json.dump(config, fh) + +# start training +trainer.run( + train_loader, + max_epochs=args.epochs, + epoch_length=args.epoch_length, +) +tb_logger.close()