# --------------------------- # 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')