added model training script
This commit is contained in:
+166
@@ -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()
|
||||
Reference in New Issue
Block a user