qic999's picture
Upload folder using huggingface_hub
28e6f98 verified
Raw
History Blame Contribute Delete
9.25 kB
# ---------------------------
# 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')