File size: 7,821 Bytes
69494e8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 | import numpy as np
from tqdm import tqdm
from medpy import metric
from torch.cuda.amp import autocast as autocast
import torch
from utils import test_score
from colorama import Fore
import time
from sklearn.metrics import confusion_matrix
def train_one_epoch(train_loader,
model,
criterion,
optimizer,
scheduler,
epoch,
logger,
config,
MultiScaleLoss = True,
scaler=None):
'''
train model for one epoch
'''
stime = time.time()
model.train()
loss_list = []
for iter, data in enumerate(train_loader):
optimizer.zero_grad()
images, targets = data['image'], data['label']
images, targets = images.cuda(non_blocking=True).float().permute(0,3,1,2), targets.cuda(non_blocking=True).float()
if config.amp:
with autocast():
out, dec_outputs = model(images)
loss = criterion(out, targets)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
elif MultiScaleLoss is True:
out, dec_outputs = model(images)
loss = criterion(out, targets, dec_outputs)
loss.backward()
optimizer.step()
else:
out, dec_outputs = model(images)
loss = criterion(out, targets)
loss.backward()
optimizer.step()
loss_list.append(loss.item())
now_lr = optimizer.state_dict()['param_groups'][0]['lr']
if iter % config.print_interval == 0:
log_info = f'train: epoch {epoch}, iter:{iter}, loss: {loss.item():.4f}, lr: {now_lr}'
print(log_info)
logger.info(log_info)
scheduler.step()
mean_loss = np.mean(loss_list)
etime = time.time()
log_info = f'Finish one epoch train: epoch {epoch}, loss: {mean_loss:.4f}, time(s): {etime-stime:.2f}'
print(log_info)
logger.info(log_info)
return mean_loss
def val_one_epoch(test_loader, model, epoch, logger, config):
stime = time.time()
model.eval()
with torch.no_grad():
metric_list = 0.0
for data in tqdm(test_loader):
img, msk = data['image'], data['label']#B H W
metric_i = test_score(img, msk, model, classes=config.num_classes,
patch_size=[config.input_size_h, config.input_size_w])
metric_list += np.array(metric_i)
metric_list = metric_list[:, :-1] / metric_list[:, -1].reshape(-1, 1)
for i in range(1, config.num_classes):
logger.info('Mean class %d mean_dice %f mean_hd95 %f mean_recall %f mean_IOU %f mean_acc %f mean_spe %f' % (i, metric_list[i-1][0], metric_list[i-1][1], metric_list[i-1][2], metric_list[i-1][3], metric_list[i-1][4], metric_list[i-1][5]))
performance = np.mean(metric_list, axis=0)[0]
mean_hd95 = np.mean(metric_list, axis=0)[1]
mean_recall = np.mean(metric_list, axis=0)[2]
mean_IOU = np.mean(metric_list, axis=0)[3]
mean_acc = np.mean(metric_list, axis=0)[4]
mean_spe = np.mean(metric_list, axis=0)[5]
etime = time.time()
log_info = f'val epoch: {epoch}, mean_dice: {performance}, mean_hd95: {mean_hd95}, mean_recall: {mean_recall}, mean_IOU: {mean_IOU}, mean_acc: {mean_acc}, mean_spe: {mean_spe}, time(s): {etime-stime:.2f}'
print(log_info)
logger.info(log_info)
return performance, mean_hd95
def val_one_epochV2(test_loader, model, epoch, logger, config):
stime = time.time()
model.eval()
with torch.no_grad():
total_pred = []
total_target = []
for data in tqdm(test_loader):
img, msk = data['image'], data['label']#B H W
model.eval()
outputs, _ = model(img.permute(0, 3, 1, 2))
outputs = torch.argmax(torch.softmax(outputs, dim=1), dim=1).squeeze(0)
outputs = outputs.cpu().detach().numpy()#H W
msk = msk.cpu().detach().numpy()
total_pred.append(outputs)
total_target.append(msk)
dsc_avg = 0
hd95_avg = 0
sen_avg = 0
miou_avg = 0
acc_avg = 0
spe_avg = 0
TP_total = 0
TN_total = 0
FP_total = 0
FN_total = 0
for cur_num_classes in range(1, config.num_classes):
TP = 0
TN = 0
FP = 0
FN = 0
hd95_total = 0.0
num = 0
for pred_batch, target_batch in zip(total_pred, total_target):
for i in range(pred_batch.shape[0]):
pred = pred_batch[i]
target = target_batch[i]
pred_flat = pred.ravel()
target_flat = target.ravel()
TP += np.sum((pred_flat == cur_num_classes) & (target_flat == cur_num_classes))
TN += np.sum((pred_flat != cur_num_classes) & (target_flat != cur_num_classes))
FP += np.sum((pred_flat == cur_num_classes) & (target_flat != cur_num_classes))
FN += np.sum((pred_flat != cur_num_classes) & (target_flat == cur_num_classes))
TP_total += TP
TN_total += TN
FP_total += FP
FN_total += FN
pred = (pred == cur_num_classes)
target = (target == cur_num_classes)
if pred.sum() == 0 and target.sum() > 0:
hd95_total += config.input_size_h * 1.414
num +=1
elif pred.sum() > 0 and target.sum() == 0:
hd95_total += config.input_size_h * 1.414
num +=1
elif pred.sum() > 0 and target.sum() > 0:
hd95_total += metric.binary.hd95(pred, target)
num +=1
else:
hd95_total += 0
num += 1
epsilon = 1e-8
dsc = (2 * TP) / (2 * TP + FP + FN + epsilon)
dsc_avg += dsc
sen = TP / (TP + FN + epsilon)
sen_avg += sen
spe = TN / (TN + FP + epsilon)
spe_avg += spe
acc = (TP + TN) / (TP + TN + FP + FN + epsilon)
acc_avg += acc
iou_foreground = TP / (TP + FP + FN + epsilon)
miou = iou_foreground
miou_avg += miou
hd95 = hd95_total / num
hd95_avg += hd95
logger.info('Mean class %d mean_dice %f mean_hd95 %f mean_recall %f mean_IOU %f mean_acc %f mean_spe %f' % (cur_num_classes, dsc, hd95, sen, miou, acc, spe))
etime = time.time()
log_info = f'val epoch: {epoch}, mean_dice: {(2 * TP_total) / (2 * TP_total + FP_total + FN_total + epsilon)}, mean_hd95: {hd95_avg/ (config.num_classes - 1)}, mean_recall: {TP_total / (TP_total + FN_total + epsilon)}, mean_IOU: {TP_total / (TP_total + FP_total + FN_total + epsilon)}, mean_acc: {(TP_total + TN_total) / (TP_total + TN_total + FP_total + FN_total + epsilon)}, mean_spe: {TN_total / (TN_total + FP_total + epsilon)}, time(s): {etime-stime:.2f}'
log_info = f'val epoch: {epoch}, mean_dice: {dsc_avg / (config.num_classes - 1)}, mean_hd95: {hd95_avg/ (config.num_classes - 1)}, mean_recall: {sen_avg / (config.num_classes - 1)}, mean_IOU: {miou_avg / (config.num_classes - 1)}, mean_acc: {acc_avg / (config.num_classes - 1)}, mean_spe: {spe_avg / (config.num_classes - 1)}, time(s): {etime-stime:.2f}'
print(log_info)
logger.info(log_info)
return dsc, hd95
|