stereoid's picture
Add files using upload-large-folder tool
af46737 verified
Raw
History Blame Contribute Delete
11.4 kB
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)