| import os
|
| from tqdm import tqdm
|
| import time
|
| import argparse
|
| import json
|
| import torch
|
| from torch.utils.data import DataLoader
|
| from torch import nn, optim
|
| from torchvision import transforms
|
| from torchmetrics.classification import \
|
| MulticlassF1Score, MulticlassAccuracy, MulticlassPrecision, MulticlassRecall, MulticlassAUROC
|
| from torch.utils.tensorboard import SummaryWriter
|
|
|
| import datasets
|
| from CNN import CNN
|
|
|
| seed = 123
|
|
|
| def do_train(model, train_loader, criterion, optimizer, epoch, output_dir, best_metric, device,
|
| dev_loader = None, resume = False):
|
|
|
| start_epoch = 0
|
| best_metric_value = -1
|
| ckpt = None
|
| if resume and os.path.exists(os.path.join(output_dir, 'ckpt', 'last.pth')):
|
| ckpt = os.path.join(output_dir, 'ckpt', 'last.pth')
|
|
|
| if ckpt:
|
| ckpt = torch.load(ckpt, map_location=device)
|
| assert ckpt['best_metric'] == best_metric, 'best metric mismatch'
|
| best_metric_value = ckpt['best_metric_value']
|
| model.load_state_dict(ckpt['model_state_dict'])
|
| optimizer.load_state_dict(ckpt['optimizer_state_dict'])
|
| start_epoch = ckpt['epoch'] + 1
|
| print(f'loaded ckpt from {os.path.join(output_dir, "ckpt", "last.pth")}, starting from epoch {start_epoch}')
|
|
|
| os.makedirs(os.path.join(output_dir, 'ckpt'), exist_ok=True)
|
| model.to(device)
|
| total_step = start_epoch * len(train_loader)
|
|
|
| writer = SummaryWriter(output_dir, flush_secs=10)
|
|
|
| for cur_epoch in range(start_epoch, epoch):
|
|
|
| pbar = tqdm(enumerate(train_loader, 0), total=len(train_loader))
|
| pbar.desc = '[%s: epoch %2d, batch %3d] loss: %.5f' % \
|
| (output_dir, cur_epoch, 0, 0)
|
|
|
| model.train()
|
| for i, data in pbar:
|
|
|
| inputs, labels = data
|
| inputs = inputs.to(device)
|
| labels = labels.to(device)
|
|
|
| outputs = model(inputs)
|
|
|
|
|
| loss = criterion(outputs, labels)
|
|
|
|
|
| optimizer.zero_grad()
|
|
|
| loss.backward()
|
|
|
|
|
| optimizer.step()
|
|
|
| running_loss = loss.item()
|
|
|
| pbar.desc = '[%s: epoch %2d, batch %3d] loss: %.5f' % \
|
| (output_dir, cur_epoch, i + 1, running_loss)
|
|
|
| total_step += train_loader.batch_size
|
| writer.add_scalar("train/Loss", running_loss, global_step = total_step)
|
|
|
| pbar.close()
|
|
|
| if dev_loader:
|
| print("evaluating")
|
| eval_loss, _, _, _, metrics = do_eval(model, dev_loader, device, loss_criterion = criterion)
|
|
|
| writer.add_scalar("dev/Loss", eval_loss, global_step = cur_epoch)
|
| writer.add_scalar("dev/F1", metrics['f1']['macro'], global_step = cur_epoch)
|
| writer.add_scalar("dev/Accuracy", metrics['acc']['macro'], global_step = cur_epoch)
|
| writer.add_scalar("dev/Precision", metrics['precision']['macro'], global_step = cur_epoch)
|
| writer.add_scalar("dev/Recall", metrics['recall']['macro'], global_step = cur_epoch)
|
| writer.add_scalar("dev/AUROC", metrics['auroc'], global_step = cur_epoch)
|
|
|
| if best_metric_value < metrics[best_metric]['macro']:
|
| torch.save({
|
| 'best_metric': best_metric,
|
| 'best_metric_value': metrics[best_metric],
|
| 'epoch': cur_epoch,
|
| 'model_state_dict': model.state_dict(),
|
| 'optimizer_state_dict': optimizer.state_dict()}, os.path.join(output_dir, 'ckpt', 'best.pth'))
|
| best_metric_value = metrics[best_metric]['macro']
|
|
|
| torch.save({
|
| 'best_metric': best_metric,
|
| 'best_metric_value': metrics[best_metric],
|
| 'epoch': cur_epoch,
|
| 'model_state_dict': model.state_dict(),
|
| 'optimizer_state_dict': optimizer.state_dict()}, os.path.join(output_dir, 'ckpt', 'last.pth'))
|
|
|
| writer.close()
|
|
|
|
|
| def do_eval(model, eval_loader, device, ckpt_path = None, loss_criterion = None):
|
| torch.cuda.empty_cache()
|
|
|
| if ckpt_path:
|
| ckpt = torch.load(ckpt_path)
|
| model.load_state_dict(ckpt['model_state_dict'])
|
|
|
| model.to(device)
|
|
|
| batch_count = 0
|
| loss = 0.0
|
| pred_all = torch.tensor([]).to(device)
|
| label_all = torch.tensor([]).to(device)
|
| output_all = torch.tensor([]).to(device)
|
|
|
| model.eval()
|
| with torch.no_grad():
|
| for data in tqdm(eval_loader):
|
| images, labels = data
|
| images = images.to(device)
|
| labels = labels.to(device)
|
| outputs = model(images)
|
| if loss_criterion:
|
| loss += loss_criterion(outputs, labels).item()
|
| batch_count += 1
|
| pred = outputs.argmax(dim=1)
|
| pred_all = torch.cat((pred_all, pred), dim=0)
|
| label_all = torch.cat((label_all, labels), dim=0)
|
| output_all = torch.cat((output_all, outputs), dim=0)
|
|
|
| if loss_criterion and batch_count > 0:
|
| loss /= batch_count
|
|
|
| ma_f1_metric = MulticlassF1Score(model.num_classes, average='macro').to(device)
|
| mi_f1_metric = MulticlassF1Score(model.num_classes, average='micro').to(device)
|
| ma_p_metric = MulticlassPrecision(model.num_classes, average='macro').to(device)
|
| mi_p_metric = MulticlassPrecision(model.num_classes, average='micro').to(device)
|
| ma_r_metric = MulticlassRecall(model.num_classes, average='macro').to(device)
|
| mi_r_metric = MulticlassRecall(model.num_classes, average='micro').to(device)
|
| ma_acc_metric = MulticlassAccuracy(model.num_classes, average='macro').to(device)
|
| mi_acc_metric = MulticlassAccuracy(model.num_classes, average='micro').to(device)
|
| auroc_metric = MulticlassAUROC(model.num_classes, average='macro', thresholds=10).to(device)
|
|
|
| label_all = label_all.long()
|
| metrics = {}
|
| metrics['f1'] = {'macro': ma_f1_metric(pred_all, label_all).item(), 'micro': mi_f1_metric(pred_all, label_all).item()}
|
| metrics['acc'] = {'macro': ma_acc_metric(pred_all, label_all).item(), 'micro': mi_acc_metric(pred_all, label_all).item()}
|
| metrics['precision'] = {'macro': ma_p_metric(pred_all, label_all).item(), 'micro': mi_p_metric(pred_all, label_all).item()}
|
| metrics['recall'] = {'macro': ma_r_metric(pred_all, label_all).item(), 'micro': mi_r_metric(pred_all, label_all).item()}
|
| metrics['auroc'] = auroc_metric(output_all, label_all).item()
|
|
|
| return loss, label_all, pred_all, output_all, metrics
|
|
|
|
|
| def predict(model, pred_loader, device, ckpt_path = None):
|
|
|
| if ckpt_path:
|
| ckpt = torch.load(ckpt_path)
|
| model.load_state_dict(ckpt['model_state_dict'])
|
|
|
| model.to(device)
|
|
|
| pred_all = torch.tensor([]).to(device)
|
| output_all = torch.tensor([]).to(device)
|
|
|
| model.eval()
|
| start_time = time.time()
|
| with torch.no_grad():
|
| for data in tqdm(pred_loader):
|
| images = data
|
| images = images.to(device)
|
| outputs = model(images)
|
| pred = outputs.argmax(dim=1)
|
| pred_all = torch.cat((pred_all, pred), dim=0)
|
| output_all = torch.cat((output_all, outputs), dim=0)
|
| end_time = time.time()
|
|
|
| return pred_all, output_all, end_time - start_time
|
|
|
|
|
|
|
| def main(args):
|
| assert args.best_metric in ['acc', 'auroc', 'f1', 'precision', 'recall'], 'best metric must be one of acc, auroc, f1, precision, recall'
|
|
|
|
|
|
|
|
|
| torch.manual_seed(seed)
|
| torch.cuda.manual_seed_all(seed)
|
|
|
| torch.backends.cudnn.deterministic = True
|
| torch.backends.cudnn.benchmark = False
|
|
|
| resize = (64, 64)
|
|
|
| trans = [
|
| transforms.ToTensor(),
|
| transforms.Resize(resize, antialias=True)
|
| ]
|
|
|
| data_transform = transforms.Compose(trans)
|
|
|
| model = CNN(args.task, softmax=False)
|
| criterion = nn.CrossEntropyLoss()
|
| optimizer = optim.Adadelta(model.parameters(), lr=args.lr)
|
|
|
| train_dataset = datasets.TrainDataset(args.dataset_root, transform=data_transform)
|
| train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=8)
|
| dev_dataset = datasets.TrainDataset(args.dataset_root, is_test=True, transform=data_transform)
|
| dev_loader = DataLoader(dev_dataset, batch_size=args.batch_size, shuffle=False, num_workers=8)
|
|
|
| os.environ['CUDA_VISIBLE_DEVICES'] = args.device
|
| device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
|
|
| os.makedirs(args.output_dir, exist_ok=True)
|
| T1 = time.time()
|
| do_train(model, train_loader, criterion, optimizer, args.epoch, args.output_dir, args.best_metric, device,
|
| dev_loader = dev_loader,
|
| resume = True)
|
| T2 = time.time()
|
| print('Time elapsed: %.5f s' % (T2-T1))
|
|
|
| with open(os.path.join(args.output_dir, 'config.txt'), 'w') as f:
|
| f.write(f'batch_size: {args.batch_size}\n')
|
| f.write(f'epoch: {args.epoch}\n')
|
| f.write(f'lr: {args.lr}\n')
|
| f.write(f'best_metric: {args.best_metric}\n')
|
|
|
| pred_dataset = datasets.PredDataset(args.dataset_root, os.path.join(args.output_dir, 'first_stage.json'), transform=data_transform)
|
| pred_loader = DataLoader(pred_dataset, batch_size=args.batch_size, shuffle=False, num_workers=8)
|
|
|
| print('predicting')
|
| model = CNN(args.task, softmax=True)
|
| test_pred, test_output, inference_time = predict(model, pred_loader, device, ckpt_path = os.path.join(args.output_dir, 'ckpt', 'best.pth'))
|
|
|
| with open(os.path.join(args.output_dir, 'first_stage.json'), 'r') as f:
|
| first_stage = json.load(f)
|
| for i, pred in enumerate(test_pred):
|
| first_stage[i]['category_id'] = int(pred.item() + 1)
|
| first_stage[i]['score'] = test_output[i][first_stage[i]['category_id'] - 1].item()
|
|
|
| with open(os.path.join(args.output_dir, 'results.json'), 'w') as f:
|
| json.dump(first_stage, f)
|
|
|
| with open(os.path.join(args.output_dir, 'inference_time.txt'), 'w') as f:
|
| f.write(f'total inference time: {inference_time} s\n')
|
|
|
|
|
|
|
| if __name__ == '__main__':
|
| parser = argparse.ArgumentParser()
|
| parser.add_argument('--task', type=str, required=True, help='task type')
|
| parser.add_argument('--output_dir', type=str, required=True, help='output directory')
|
| parser.add_argument('--dataset_root', type=str, required=True, help='dataset root directory')
|
| parser.add_argument('--batch_size', type=int, default=64, help='batch size')
|
| parser.add_argument('--epoch', type=int, default=30, help='epoch')
|
| parser.add_argument('--lr', type=float, default=0.02, help='learning rate')
|
| parser.add_argument('--resume', action='store_true', help='resume training')
|
| parser.add_argument('--best_metric', type=str, default='acc', help='metric to determine best ckpt')
|
| parser.add_argument('--device', type=str, default='0', help='device id')
|
| args = parser.parse_args()
|
| main(args)
|
|
|