qic999's picture
Upload folder using huggingface_hub
28e6f98 verified
Raw
History Blame Contribute Delete
9.19 kB
import os
import sys
import logging
from skimage import io
from skimage import img_as_ubyte
from torch.utils.data import DataLoader
from networks.mynet import TwoBranch
from utils.option import args
from tqdm import tqdm
from utils.metric import nmse, psnr, ssim
from collections import defaultdict
from networks_time.mynet import DiffTwoBranch
test_data_path = args.root_path
snapshot_path = "model/" + args.exp + "/"
os.environ['CUDA_VISIBLE_DEVICES'] = args.gpu
def normalize_output(out_img):
out_img = (out_img - out_img.min())/(out_img.max() - out_img.min() + 1e-8)
return out_img
from frequency_diffusion.degradation.k_degradation import apply_tofre, apply_to_spatial
from utils.utils import *
DEBUG = False
use_kspace = args.use_kspace
use_time_model = args.use_time_model
num_timesteps = args.num_timesteps
image_size = args.image_size
snapshot_path=args.snapshot_path
from frequency_diffusion.degradation.k_degradation import get_ksu_kernel, apply_tofre, apply_to_spatial
kspace_masks = np.load(f"./dataloaders/example_mask/kspace_{args.ACCELERATIONS[0]}_mask.npy")
kspace_masks = torch.from_numpy(np.asarray(kspace_masks)).cuda()
kspace_masks = get_ksu_kernel(num_timesteps, image_size,
ksu_routine="LogSamplingRate",
accelerated_factor=args.ACCELERATIONS[0]
)
kspace_masks = torch.from_numpy(np.asarray(kspace_masks)).cuda()
print("kspace_masks shape: ", kspace_masks.shape)
@torch.no_grad()
def evaluate(model, data_loader, device, save_path):
os.makedirs(save_path, exist_ok=True)
model.eval()
nmse_meter, psnr_meter, ssim_meter = [], [], []
direct_nmse, direct_psnr, direct_ssim = [], [], []
output_dic = defaultdict(dict)
target_dic = defaultdict(dict)
input_dic = defaultdict(dict)
direct_recon_dic = defaultdict(dict)
flag=0
last_name='no'
print("len of data_loader: ", len(data_loader))
for data in tqdm(data_loader):
pd, pdfs, _ = data
name = os.path.basename(pdfs[4][0]).split('.')[0]
target = pdfs[1].to(device)
mean, std = pdfs[2], pdfs[3]
fname = pdfs[4]
slice_num = pdfs[5]
mean = mean.unsqueeze(1).unsqueeze(2).to(device)
std = std.unsqueeze(1).unsqueeze(2).to(device)
pd_img = pd[1].unsqueeze(1).to(device)
pdfs_img = pdfs[0].unsqueeze(1).to(device)
pdfs_img_origin = pdfs_img.clone()
# Degradation
if use_kspace:
b = pdfs_img.shape[0]
t = torch.randint(num_timesteps - 1, num_timesteps, (b,), device=device).long() # t-1
mask = kspace_masks[t]
fft, mask = apply_tofre(target.clone(), mask)
fft = fft * mask + 0.0
pdfs_img = apply_to_spatial(fft)
while t >= 0:
if use_time_model:
outputs = network(pdfs_img, pd_img)['img_out']
else:
outputs = network(pdfs_img, pd_img, t)['img_out']
if t == num_timesteps - 1:
direct_recon = outputs
if t == 0:
pdfs_img = outputs
else:
k_full = kspace_masks[-1]
faded_recon_sample_fre, k_full = apply_tofre(pdfs_img, k_full)
with torch.no_grad():
kt_sub_1 = kspace_masks[t-1] #get_kspace_kernels(t - 2).cuda()
kt = kspace_masks[t] #self.get_kspace_kernels(t - 1).cuda() # last one
k_residual = kt_sub_1 - kt
recon_sample_fre, k_residual = apply_tofre(outputs, k_residual)
fre_amend = recon_sample_fre * k_residual
faded_recon_sample_fre = faded_recon_sample_fre + fre_amend # * (1-k_residual)
# faded_recon_sample_fre = faded_recon_sample_fre * kt_sub_1
outputs = apply_to_spatial(faded_recon_sample_fre)
pdfs_img = outputs
t = t-1
else:
outputs = network(pdfs_img, pd_img)['img_out']
outputs = outputs.squeeze(1)
direct_recon = direct_recon.squeeze(1)
outputs_save = outputs[0].cpu().clone().numpy()/6.0
outputs_save = np.clip(outputs_save, a_min=-1, a_max=1)
target_save = target[0].cpu().clone().numpy()/6.0
in_save = pdfs_img_origin[0][0].cpu().clone().numpy()/6.0
# Not sure if it was correct to convert to ubyte
outputs_save = img_as_ubyte(outputs_save)
target_save = img_as_ubyte(target_save)
in_save = img_as_ubyte(in_save)
io.imsave(save_path + str(name) + '_' + str(slice_num[0].cpu().numpy()) + '.png', target_save)
io.imsave(save_path + str(name) + '_' + str(slice_num[0].cpu().numpy()) + '_in.png', in_save)
io.imsave(save_path + str(name) + '_' + str(slice_num[0].cpu().numpy()) + '_out.png', outputs_save)
outputs = outputs * std + mean
target = target * std + mean
inputs = pdfs_img_origin.squeeze(1) * std + mean
direct_recon = direct_recon * std + mean
output_dic[fname[0]][slice_num[0]] = outputs[0]
target_dic[fname[0]][slice_num[0]] = target[0]
input_dic[fname[0]][slice_num[0]] = inputs[0]
direct_recon_dic[fname[0]][slice_num[0]] = direct_recon[0]
# print("target/outputs shape: ", target.shape, outputs.shape)
our_nmse = nmse(target[0].cpu().numpy(), outputs[0].cpu().numpy())
our_psnr = psnr(target[0].cpu().numpy(), outputs[0].cpu().numpy())
our_ssim = ssim(target[0].cpu().numpy(), outputs[0].cpu().numpy())
# print('name:{}, slice:{}, nmse:{}, psnr:{}, ssim:{}'.format(name, slice_num[0], our_nmse, our_psnr, our_ssim))
for name in output_dic.keys():
f_output = torch.stack([v for _, v in output_dic[name].items()])
f_target = torch.stack([v for _, v in target_dic[name].items()])
our_nmse = nmse(f_target.cpu().numpy(), f_output.cpu().numpy())
our_psnr = psnr(f_target.cpu().numpy(), f_output.cpu().numpy())
our_ssim = ssim(f_target.cpu().numpy(), f_output.cpu().numpy())
nmse_meter.append(our_nmse)
psnr_meter.append(our_psnr)
ssim_meter.append(our_ssim)
direct_nmse.append(nmse(f_target.cpu().numpy(), torch.stack([v for _, v in direct_recon_dic[name].items()]).cpu().numpy()))
direct_psnr.append(psnr(f_target.cpu().numpy(), torch.stack([v for _, v in direct_recon_dic[name].items()]).cpu().numpy()))
direct_ssim.append(ssim(f_target.cpu().numpy(), torch.stack([v for _, v in direct_recon_dic[name].items()]).cpu().numpy()))
nmse_meter_score = np.array(nmse_meter)
psnr_meter_score = np.array(psnr_meter)
ssim_meter_score = np.array(ssim_meter)
direct_nmse_score = np.array(direct_nmse)
direct_psnr_score = np.array(direct_psnr)
direct_ssim_score = np.array(direct_ssim)
print("===> Evaluate Metric <===")
print("Direct Results")
print("-" * 36)
print(f"NMSE: {np.mean(direct_nmse_score) * 100:.4f} ± {np.std(direct_nmse_score) * 100:.4f}")
print(f"PSNR: {np.mean(direct_psnr_score):.4f} ± {np.std(direct_psnr_score):.4f}")
print(f"SSIM: {np.mean(direct_ssim_score):.4f} ± {np.std(direct_ssim_score):.4f}")
print("-" * 36)
print("===> Evaluate Metric <===")
print("Results")
print("-" * 36)
print(f"NMSE: {np.mean(nmse_meter_score) * 100:.4f} ± {np.std(nmse_meter_score) * 100:.4f}")
print(f"PSNR: {np.mean(psnr_meter_score):.4f} ± {np.std(psnr_meter_score):.4f}")
print(f"SSIM: {np.mean(ssim_meter_score):.4f} ± {np.std(ssim_meter_score):.4f}")
print("-" * 36)
print(f"Save Path: {save_path}")
model.train()
return {'NMSE': np.mean(nmse_meter_score), 'PSNR': np.mean(psnr_meter_score), 'SSIM': np.mean(ssim_meter_score)}
from dataloaders.fastmri import build_dataset
if __name__ == "__main__":
if use_time_model:
network = DiffTwoBranch(args).cuda()
else:
network = TwoBranch(args).cuda()
device = torch.device('cuda')
network.to(device)
if len(args.gpu.split(',')) > 1:
network = nn.DataParallel(network)
db_test = build_dataset(args, mode='val', use_kspace=use_kspace)
testloader = DataLoader(db_test, batch_size=1, shuffle=False, num_workers=4, pin_memory=True)
if args.phase == 'test':
save_mode_path = os.path.join(snapshot_path, 'best_checkpoint.pth')
# save_mode_path = os.path.join(snapshot_path, 'iter_100000.pth')
print('load weights from ' + save_mode_path)
checkpoint = torch.load(save_mode_path)
weights_dict = {}
for k, v in checkpoint['network'].items():
new_k = k.replace('module.', '') if 'module' in k else k
weights_dict[new_k] = v
network.load_state_dict(weights_dict)
network.eval()
eval_result = evaluate(network, testloader, device, save_path = snapshot_path + '/result_case/')