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 # Results 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 # Ksample | ColdDiffusion | DDPM 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 = {} # defaultdict(dict) target_dic = {} # efaultdict(dict) input_dic = {} # defaultdict(dict) 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() # Degradation if use_kspace: b = 1 t = torch.randint(num_timesteps - 1, num_timesteps, (b,), device=device).long() # t-1 mask = kspace_masks[t] fft, mask = apply_tofre(t2_in.clone(), mask) # t2_img fft = fft * mask + 0.0 t2_in = apply_to_spatial(fft) t2_in_origin = t2_in.clone() while t >= 0: # outputs = model(t2_in, t1_img)['img_out'] 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] # last one t2_in = outputs else: if test_sample == "Ksample": # Ksample | ColdDiffusion | DDPM 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] # get_kspace_kernels(t - 2).cuda() 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'] # print("outputs shape: ", outputs.shape, outputs.min(), outputs.max()) # print("t2_img shape: ", t2_img.shape, t2_img.min(), t2_img.max()) 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() # outputs_save = np.clip(outputs_save, a_min=-1, a_max=1) 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) # 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) # print("outputs_save shape: ", outputs_save.shape, outputs_save.min(), outputs_save.max()) # print("target_save shape: ", target_save.shape, target_save.min(), target_save.max()) # print("in_save shape: ", in_save.shape, in_save.min(), in_save.max()) 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]) # 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)) nmse_meter_all.append(our_nmse) psnr_meter_all.append(our_psnr) ssim_meter_all.append(our_ssim) # print("psnr_meter_all: ", np.mean(psnr_meter_all)) for name in output_dic.keys(): print("name: ", name, len(output_dic[name])) # f_output = torch.stack([v for _, v in output_dic[name].items()]) # f_target = torch.stack([v for _, v in target_dic[name].items()]) 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) # 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) 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/')