| """ |
| 2023/10/16, |
| preprocess kspace data with the undersampling mask in the fastMRI project. |
| """ |
|
|
| import contextlib |
|
|
| import numpy as np |
| import torch |
|
|
|
|
| @contextlib.contextmanager |
| def temp_seed(rng, seed): |
| state = rng.get_state() |
| rng.seed(seed) |
| try: |
| yield |
| finally: |
| rng.set_state(state) |
|
|
|
|
| def create_mask_for_mask_type(mask_type_str, center_fractions, accelerations): |
| if mask_type_str == "random": |
| return RandomMaskFunc(center_fractions, accelerations) |
| elif mask_type_str == "equispaced": |
| return EquispacedMaskFunc(center_fractions, accelerations) |
| else: |
| raise Exception(f"{mask_type_str} not supported") |
|
|
|
|
| |
| def mri_fourier_transform_2d(image, mask): |
| ''' |
| image: input tensor [B, H, W, C] |
| mask: mask tensor [H, W] |
| ''' |
| spectrum = torch.fft.fftn(image, dim=(1, 2), norm='ortho') |
| |
| spectrum = torch.fft.fftshift(spectrum, dim=(1, 2)) |
| |
| masked_spectrum = spectrum * mask[None, :, :, None] |
| return spectrum, masked_spectrum |
|
|
|
|
| |
| def mri_inver_fourier_transform_2d(spectrum): |
| ''' |
| image: input tensor [B, H, W, C] |
| ''' |
| spectrum = torch.fft.ifftshift(spectrum, dim=(1, 2)) |
| image = torch.fft.ifftn(spectrum, dim=(1, 2), norm='ortho') |
| |
|
|
| return image |
|
|
|
|
| def add_gaussian_noise(kspace, snr): |
| |
| num_pixels = kspace.shape[0]*kspace.shape[1]*kspace.shape[2]*kspace.shape[3] |
| psr = torch.sum(torch.abs(kspace.real)**2)/num_pixels |
| pnr = psr/(np.power(10, snr/10)) |
| noise_r = torch.randn_like(kspace.real)*np.sqrt(pnr) |
|
|
| psim = torch.sum(torch.abs(kspace.imag)**2)/num_pixels |
| pnim = psim/(np.power(10, snr/10)) |
| noise_im = torch.randn_like(kspace.imag)*np.sqrt(pnim) |
|
|
| noise = noise_r + 1j*noise_im |
| noisy_kspace = kspace + noise |
|
|
| return noisy_kspace |
|
|
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
|
|
| |
| |
|
|
| |
| |
|
|
| def mri_fft_m4raw(lq_mri, hq_mri): |
| |
| lq_mri = torch.tensor(lq_mri[0])[None, :, :, None].to(torch.float32) |
| lq_mri_spectrum = torch.fft.fftn(lq_mri, dim=(1, 2), norm='ortho') |
| lq_mri_spectrum = torch.fft.fftshift(lq_mri_spectrum, dim=(1, 2)) |
|
|
| lq_mri = mri_inver_fourier_transform_2d(lq_mri_spectrum) |
| lq_mri = torch.cat([torch.real(lq_mri), torch.imag(lq_mri)], dim=-1) |
| lq_kspace = torch.cat([torch.real(lq_mri_spectrum), torch.imag(lq_mri_spectrum)], dim=-1) |
|
|
|
|
| hq_mri = torch.tensor(hq_mri[0])[None, :, :, None].to(torch.float32) |
| hq_mri_spectrum = torch.fft.fftn(hq_mri, dim=(1, 2), norm='ortho') |
| hq_mri_spectrum = torch.fft.fftshift(hq_mri_spectrum, dim=(1, 2)) |
|
|
| hq_mri = mri_inver_fourier_transform_2d(hq_mri_spectrum) |
| hq_mri = torch.cat([torch.real(hq_mri), torch.imag(hq_mri)], dim=-1) |
| hq_kspace = torch.cat([torch.real(hq_mri_spectrum), torch.imag(hq_mri_spectrum)], dim=-1) |
|
|
| |
| return lq_kspace[0], lq_mri[0].permute(2, 0, 1), \ |
| hq_kspace[0], hq_mri[0].permute(2, 0, 1) |
|
|
| def mri_fft(raw_mri, _SNR): |
| mri = torch.tensor(raw_mri)[None, :, :, None].to(torch.float32) |
| spectrum = torch.fft.fftn(mri, dim=(1, 2), norm='ortho') |
| |
| kspace = torch.fft.fftshift(spectrum, dim=(1, 2)) |
|
|
| if _SNR > 0: |
| noisy_kspace = add_gaussian_noise(kspace, _SNR) |
| else: |
| noisy_kspace = kspace |
|
|
| noisy_mri = mri_inver_fourier_transform_2d(noisy_kspace) |
| noisy_mri = torch.cat([torch.real(noisy_mri), torch.imag(noisy_mri)], dim=-1) |
| noisy_kspace = torch.cat([torch.real(noisy_kspace), torch.imag(noisy_kspace)], dim=-1) |
|
|
| raw_ksapce = torch.cat([torch.real(kspace), torch.imag(kspace)], dim=-1) |
| raw_mri = mri_inver_fourier_transform_2d(kspace) |
| raw_mri = torch.cat([torch.real(raw_mri), torch.imag(raw_mri)], dim=-1) |
|
|
| |
| return noisy_kspace[0], noisy_mri[0].permute(2, 0, 1), \ |
| raw_ksapce[0], raw_mri[0].permute(2, 0, 1) |
|
|
| def undersample_mri(raw_mri, _MRIDOWN, _SNR): |
| mri = torch.tensor(raw_mri)[None, :, :, None].to(torch.float32) |
| raw_spectrum = torch.fft.fftn(mri, dim=(1, 2), norm='ortho') |
| |
| raw_kspace = torch.fft.fftshift(raw_spectrum, dim=(1, 2)) |
|
|
| if not _MRIDOWN == "0X": |
| if _MRIDOWN == "4X": |
| mask_type_str, center_fraction, MRIDOWN = "random", 0.1, 4 |
| elif _MRIDOWN == "8X": |
| mask_type_str, center_fraction, MRIDOWN = "equispaced", 0.04, 8 |
|
|
| ff = create_mask_for_mask_type(mask_type_str, [center_fraction], [MRIDOWN]) |
| shape = [240, 240, 1] |
| mask = ff(shape, seed=1337) |
| mask = mask[:, :, 0] |
| |
| |
| kspace, masked_kspace = mri_fourier_transform_2d(mri, mask) |
|
|
| |
| if _SNR > 0: |
| noisy_kspace = add_gaussian_noise(masked_kspace, _SNR) |
| else: |
| noisy_kspace = masked_kspace |
|
|
| else: |
| if _SNR > 0: |
| noisy_kspace = add_gaussian_noise(raw_kspace, _SNR) |
| else: |
| noisy_kspace = raw_kspace |
| |
| mask = torch.ones([1,240]) |
| |
| |
| noisy_mri = mri_inver_fourier_transform_2d(noisy_kspace) |
| noisy_mri = torch.cat([torch.real(noisy_mri), torch.imag(noisy_mri)], dim=-1) |
|
|
| noisy_kspace_ = torch.cat([torch.real(noisy_kspace), torch.imag(noisy_kspace)], dim=-1) |
|
|
| raw_mri = mri_inver_fourier_transform_2d(raw_kspace) |
| raw_mri = torch.cat([torch.real(raw_mri), torch.imag(raw_mri)], dim=-1) |
| raw_kspace = torch.cat([torch.real(raw_kspace), torch.imag(raw_kspace)], dim=-1) |
|
|
| return noisy_kspace_[0], noisy_mri[0].permute(2, 0, 1), \ |
| raw_kspace[0], raw_mri[0].permute(2, 0, 1), mask.unsqueeze(-1) |
|
|
| |
| |
| |
| |
| |
| |
|
|
| |
|
|
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
|
|
| |
| |
|
|
|
|
|
|
| class MaskFunc(object): |
| """ |
| An object for GRAPPA-style sampling masks. |
| |
| This crates a sampling mask that densely samples the center while |
| subsampling outer k-space regions based on the undersampling factor. |
| """ |
|
|
| def __init__(self, center_fractions, accelerations): |
| """ |
| Args: |
| center_fractions (List[float]): Fraction of low-frequency columns to be |
| retained. If multiple values are provided, then one of these |
| numbers is chosen uniformly each time. |
| accelerations (List[int]): Amount of under-sampling. This should have |
| the same length as center_fractions. If multiple values are |
| provided, then one of these is chosen uniformly each time. |
| """ |
| if len(center_fractions) != len(accelerations): |
| raise ValueError( |
| "Number of center fractions should match number of accelerations" |
| ) |
|
|
| self.center_fractions = center_fractions |
| self.accelerations = accelerations |
| self.rng = np.random |
|
|
| def choose_acceleration(self): |
| """Choose acceleration based on class parameters.""" |
| choice = self.rng.randint(0, len(self.accelerations)) |
| center_fraction = self.center_fractions[choice] |
| acceleration = self.accelerations[choice] |
|
|
| return center_fraction, acceleration |
|
|
|
|
| class RandomMaskFunc(MaskFunc): |
| """ |
| RandomMaskFunc creates a sub-sampling mask of a given shape. |
| |
| The mask selects a subset of columns from the input k-space data. If the |
| k-space data has N columns, the mask picks out: |
| 1. N_low_freqs = (N * center_fraction) columns in the center |
| corresponding to low-frequencies. |
| 2. The other columns are selected uniformly at random with a |
| probability equal to: prob = (N / acceleration - N_low_freqs) / |
| (N - N_low_freqs). This ensures that the expected number of columns |
| selected is equal to (N / acceleration). |
| |
| It is possible to use multiple center_fractions and accelerations, in which |
| case one possible (center_fraction, acceleration) is chosen uniformly at |
| random each time the RandomMaskFunc object is called. |
| |
| For example, if accelerations = [4, 8] and center_fractions = [0.08, 0.04], |
| then there is a 50% probability that 4-fold acceleration with 8% center |
| fraction is selected and a 50% probability that 8-fold acceleration with 4% |
| center fraction is selected. |
| """ |
|
|
| def __call__(self, shape, seed=None): |
| """ |
| Create the mask. |
| |
| Args: |
| shape (iterable[int]): The shape of the mask to be created. The |
| shape should have at least 3 dimensions. Samples are drawn |
| along the second last dimension. |
| seed (int, optional): Seed for the random number generator. Setting |
| the seed ensures the same mask is generated each time for the |
| same shape. The random state is reset afterwards. |
| |
| Returns: |
| torch.Tensor: A mask of the specified shape. |
| """ |
| if len(shape) < 3: |
| raise ValueError("Shape should have 3 or more dimensions") |
|
|
| with temp_seed(self.rng, seed): |
| num_cols = shape[-2] |
| center_fraction, acceleration = self.choose_acceleration() |
|
|
| |
| num_low_freqs = int(round(num_cols * center_fraction)) |
| prob = (num_cols / acceleration - num_low_freqs) / ( |
| num_cols - num_low_freqs |
| ) |
| mask = self.rng.uniform(size=num_cols) < prob |
| pad = (num_cols - num_low_freqs + 1) // 2 |
| mask[pad : pad + num_low_freqs] = True |
|
|
| |
| mask_shape = [1 for _ in shape] |
| mask_shape[-2] = num_cols |
| mask = torch.from_numpy(mask.reshape(*mask_shape).astype(np.float32)) |
|
|
| return mask |
|
|
|
|
| class EquispacedMaskFunc(MaskFunc): |
| """ |
| EquispacedMaskFunc creates a sub-sampling mask of a given shape. |
| |
| The mask selects a subset of columns from the input k-space data. If the |
| k-space data has N columns, the mask picks out: |
| 1. N_low_freqs = (N * center_fraction) columns in the center |
| corresponding tovlow-frequencies. |
| 2. The other columns are selected with equal spacing at a proportion |
| that reaches the desired acceleration rate taking into consideration |
| the number of low frequencies. This ensures that the expected number |
| of columns selected is equal to (N / acceleration) |
| |
| It is possible to use multiple center_fractions and accelerations, in which |
| case one possible (center_fraction, acceleration) is chosen uniformly at |
| random each time the EquispacedMaskFunc object is called. |
| |
| Note that this function may not give equispaced samples (documented in |
| https://github.com/facebookresearch/fastMRI/issues/54), which will require |
| modifications to standard GRAPPA approaches. Nonetheless, this aspect of |
| the function has been preserved to match the public multicoil data. |
| """ |
|
|
| def __call__(self, shape, seed): |
| """ |
| Args: |
| shape (iterable[int]): The shape of the mask to be created. The |
| shape should have at least 3 dimensions. Samples are drawn |
| along the second last dimension. |
| seed (int, optional): Seed for the random number generator. Setting |
| the seed ensures the same mask is generated each time for the |
| same shape. The random state is reset afterwards. |
| |
| Returns: |
| torch.Tensor: A mask of the specified shape. |
| """ |
| if len(shape) < 3: |
| raise ValueError("Shape should have 3 or more dimensions") |
|
|
| with temp_seed(self.rng, seed): |
| center_fraction, acceleration = self.choose_acceleration() |
| num_cols = shape[-2] |
| num_low_freqs = int(round(num_cols * center_fraction)) |
|
|
| |
| mask = np.zeros(num_cols, dtype=np.float32) |
| pad = (num_cols - num_low_freqs + 1) // 2 |
| mask[pad : pad + num_low_freqs] = True |
|
|
| |
| adjusted_accel = (acceleration * (num_low_freqs - num_cols)) / ( |
| num_low_freqs * acceleration - num_cols |
| ) |
| offset = self.rng.randint(0, round(adjusted_accel)) |
|
|
| accel_samples = np.arange(offset, num_cols - 1, adjusted_accel) |
| accel_samples = np.around(accel_samples).astype(np.uint) |
| mask[accel_samples] = True |
|
|
| |
| mask_shape = [1 for _ in shape] |
| mask_shape[-2] = num_cols |
| mask = torch.from_numpy(mask.reshape(*mask_shape).astype(np.float32)) |
|
|
| return mask |
|
|