msbackup / MRI_recon /code /Frequency-Diffusion /FSMNet /frequency_diffusion /degradation /kspace_test.py
| # --------------------------- | |
| # Fade kernels | |
| # --------------------------- | |
| import cv2, torch | |
| import numpy as np | |
| import torchgeometry as tgm | |
| from torch.fft import fft2, ifft2, fftshift, ifftshift | |
| import matplotlib.pyplot as plt | |
| from mask_utils import RandomMaskFunc, EquispacedMaskFunc | |
| try: | |
| from mask_utils import RandomMaskFunc, EquispacedMaskFunc | |
| except: | |
| from mask_utils import RandomMaskFunc, EquispacedMaskFractionFunc, EquispacedMaskFunc, RandomPatchFunc | |
| try: | |
| from .k_degradation import get_fade_kernel, get_fade_kernels, get_mask_func, high_fre_mask, get_ksu_kernel | |
| except: | |
| from k_degradation import get_fade_kernel, get_fade_kernels, get_mask_func, high_fre_mask, get_ksu_kernel | |
| use_fix_center_ratio = False | |
| def get_ksu_mask(mask_method, af, cf, pe, fe, seed=0, is_training=False): | |
| mask_func = get_mask_func(mask_method, af, cf) # acceleration factor, center fraction | |
| if mask_method in ['cartesian_regular', 'cartesian_random']: | |
| mask, num_low_freq = mask_func((1, pe, 1), seed=seed) # mask (torch): (1, pe, 1) | |
| # Extent the stride | |
| mask = mask.permute(0, 2, 1).repeat(1, fe, 1) # mask (torch): (1, pe, 1) --> (1, 1, pe) --> (1, fe, pe) | |
| if is_training: # Step 1, Random Drop elements to learn intra-stride effects, patch-wise | |
| af_new = 1.0 + (af - 1.0) / 2 | |
| # af_new = max(af_new, 1.0) | |
| patch_mask = get_mask_func("randompatch", af_new, cf) | |
| mask_, _ = patch_mask((fe, pe, 1), seed=seed) # mask (numpy): (fe, pe) | |
| mask = mask_ * mask | |
| elif mask_method == 'gaussian_2d': | |
| mask, _ = mask_func(resolution=(fe, pe), accel=af, sigma=100, seed=seed) # mask (numpy): (fe, pe) | |
| mask = torch.from_numpy(mask[np.newaxis, :, :]) # mask (torch): (fe, pe) --> (1, fe, pe) | |
| elif mask_method in ['radial_add', 'radial_sub', 'spiral_add', 'spiral_sub']: | |
| sr = 1 / af | |
| mask = mask_func(mask_type=mask_method, mask_sr=sr, res=pe, seed=seed) # mask (numpy): (pe, pe) | |
| mask = torch.from_numpy(mask[np.newaxis, :, :]) # mask (torch): (pe, pe) --> (1, pe, pe) | |
| else: | |
| raise NotImplementedError | |
| # print("return mask = ", mask.shape) | |
| return mask | |
| # ksu_masks = get_ksu_kernels() | |
| # (C, H, W) --> (B, C, H, W) | |
| high_fre_mask_cls = high_fre_mask() | |
| def apply_ksu_kernel(x_start, mask, params_dict=None, pixel_range='mean_std', | |
| use_fre_noise=False, return_mask=False): | |
| fft, mask = apply_tofre(x_start, mask, params_dict, pixel_range) | |
| # Use the high frequency mask to add noise | |
| if use_fre_noise: | |
| fft = fft * mask | |
| fft_magnitude = torch.abs(fft) # 幅度 | |
| fft_phase = torch.angle(fft) # 相位 | |
| _, _, H, W = fft.shape | |
| high_freq_mask = high_fre_mask_cls(H, W).to(fft.device) | |
| high_freq_mask = high_freq_mask.unsqueeze(0).unsqueeze(0).repeat(fft.shape[0], 1, 1, 1) | |
| # Background Noise | |
| sigma = 0.2 | |
| noise = torch.randn_like(fft_magnitude) * sigma | |
| mean_mag = fft_magnitude.sum() / (mask.sum() + 1) | |
| noise_magnitude_high = noise * (mean_mag) * (1 - mask) # high_freq_mask | |
| sigma = 0.1 | |
| noise = torch.randn_like(fft_magnitude) * sigma | |
| noise_magnitude_low = noise * fft_magnitude * mask # (1 - high_freq_mask) | |
| # fft_noisy_magnitude = fft_magnitude * mask + noise_magnitude * high_freq_mask * (1 - mask) | |
| fft_noisy_magnitude = fft_magnitude * mask | |
| fft_noisy_magnitude += noise_magnitude_high + noise_magnitude_low | |
| fft_noisy_magnitude = torch.clamp(fft_noisy_magnitude, min=0.0) | |
| fft = fft_noisy_magnitude * torch.exp(1j * fft_phase) | |
| else: | |
| fft = fft * mask | |
| x_ksu = apply_to_spatial(fft, params_dict, pixel_range) | |
| if return_mask: | |
| return x_ksu, fft, fft_magnitude | |
| return x_ksu | |
| def apply_tofre(x_start, mask, params_dict=None, pixel_range='mean_std'): | |
| fft = fftshift(fft2(x_start)) | |
| mask = mask.to(fft.device) | |
| return fft, mask # , _min, _max | |
| def apply_to_spatial(fft, params_dict=None, pixel_range='mean_std'): | |
| x_ksu = ifft2(ifftshift(fft)) | |
| x_ksu = torch.abs(x_ksu) | |
| return x_ksu | |
| if __name__ == "__main__": | |
| # First STEP | |
| import SimpleITK as sitk | |
| import numpy as np | |
| import os | |
| image_size = 240 | |
| batch_size = 1 | |
| t = 5 | |
| use_linux = True | |
| # Load MRI back here | |
| if use_linux: | |
| root = "/gamedrive/Datasets/medical/Brain/brats/ASNR-MICCAI-BraTS2023-GLI-Challenge-TrainingData" | |
| p_id = 639 | |
| modality = "T1C" | |
| filename = f"{root}/BraTS-GLI-{p_id:05d}-000/BraTS-GLI-{p_id:05d}-000-{modality.lower()}.nii.gz" | |
| img_obj = sitk.ReadImage(filename) | |
| img_array = sitk.GetArrayFromImage(img_obj) | |
| slice = img_array.shape[0] // 2 | |
| img = img_array[slice, ...] | |
| plt.imshow(img, cmap="gray") | |
| plt.show() | |
| img = (img - img.min()) / (img.max() - img.min()) | |
| plt.imsave("visualization/original.png", img, cmap="gray") | |
| else: | |
| # Or use PNG | |
| img = plt.imread( | |
| "/Users/haochen/Documents/GitHub/DiffDecomp/Cold-Diffusion/generation-diffusion-pytorch/diffusion_pytorch/assets/img.png") | |
| img = np.transpose(img, (2, 0, 1))[0] | |
| img = cv2.resize(img, (image_size, image_size), interpolation=cv2.INTER_LINEAR) | |
| print("img shape=", img.shape) | |
| img = np.expand_dims(img, axis=0) | |
| img = torch.from_numpy(img).unsqueeze(0).float() | |
| ksu_routine = "LogSamplingRate" # "LinearSamplingRate" # | |
| kspace_kernels, patch_drop_masks = get_ksu_kernel(t, image_size, | |
| ksu_routine=ksu_routine, is_training=True, | |
| example_frequency_img=example) | |
| kspace_kernels = torch.stack(kspace_kernels).squeeze(1) | |
| # all k_space | |
| rand_kernels = [] | |
| rand_x = torch.randint(0, image_size + 1, (batch_size,)).long() | |
| for i in range(batch_size): | |
| # k = kspace_kernels[j].clone() | |
| rand_kernels.append(torch.stack( | |
| [kspace_kernels[j][: image_size, # rand_x[i]:rand_x[i] + | |
| : image_size] for j in | |
| range(len(kspace_kernels))])) | |
| rand_kernels = torch.stack(rand_kernels) | |
| # rand_kernels shape: torch.Size([24, 5, 128, 128]) | |
| ori_img = img | |
| masked_img = [] | |
| masks = [] | |
| for i in range(t): | |
| k = torch.stack([rand_kernels[:, i]], 1)[0] | |
| masks.append(k) | |
| img = apply_ksu_kernel(img, k, pixel_range='0_1') | |
| masked_img.append(img) | |
| # Visualize the masks | |
| masks = np.concatenate(masks, axis=-1)[0] | |
| masked_img = (torch.concat(masked_img, dim=-1).numpy() + 1) * 0.5 | |
| masked_img = np.transpose(masked_img, (0, 2, 3, 1))[0, ..., 0] | |
| # Save individually | |
| print("masks / masked_img=", masks.max(), masked_img.max()) | |
| # img = np.concatenate([masks, masked_img], axis=0) | |
| plt.imsave("visualization/sample_masks.png", masks, cmap='gray') | |
| # masked_img = (masked_img - masked_img.min())/(masked_img) | |
| # masked_img = np.concatenate([masked_img, 1-masked_img], axis=0) | |
| plt.imsave("visualization/sample_images.png", masked_img, cmap='gray') | |
| w = masked_img.shape[0] | |
| pr_folder = "visualization/progressive" | |
| os.makedirs(pr_folder, exist_ok=True) | |
| # Progressive | |
| print() | |
| for i in range(t): | |
| plt.imsave(f"{pr_folder}/{i}_masked_img.png", masked_img[:, i * w: (i + 1) * w], cmap='gray') | |
| plt.imsave(f"{pr_folder}/{i}_masks.png", masks[:, i * w: (i + 1) * w], cmap='gray') | |
| img = ori_img | |
| # use_fre_noise=False, return_mask=False | |
| masked_img = [] | |
| masks = [] | |
| fft = [] | |
| ks = [] | |
| for i in range(t): | |
| k = torch.stack([rand_kernels[:, i]], 1)[0] | |
| ks.append(k) | |
| img, k, fft_original = apply_ksu_kernel(img, k, pixel_range='0_1', use_fre_noise=True, return_mask=True) | |
| # k -> fft | |
| fft_magnitude = np.abs(k) # 幅度 | |
| # fft_phase = torch.angle(k) # 相位 | |
| mag = np.log(fft_magnitude[0]) | |
| masks.append(mag) | |
| fft.append(np.log(fft_original[0])) | |
| masked_img.append(img) | |
| # Visualize the masks | |
| masks = np.concatenate(masks, axis=-1)[0] | |
| ks = np.concatenate(ks, axis=-1)[0] | |
| masked_img = (torch.concat(masked_img, dim=-1).numpy() + 1) * 0.5 | |
| masked_img = np.transpose(masked_img, (0, 2, 3, 1))[0, ..., 0] | |
| fft = np.concatenate(fft, axis=-1)[0] | |
| plt.imsave("visualization/sample_noisy_mask.png", masks, cmap='gray') | |
| # masked_img = np.concatenate([masked_img, 1 - masked_img], axis=0) | |
| plt.imsave("visualization/sample_noisy_image.png", masked_img, cmap='gray') | |
| # print("masked_img shape=", masked_img.shape, w) | |
| # Progressive | |
| for i in range(t): | |
| # print("masked_img[:, t*w: (t+1)*w] = ", masked_img[:, t*w: (t+1)*w].shape, t*w) | |
| plt.imsave(f"{pr_folder}/{i}_n_masked_img.png", masked_img[:, i * w: (i + 1) * w], cmap='gray') | |
| plt.imsave(f"{pr_folder}/{i}_n_masks.png", masks[:, i * w: (i + 1) * w], cmap='gray') | |
| plt.imsave(f"{pr_folder}/{i}_n_fft.png", fft[:, i * w: (i + 1) * w], cmap='gray') | |
| plt.imsave(f"{pr_folder}/{i}_n_ks.png", ks[:, i * w: (i + 1) * w], cmap='gray') | |