qic999's picture
Upload folder using huggingface_hub
28e6f98 verified
Raw
History Blame Contribute Delete
2.38 kB
"""
读取图像的ground truth和每个round重建结果, 并且绘制error map.
"""
import os
import numpy as np
from PIL import Image
from skimage import io
from matplotlib import pyplot as plt
def normalize_image(image):
# return (image - image.min())/(image.max() - image.min())
image = image[40:200, 55:215]
# image = image[80:160, 95:175]
print("image shape:", image.shape)
return image/255.0
# return (image - image.min())/(image.max() - image.min())
def viz_diff_img(image, test_outputdir, image_name):
print("image range:", image.max(), image.min())
plt.axis('off')
# plt.imshow(image, cmap='jet',vmin=0, vmax=50)
plt.imshow(image, cmap='jet',vmin=0, vmax=30)
# plt.colorbar()
plt.savefig(os.path.join(test_outputdir, f'{image_name}'), bbox_inches='tight',pad_inches = 0)
# baseline = 'UNet_4X'
# baseline_list = ['DCAMSR_4X', 'MCCA_4X', 'MINet_4X', 'MTrans_4X', 'swinir_4X_']
baseline_list = ['DCAMSR_8X', 'MCCA_8X', 'MINet_8X', 'MTrans_8X', 'swinir_8X_']
# baseline_list = ['swinir_8X_']
baseline_list = ['our']
for baseline in baseline_list:
# root_dir = f"/data/qic99/recon_code/recon_2M/BRATS_baseline/model/{baseline}/result_case/"
root_dir = '/data/qic99/recon_code/recon_2M/BRATS_freq_multi_fusion_2_2_4/model/unet_wo_kspace_4X_lr1e-4/result_case/'
# root_dir = '/data/qic99/recon_code/recon_2M/BRATS_freq_multi_fusion_2_2_4/model/unet_wo_kspace_8X_lr1e-4/result_case/'
image_name = "301_t2"
dst_dir = "./error_map_8X"
# dst_dir = "./error_map_4X"
os.makedirs(dst_dir, exist_ok=True)
img_gt = normalize_image(np.array(Image.open(root_dir + image_name + ".png")))
img_in = normalize_image(np.array(Image.open(root_dir + image_name + "_out.png")))
img_lq = normalize_image(np.array(Image.open(root_dir + image_name + "_in.png")))
print(img_gt.max(), img_gt.min())
print(img_in.max(), img_in.min())
io.imsave(os.path.join(dst_dir, image_name + "_lq.png"), (img_lq*255).astype(np.uint8))
io.imsave(os.path.join(dst_dir, image_name + ".png"), (img_gt*255).astype(np.uint8))
io.imsave(os.path.join(dst_dir, baseline+'_'+image_name + "_out.png"), (img_in*255).astype(np.uint8))
viz_diff_img(np.abs(img_gt - img_in)*255, dst_dir, baseline+'_'+image_name + "_error_map.png")
print("input error:", np.mean(np.abs(img_gt - img_in)))
# break