| 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 |
|
|
| os.environ['CUDA_VISIBLE_DEVICES'] = args.gpu |
| use_new_dataloader = True |
|
|
|
|
| |
|
|
|
|
| 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, apply_ksu_kernel |
| from utils.utils import * |
|
|
|
|
| num_timesteps = args.num_timesteps |
| image_size = args.image_size |
| distortion_sigma = 10 / 255 |
| use_kspace = args.use_kspace |
| use_time_model = args.use_time_model |
| DEBUG = args.DEBUG |
| snapshot_path=args.snapshot_path |
|
|
| kspace_masks = np.load(f"./dataloaders/example_mask/m4raw_{args.ACCELERATIONS[0]}_mask.npy") |
| kspace_masks = torch.from_numpy(np.asarray(kspace_masks)).cuda() |
|
|
| test_sample = args.test_sample |
| frequency_distortion = True |
|
|
| @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 = [] |
| nmse_meter_all = [] |
| psnr_meter_all = [] |
| ssim_meter_all = [] |
| output_dic = {} |
| target_dic = {} |
| input_dic = {} |
|
|
| flag = 0 |
| last_name = 'no' |
|
|
| print("len of data_loader: ", len(data_loader)) |
|
|
| for sampled_batch in tqdm(data_loader): |
| t1_img, t1_in = sampled_batch['t1'], sampled_batch['t1_in'] |
| t2_img, t2_in = sampled_batch['t2'], sampled_batch['t2_in'] |
|
|
| t1_img = t1_img.to(device) |
| t1_in = t1_in.to(device) |
| t2_img = t2_img.to(device) |
| t2_in = t2_in.to(device) |
|
|
| mean, std = sampled_batch['t2_mean'], sampled_batch['t2_std'] |
|
|
| name = sampled_batch['fname'] |
| fname = [name] |
| slice_num = sampled_batch['slice'] |
|
|
| mean = mean.unsqueeze(1).unsqueeze(2).to(device) |
| std = std.unsqueeze(1).unsqueeze(2).to(device) |
|
|
| t2_in_origin = t2_in.clone() |
|
|
| |
| if use_kspace: |
| b = 1 |
| t = torch.randint(num_timesteps - 1, num_timesteps, (b,), device=device).long() |
| mask = kspace_masks[t] |
| fft, mask = apply_tofre(t2_in.clone(), mask) |
| fft = fft * mask + 0.0 |
| t2_in = apply_to_spatial(fft) |
| t2_in_origin = t2_in.clone() |
|
|
| while t >= 0: |
| |
| if use_time_model: |
| outputs = model(t2_in, t1_img, t)['img_out'] |
| else: |
| outputs = model(t2_in, t1_img)['img_out'] |
|
|
| if t == 0: |
| mask = kspace_masks[0] |
| t2_in = outputs |
|
|
| else: |
| if test_sample == "Ksample": |
| k_full = kspace_masks[-1] |
| faded_recon_sample_fre, k_full = apply_tofre(t2_in, k_full) |
|
|
| with torch.no_grad(): |
|
|
| kt_sub_1 = kspace_masks[t - 1] |
| kt = kspace_masks[t] |
| 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 |
|
|
| outputs = apply_to_spatial(faded_recon_sample_fre) |
| t2_in = outputs |
|
|
|
|
| elif test_sample == "ColdDiffusion": |
| with torch.no_grad(): |
|
|
| kt_sub_1 = kspace_masks[t - 1] |
| kt = kspace_masks[t] |
|
|
| x_t_hat = apply_ksu_kernel(outputs, kt) |
| x_t_sub_1_hat = apply_ksu_kernel(outputs, kt_sub_1) |
|
|
| outputs = t2_in - x_t_hat + x_t_sub_1_hat |
|
|
| t2_in = outputs |
|
|
| elif test_sample == "DDPM": |
|
|
| with torch.no_grad(): |
|
|
| kt_sub_1 = kspace_masks[t - 1] |
| outputs = apply_ksu_kernel(kt_sub_1, kt_sub_1) |
| t2_in = outputs |
|
|
|
|
|
|
|
|
| t = t - 1 |
|
|
| else: |
| outputs = model(t2_in, t1_img)['img_out'] |
|
|
| |
| |
|
|
| target = t2_img.clone().squeeze(1) * std + mean |
| inputs = t2_in_origin.clone().squeeze(1) * std + mean |
| outputs_save = outputs.clone().squeeze(1) * std + mean |
|
|
| outputs_save = outputs_save.cpu().numpy() |
| |
| target_save = target.cpu().numpy() |
| in_save = inputs.cpu().numpy() |
|
|
| _min, _max = target_save.min(), target_save.max() |
| target_save = (((target_save - _min) / (_max - _min)) * 255).astype(np.uint8) |
| in_save = (((in_save - _min) / (_max - _min)) * 255).astype(np.uint8) |
| outputs_save = (((outputs_save - _min) / (_max - _min)) * 255).astype(np.uint8) |
|
|
| |
| outputs_save = img_as_ubyte(outputs_save) |
| target_save = img_as_ubyte(target_save) |
| in_save = img_as_ubyte(in_save) |
|
|
| |
| |
| |
|
|
| if len(outputs_save.shape) > 3: |
| outputs_save = outputs_save.squeeze(0) |
| target_save = target_save.squeeze(0) |
| in_save = in_save.squeeze(0) |
|
|
| if len(outputs_save.shape) > 3: |
| outputs_save = outputs_save.squeeze(0) |
| target_save = target_save.squeeze(0) |
| in_save = in_save.squeeze(0) |
|
|
| name = name[0].numpy() |
| name_int = int(name) |
|
|
| 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.squeeze(1) * std + mean |
| target = t2_img.squeeze(1) * std + mean |
| inputs = t2_in_origin.squeeze(1) * std + mean |
|
|
| if name_int not in output_dic.keys(): |
| output_dic[name_int] = [] |
| target_dic[name_int] = [] |
| input_dic[name_int] = [] |
|
|
| output_dic[name_int].append(outputs[0]) |
| target_dic[name_int].append(target[0]) |
| input_dic[name_int].append(inputs[0]) |
|
|
| |
| 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)) |
|
|
| nmse_meter_all.append(our_nmse) |
| psnr_meter_all.append(our_psnr) |
| ssim_meter_all.append(our_ssim) |
| |
|
|
| for name in output_dic.keys(): |
| print("name: ", name, len(output_dic[name])) |
| |
| |
| f_output = torch.stack(list(output_dic[name])) |
| f_target = torch.stack(list(target_dic[name])) |
|
|
| print("f_output shape: ", f_output.shape) |
|
|
| if len(f_output.shape) > 3: |
| f_output = f_output.squeeze(1) |
| f_target = f_target.squeeze(1) |
|
|
| 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) |
|
|
| nmse_meter_score = np.array(nmse_meter) |
| psnr_meter_score = np.array(psnr_meter) |
| ssim_meter_score = np.array(ssim_meter) |
|
|
| nmse_meter_all_score = np.array(nmse_meter_all) |
| psnr_meter_all_score = np.array(psnr_meter_all) |
| ssim_meter_all_score = np.array(ssim_meter_all) |
|
|
| print("===> Evaluate Metric <===") |
| print("Results") |
| print("-" * 36) |
| print(f"{test_sample} NMSE: {np.mean(nmse_meter_score) * 100:.4f} ± {np.std(nmse_meter_score) * 100:.4f}") |
| print(f"{test_sample} PSNR: {np.mean(psnr_meter_score):.4f} ± {np.std(psnr_meter_score):.4f}") |
| print(f"{test_sample} SSIM: {np.mean(ssim_meter_score):.4f} ± {np.std(ssim_meter_score):.4f}") |
| print("-" * 36) |
| print(f"All NMSE: {np.mean(nmse_meter_all_score) * 100:.4f} ± {np.std(nmse_meter_all_score) * 100:.4f}") |
| print(f"All PSNR: {np.mean(psnr_meter_all_score):.4f} ± {np.std(psnr_meter_all_score):.4f}") |
| print(f"All SSIM: {np.mean(ssim_meter_all_score):.4f} ± {np.std(ssim_meter_all_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.m4raw_std_dataloader import M4Raw_TestSet as M4Raw_TestSet_new, M4Raw_TrainSet as M4Raw_TrainSet_new |
|
|
| from dataloaders.m4raw_dataloader import M4Raw_TestSet, M4Raw_TrainSet |
|
|
|
|
|
|
|
|
| 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) |
|
|
| if use_new_dataloader: |
| db_test = M4Raw_TestSet_new(args, use_kspace=use_kspace) |
| else: |
| db_test = M4Raw_TestSet(args.root_path, args.MRIDOWN, 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') |
| |
| print('load weights from ' + save_mode_path) |
|
|
| try: |
| checkpoint = torch.load(save_mode_path) |
| except: |
| print("Missing keys:", set(model_state_dict.keys()) - set(loaded_state_dict.keys())) |
|
|
|
|
| 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/') |
|
|
|
|