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: # get the inputs; data is a list of [inputs, labels] inputs, labels = data inputs = inputs.to(device) labels = labels.to(device) outputs = model(inputs) # calculate loss loss = criterion(outputs, labels) # zero the parameter gradients optimizer.zero_grad() # backpropagation loss.backward() # update parameters 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' # random.seed(seed) # python random generator # np.random.seed(seed) # numpy random generator 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)