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 torch import nn from torch.optim.lr_scheduler import ExponentialLR from torch.utils.data import DataLoader from model.src.utils import VCDADataset, get_class from shared.datastore 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, # type: ignore 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()