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