| import copy, time |
| import gc |
|
|
| import torch |
| from torch import nn |
| import torch.nn.functional as func |
| from inspect import isfunction |
| from functools import partial |
|
|
| from torch.utils import data |
| from pathlib import Path |
| from torch.optim import AdamW, lr_scheduler |
| from torchvision import utils |
| import torch.nn.functional as F |
| |
|
|
| import os |
| import errno |
| from PIL import Image |
| |
| import cv2 |
| import numpy as np |
| import imageio |
| from skimage.metrics import mean_squared_error, peak_signal_noise_ratio, structural_similarity |
|
|
| |
| from diffusion_pytorch.degradation.k_degradation import get_fade_kernels, get_ksu_kernel, apply_ksu_kernel, apply_tofre, apply_to_spatial |
| from dataset import Dataset, Dataset_Aug1, BrainDataset |
|
|
| |
| from skimage.metrics import peak_signal_noise_ratio as psnr |
|
|
| from torchmetrics.image import StructuralSimilarityIndexMeasure |
|
|
| from metrics.lpips import LPIPS |
| from metrics.fid import calculate_fid |
| from metrics.fid_3d import calculate_fid_3d |
| from metrics.nmse import nmse |
| from diffusion_pytorch.degradation.k_degradation import get_noisy_patches |
| import torch.amp as amp |
| from torch.cuda.amp import GradScaler, autocast |
| from metrics.frequency_loss import AMPLoss |
| scaler = GradScaler() |
|
|
| |
| |
| |
| |
| |
| |
|
|
|
|
| def exists(x): |
| return x is not None |
|
|
|
|
| def default(val, d): |
| if exists(val): |
| return val |
| return d() if isfunction(d) else d |
|
|
|
|
| def create_folder(path): |
| try: |
| os.mkdir(path) |
| except OSError as exc: |
| if exc.errno != errno.EEXIST: |
| raise |
| pass |
|
|
|
|
| def cycle(dl): |
| while True: |
| for inputs in dl: |
| yield inputs |
|
|
|
|
| def num_to_groups(num, divisor): |
| groups = num // divisor |
| remainder = num % divisor |
| arr = [divisor] * groups |
| if remainder > 0: |
| arr.append(remainder) |
| return arr |
|
|
|
|
| def loss_backwards(fp16, loss, optimizer, **kwargs): |
| if fp16: |
| scaler.scale(loss).backward(**kwargs) |
| else: |
| loss.backward(**kwargs) |
|
|
|
|
|
|
| class EMA: |
| def __init__(self, beta): |
| super().__init__() |
| self.beta = beta |
|
|
| def update_model_average(self, ma_model, current_model): |
| for current_params, ma_params in zip(current_model.parameters(), ma_model.parameters()): |
| old_weight, up_weight = ma_params.data, current_params.data |
| ma_params.data = self.update_average(old_weight, up_weight) |
|
|
| def update_average(self, old, new): |
| if old is None: |
| return new |
| return old * self.beta + (1 - self.beta) * new |
|
|
|
|
|
|
| class GaussianDiffusion(nn.Module): |
| def __init__( |
| self, |
| diffusion_type, |
| restore_fn, |
| *, |
| image_size, |
| device_of_kernel, |
| channels=3, |
| timesteps=1000, |
| loss_type='l1', |
| kernel_std=0.1, |
| initial_mask=11, |
| fade_routine='Incremental', |
| sampling_routine='default', |
| discrete=False, |
| accelerate_factor=4, |
| fp16=False, |
| normalizer="mean_std", |
| example_frequency_img=None |
| |
| ): |
| super().__init__() |
| self.fp16 = fp16 |
| self.channels = channels |
| self.image_size = image_size |
| self.restore_fn = restore_fn |
| self.accelerate_factor = accelerate_factor |
|
|
| |
| self.example_frequency_img = example_frequency_img |
| self.device_of_kernel = device_of_kernel |
| self.num_timesteps = int(timesteps) |
| self.loss_type = loss_type |
| self.kernel_std = kernel_std |
| self.initial_mask = initial_mask |
| self.fade_routine = fade_routine |
| self.backbone = diffusion_type.split('_')[0] |
| self.degradation_type = diffusion_type.split('_')[1] |
|
|
| self.sampling_routine = sampling_routine |
| self.discrete = discrete |
|
|
| |
| if self.backbone == 'twobranch' or self.backbone == 'twounet': |
| self.amploss = AMPLoss() |
| self.kl_loss = torch.nn.KLDivLoss(reduction='sum') |
|
|
| self.lpips = LPIPS().eval().cuda() |
| self.ssim = StructuralSimilarityIndexMeasure() |
|
|
| self.use_fre_loss = True |
|
|
| self.use_ssim = False |
| self.use_lpips = True |
|
|
| self.clamp_every_sample = False |
| if normalizer == "min_max": |
| self.clamp_every_sample = True |
|
|
| self.use_fre_noise = True |
|
|
| self.update_kernel = False |
| self.use_patch_kernel = False |
| self.use_kl = False |
|
|
|
|
| if self.degradation_type == 'fade': |
| self.fade_kernels = get_fade_kernels(fade_routine, self.num_timesteps, image_size, kernel_std, initial_mask) |
| |
|
|
| elif self.degradation_type == "kspace": |
| self.get_new_kspace(is_training=True) |
| |
| else: |
| raise NotImplementedError() |
|
|
| def get_new_kspace(self, is_training=False): |
| |
| self.kspace_kernels, self.noisy_kspace_kernels = get_ksu_kernel(self.num_timesteps, self.image_size, |
| ksu_routine="LogSamplingRate", is_training=is_training, |
| accelerated_factor=self.accelerate_factor, |
| example_frequency_img=self.example_frequency_img) |
|
|
| self.kspace_kernels = torch.stack(self.kspace_kernels).squeeze(1).cuda() |
| self.noisy_kspace_kernels =self.kspace_kernels.cuda() |
|
|
|
|
| def get_kspace_kernels(self, index): |
| k = torch.stack([self.kspace_kernels[index]], 0).unsqueeze(0) |
| return k |
|
|
| @torch.no_grad() |
| def sample(self, batch_size=16, faded_recon_sample=None, aux=None, |
| t=None, params_dict=None, sample_routine=None): |
| |
| self.restore_fn.eval() |
|
|
| if not sample_routine: |
| sample_routine = self.sampling_routine |
|
|
| sample_device = faded_recon_sample.device |
| batch_size = faded_recon_sample.size(0) |
|
|
|
|
|
|
| if t is None: |
| t = self.num_timesteps |
|
|
| |
| |
|
|
| |
| with torch.no_grad(): |
| k = torch.stack([self.kspace_kernels[[t - 1]]], 1) |
| |
| faded_recon_sample = apply_ksu_kernel(faded_recon_sample, k, params_dict) |
|
|
| return_k = k.repeat( batch_size, 1, 1, 1) |
|
|
| xt = faded_recon_sample |
| |
|
|
|
|
| direct_recons = None |
| recon_sample = None |
| all_recons = [] |
| all_recons_fre = [] |
| all_masks = [] |
|
|
| k_known_mask = torch.zeros_like(self.get_kspace_kernels(-1)).cuda() |
|
|
| while t: |
| step = torch.full((batch_size,), t - 1, dtype=torch.long).cuda() |
| if self.backbone == "unet": |
| recon_sample = self.restore_fn(faded_recon_sample, step) |
|
|
| elif self.backbone == "twounet": |
| recon_sample = self.restore_fn(faded_recon_sample, aux, k, step) |
|
|
| elif self.backbone == "twobranch": |
| recon_sample, recon_fre = self.restore_fn(faded_recon_sample, aux, step) |
| all_recons_fre.append(recon_fre) |
|
|
| if direct_recons is None: |
| direct_recons = recon_sample |
| all_recons.append(recon_sample) |
|
|
| if self.degradation_type == 'kspace': |
| |
|
|
| if sample_routine == 'default': |
| all_recons.append(recon_sample) |
| with torch.no_grad(): |
| if t >=1: |
| k = self.get_kspace_kernels(t - 1) |
| recon_sample = apply_ksu_kernel(recon_sample, k, params_dict) |
|
|
| all_masks.append(k) |
| faded_recon_sample = recon_sample |
|
|
|
|
| elif sample_routine == 'x0_step_down': |
| all_recons.append(recon_sample) |
| if t <= 1: |
| if t == 1: |
| |
| |
| |
| |
| faded_recon_sample = recon_sample |
|
|
| else: |
| faded_recon_sample = recon_sample |
| all_masks.append(k) |
| else: |
| with torch.no_grad(): |
| k = self.get_kspace_kernels(t - 2, self.kspace_kernels) |
| recon_sample_sub_1 = apply_ksu_kernel(recon_sample, k, params_dict) |
|
|
| k = self.get_kspace_kernels(t - 1, self.kspace_kernels) |
| recon_sample = apply_ksu_kernel(recon_sample, k, params_dict) |
|
|
| faded_recon_sample = faded_recon_sample - recon_sample + recon_sample_sub_1 |
|
|
| if self.clamp_every_sample: |
| faded_recon_sample = faded_recon_sample.clamp(-1, 1) |
| all_masks.append(k) |
|
|
| elif sample_routine == 'x0_step_down_fre': |
| all_recons.append(recon_sample) |
| if t <= 1: |
| kt = self.get_kspace_kernels(0).cuda() |
| faded_recon_sample = recon_sample |
| k_residual = torch.ones_like(kt).cuda() |
|
|
|
|
| else: |
| k_full = self.get_kspace_kernels(- 1) |
| faded_recon_sample_fre, k_full = apply_tofre(faded_recon_sample, k_full, params_dict) |
| |
|
|
| with torch.no_grad(): |
|
|
| kt_sub_1 = self.get_kspace_kernels(t - 1).cuda() |
| kt = self.get_kspace_kernels(t - 0).cuda() |
| k_residual = kt_sub_1 - kt |
| recon_sample_fre, k_residual = apply_tofre(recon_sample, k_residual, params_dict) |
|
|
| fre_amend = recon_sample_fre * k_residual |
| faded_recon_sample_fre = faded_recon_sample_fre + fre_amend |
|
|
| faded_recon_sample_fre = faded_recon_sample_fre * kt_sub_1 |
|
|
| faded_recon_sample = apply_to_spatial(faded_recon_sample_fre, params_dict) |
|
|
| k_known_mask += k_residual |
| all_masks.append(k_known_mask.cpu().clone()) |
|
|
| if self.clamp_every_sample: |
| faded_recon_sample = faded_recon_sample.clamp(-1, 1) |
|
|
|
|
| elif sample_routine == 'fre_progressive': |
| all_recons.append(recon_sample) |
| if t == 1: |
| recon_sample_sub_1 = recon_sample |
| k = self.get_kspace_kernels(0, self.kspace_kernels) |
|
|
| recon_sample = apply_ksu_kernel(recon_sample, k, params_dict) |
| faded_recon_sample = faded_recon_sample - recon_sample + recon_sample_sub_1 |
|
|
| elif t == 0: |
| faded_recon_sample = recon_sample |
| all_masks.append(k) |
| else: |
| k_full = self.get_kspace_kernels(- 1, self.kspace_kernels) |
| faded_recon_sample_fre, k_full = apply_tofre(faded_recon_sample, k_full, params_dict) |
|
|
| with torch.no_grad(): |
|
|
| kt_sub_1 = self.get_kspace_kernels(t - 2, self.kspace_kernels).cuda() |
| kt = self.get_kspace_kernels(t - 1, self.kspace_kernels).cuda() |
| k_residual = kt_sub_1 - k_full |
| recon_sample_fre, k_residual = apply_tofre(recon_sample, k_residual, params_dict) |
|
|
|
|
| fre_amend = recon_sample_fre * k_residual |
| faded_recon_sample_fre = faded_recon_sample_fre * k_full + fre_amend |
|
|
| |
|
|
| faded_recon_sample = apply_to_spatial(faded_recon_sample_fre, params_dict) |
|
|
| k_known_mask += k_residual |
| all_masks.append(k_known_mask.cpu().clone()) |
|
|
| if self.clamp_every_sample: |
| faded_recon_sample = faded_recon_sample.clamp(-1, 1) |
|
|
|
|
|
|
| recon_sample = faded_recon_sample |
| |
|
|
| t -= 1 |
|
|
| all_recons = torch.stack(all_recons) |
| all_masks = torch.stack(all_masks) |
| all_recons_fre = torch.stack(all_recons_fre) |
|
|
| return xt, direct_recons, recon_sample, return_k, all_recons, all_recons_fre, all_masks |
|
|
| @torch.no_grad() |
| def all_sample(self, batch_size=16, faded_recon_sample=None, aux=None, t=None, params_dict=None, times=None): |
| |
| print("Running into all_sample...") |
| rand_kernels = None |
| sample_device = faded_recon_sample.device |
| if self.degradation_type == 'fade': |
| if 'Random' in self.fade_routine: |
| rand_kernels = [] |
| rand_x = torch.randint(0, self.image_size + 1, (batch_size,), device=faded_recon_sample.device).long() |
| rand_y = torch.randint(0, self.image_size + 1, (batch_size,), device=faded_recon_sample.device).long() |
| for i in range(batch_size, ): |
| rand_kernels.append(torch.stack( |
| [self.fade_kernels[j][rand_x[i]:rand_x[i] + self.image_size, |
| rand_y[i]:rand_y[i] + self.image_size] for j in range(len(self.fade_kernels))])) |
| rand_kernels = torch.stack(rand_kernels) |
|
|
| elif self.degradation_type == 'kspace': |
| rand_kernels = [] |
| rand_x = torch.randint(0, self.image_size + 1, (batch_size,), device=faded_recon_sample.device).long() |
|
|
| for i in range(batch_size, ): |
| rand_kernels.append(torch.stack( |
| [self.fade_kernels[j][rand_x[i]:rand_x[i] + self.image_size, |
| : self.image_size] for j in range(len(self.fade_kernels))])) |
| rand_kernels = torch.stack(rand_kernels) |
|
|
| if t is None: |
| t = self.num_timesteps |
| if times is None: |
| times = t |
|
|
| for i in range(t): |
| with torch.no_grad(): |
| if self.degradation_type == 'fade': |
| if 'Random' in self.fade_routine: |
| faded_recon_sample = torch.stack([rand_kernels[:, i].to(sample_device), |
| rand_kernels[:, i].to(sample_device), |
| rand_kernels[:, i].to(sample_device)], 1) * faded_recon_sample |
| else: |
| faded_recon_sample = self.fade_kernels[i].to(sample_device) * faded_recon_sample |
| elif self.degradation_type == 'kspace': |
| if rand_kernels is not None: |
| |
| k = torch.stack([rand_kernels[:, i]], 1) |
| faded_recon_sample = apply_ksu_kernel(faded_recon_sample, k, params_dict) |
| else: |
| |
| k = self.kspace_kernels[i] |
| faded_recon_sample = apply_ksu_kernel(faded_recon_sample, k, params_dict) |
|
|
| if self.discrete: |
| faded_recon_sample = (faded_recon_sample + 1) * 0.5 |
| faded_recon_sample = (faded_recon_sample * 255) |
| faded_recon_sample = faded_recon_sample.int().float() / 255 |
| faded_recon_sample = faded_recon_sample * 2 - 1 |
|
|
| x0_list = [] |
| xt_list = [] |
|
|
| while times: |
| step = torch.full((batch_size,), times - 1, dtype=torch.long).cuda() |
| if self.backbone == "unet": |
| recon_sample = self.restore_fn(faded_recon_sample, step) |
| elif self.backbone == "twounet": |
| recon_sample = self.restore_fn(faded_recon_sample, aux, k, step) |
|
|
| elif self.backbone == "twobranch": |
| recon_sample, recon_fre = self.restore_fn(faded_recon_sample, aux, step) |
| recon_sample = recon_sample |
|
|
| x0_list.append(recon_sample) |
|
|
| if self.degradation_type == 'fade': |
| if self.sampling_routine == 'default': |
| for i in range(times - 1): |
| with torch.no_grad(): |
| if rand_kernels is not None: |
| recon_sample = torch.stack([rand_kernels[:, i].to(sample_device), |
| rand_kernels[:, i].to(sample_device), |
| rand_kernels[:, i].to(sample_device)], 1) * recon_sample |
| else: |
| recon_sample = self.fade_kernels[i].to(sample_device) * recon_sample |
| faded_recon_sample = recon_sample |
|
|
| elif self.sampling_routine == 'x0_step_down': |
| for i in range(t): |
| with torch.no_grad(): |
| recon_sample_sub_1 = recon_sample |
| if rand_kernels is not None: |
|
|
| recon_sample = apply_ksu_kernel(recon_sample, rand_kernels[i], params_dict) |
| else: |
| recon_sample = apply_ksu_kernel(recon_sample, self.kspace_kernels[i], params_dict) |
|
|
| faded_recon_sample = faded_recon_sample - recon_sample + recon_sample_sub_1 |
|
|
| elif self.degradation_type == 'kspace': |
| |
| if self.sampling_routine == 'default': |
| for i in range(t - 1): |
| with torch.no_grad(): |
| if rand_kernels is not None: |
| k = torch.stack([rand_kernels[:, i]], 1) |
| recon_sample = apply_ksu_kernel(recon_sample, k, params_dict) |
| else: |
| recon_sample = apply_ksu_kernel(recon_sample, self.kspace_kernels[i], params_dict) |
|
|
| faded_recon_sample = recon_sample |
|
|
| elif self.sampling_routine == 'x0_step_down': |
| for i in range(t): |
| with torch.no_grad(): |
| recon_sample_sub_1 = recon_sample |
| if rand_kernels is not None: |
| k = torch.stack([rand_kernels[:, i]], 1) |
| recon_sample = apply_ksu_kernel(recon_sample, k, params_dict) |
| else: |
| recon_sample = apply_ksu_kernel(recon_sample, self.kspace_kernels[i], params_dict) |
|
|
| faded_recon_sample = faded_recon_sample - recon_sample + recon_sample_sub_1 |
|
|
|
|
| xt_list.append(faded_recon_sample) |
| times -= 1 |
|
|
| return x0_list, xt_list |
|
|
| |
| def q_sample(self, x_start, t, params_dict=None, use_fre_noise=False): |
| x = x_start |
|
|
| with torch.no_grad(): |
| k = torch.stack([self.kspace_kernels[t]], 1) |
| x = apply_ksu_kernel(x, k, params_dict, use_fre_noise=use_fre_noise) |
|
|
| return x, k |
|
|
|
|
| def reconstruct_loss(self, x_start, x_recon): |
| if self.loss_type == 'l1': |
| loss = (x_start - x_recon).abs().mean() |
| elif self.loss_type == 'l2': |
| loss = func.mse_loss(x_start, x_recon) |
| else: |
| raise NotImplementedError() |
| return loss |
|
|
|
|
| def gaussian_kernel(self, size: int, sigma: float): |
| """Generates a 2D Gaussian kernel.""" |
| x = torch.arange(size).float() - size // 2 |
| gauss = torch.exp(-x ** 2 / (2 * sigma ** 2)) |
| kernel = gauss[:, None] @ gauss[None, :] |
| kernel /= kernel.sum() |
| return kernel.cuda() |
|
|
| def gaussian_blur(self, input_tensor, kernel_size: int, sigma: float): |
| """Applies Gaussian blur to a 4D tensor (N, C, H, W).""" |
| |
| kernel = self.gaussian_kernel(kernel_size, sigma).unsqueeze(0).unsqueeze(0) |
| kernel = kernel.expand(input_tensor.size(1), 1, kernel_size, kernel_size) |
|
|
| |
| padding = kernel_size // 2 |
| input_tensor = F.pad(input_tensor, (padding, padding, padding, padding), mode='reflect') |
|
|
| |
| blurred = F.conv2d(input_tensor, kernel, groups=input_tensor.size(1)) |
| return blurred |
|
|
| def get_frequency_elements(self, x): |
| x_fft = torch.fft.rfft2(x, norm="ortho") |
|
|
| |
| x_mag = torch.clamp(torch.abs(x_fft), min=1e-8) |
| x_phase = torch.angle(x_fft) |
|
|
| return x_mag, x_phase |
|
|
| def get_fre_kl_loss(self, pred_spa, pred_fre, target, k): |
| |
| B = pred_spa.shape[0] |
| k = k.contiguous().view(B, -1) |
| pred_spa = pred_spa.view(B, -1) |
| pred_fre = pred_fre.view(B, -1) |
| target = target.view(B, -1) |
|
|
| |
| pred_spa = pred_spa - torch.max(pred_spa, dim=1, keepdim=True).values |
| pred_fre = pred_fre - torch.max(pred_fre, dim=1, keepdim=True).values |
| target = target - torch.max(target, dim=1, keepdim=True).values |
|
|
| target_prob = F.softmax(target, dim=1) |
|
|
| ele_num = 2 |
| target_all_prob = torch.cat([target_prob for ii in range(ele_num)], dim=0) |
| k_mask = torch.cat([k for ii in range(ele_num)], dim=0) |
| k_total = torch.sum(k_mask) |
|
|
| pred_all = torch.cat([pred_spa, pred_fre]) |
| pred_all = F.log_softmax(pred_all, dim=1) |
|
|
|
|
| |
| |
| pred_spa_prob = F.softmax(pred_spa, dim=1) |
| pred_fre_prob = F.softmax(pred_fre, dim=1) |
| pred_avg_prob = 1.0 / ele_num * (pred_spa_prob + pred_fre_prob) |
| pred_avg_prob = torch.cat([pred_avg_prob for ii in range(ele_num)], dim=0).clone().detach() |
|
|
| kl_consist_loss = (self.kl_loss(pred_all, pred_avg_prob) * k_mask).sum() / k_total |
|
|
| return kl_consist_loss |
|
|
|
|
| def frequency_consistency_loss(self, pred_spa, pred_fre, target, k): |
| ''' |
| KL-term, enforcing conditional distribution remains unchanged regardless of interventions applied |
| ''' |
|
|
| W = pred_spa.shape[-1] |
| half_W = W // 2 + 1 |
| k = (1 - k.to(pred_spa.device)) |
| k = k[..., :half_W] |
|
|
| pred_spa_mag, pred_spa_pha = self.get_frequency_elements(pred_spa) |
| pred_fre_mag, pred_fre_pha = self.get_frequency_elements(pred_fre) |
| target_mag, target_pha = self.get_frequency_elements(target) |
|
|
| mag_kl_loss = self.get_fre_kl_loss(pred_spa_mag, pred_fre_mag, target_mag, k) |
| pha_kl_loss = self.get_fre_kl_loss(pred_spa_pha, pred_fre_pha, target_pha, k) |
|
|
| |
| return mag_kl_loss + pha_kl_loss |
|
|
|
|
| def p_losses(self, x_start, aux, t, params_dict): |
| self.debug_print = False |
| self.debug_time = False |
|
|
| start_time = time.time() |
|
|
| x_start_golden = x_start.clone() |
| x_mix, k = self.q_sample(x_start=x_start, t=t, params_dict=params_dict) |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
|
|
| x_mix = x_mix.detach() |
| aux = aux.detach() |
| x_start_golden = x_start_golden.detach() |
| k = k.detach() |
|
|
| if self.debug_time: |
| print("--------------------") |
| print("sample time=", time.time() - start_time) |
|
|
|
|
| if self.backbone == 'unet': |
| x_recon = self.restore_fn(x_mix, t) |
| loss = self.reconstruct_loss(x_start_golden, x_recon) |
|
|
| if self.use_lpips: |
| lpips_weight = 0.1 |
| lpips_loss = lpips_weight * self.lpips(x_recon, x_start_golden).mean() |
| loss += lpips_loss |
|
|
|
|
| if self.use_fre_loss: |
| fft_weight = 0.01 |
|
|
| fre_loss = fft_weight * self.amploss(x_recon, x_start_golden, k) |
| loss += fre_loss |
|
|
|
|
| elif self.backbone == 'twounet': |
| x_recon = self.restore_fn(x_mix, aux, k, t) |
| loss = self.reconstruct_loss(x_start_golden, x_recon) * 5.0 |
|
|
| |
| if self.use_lpips: |
| lpips_weight = 0.1 |
| lpips_loss = self.lpips(x_recon, x_start_golden).mean() |
| loss += lpips_weight * lpips_loss |
|
|
|
|
| if self.use_fre_loss: |
| fft_weight = 0.1 |
| amp = self.amploss(x_recon, x_start_golden, k) |
| loss += fft_weight * amp |
|
|
|
|
| elif self.backbone == 'twobranch': |
|
|
| if self.fp16: |
| with autocast(): |
| x_recon, x_recon_fre = self.restore_fn(x_mix, aux, t) |
| else: |
| x_recon, x_recon_fre = self.restore_fn(x_mix, aux, t) |
| if self.debug_time: |
| print("restore_fn time=", time.time() - start_time) |
|
|
| |
| |
| |
| |
| |
|
|
| loss_spatial = self.reconstruct_loss(x_start_golden, x_recon) |
| loss_freq = self.reconstruct_loss(x_start_golden, x_recon_fre) |
| loss = loss_spatial + loss_freq |
|
|
| if self.debug_time: |
| print("reconstruct_loss time=", time.time() - start_time) |
|
|
| |
| if self.use_lpips: |
| lpips_weight = 0.1 |
| lpips_loss = lpips_weight * self.lpips(x_recon, x_start_golden).mean() |
| loss += lpips_loss |
|
|
| if self.use_ssim: |
| ssim_weight = 0.1 |
| ssim_loss = 1.0 - self.ssim(x_recon, x_start_golden).mean() |
| loss += ssim_weight * ssim_loss |
|
|
| if self.use_fre_loss: |
| fft_weight = 0.01 |
|
|
| fre_loss = fft_weight * self.amploss(x_recon_fre, x_start_golden, k) |
| loss += fre_loss |
|
|
| |
| |
|
|
| if self.use_kl: |
| amp_fre = fft_weight * self.frequency_consistency_loss(x_recon, x_recon_fre, x_start_golden, k) |
| loss += amp_fre |
|
|
| if self.debug_time: |
| print("fre loss time=", time.time() - start_time) |
| print("--------------------") |
|
|
| if np.random.rand() < 0.001: |
| print("----------------------------------------\n" |
| "loss_spatial:", loss_spatial.item(), |
| "loss_freq:", loss_freq.item(), |
| |
| |
| "fre_loss:", fre_loss.item()) |
|
|
| print("----------------------------------------\n" |
| "x_recon:", x_recon.min().item(), x_recon.max().item(), |
| "x_recon_fre:", x_recon_fre.min().item(), x_recon_fre.max().item(), |
| "x_start_golden:", x_start_golden.min().item(), x_start_golden.max().item()) |
|
|
| return loss |
|
|
|
|
| def forward(self, x1, x2=None, params_dict=None, *args, **kwargs): |
| b, c, h, w, device, img_size, = *x1.shape, x1.device, self.image_size |
| assert h == img_size and w == img_size, f'height and width of image must be {img_size}' |
| t = torch.randint(0, self.num_timesteps, (b,), device=device).long() |
|
|
| loss = self.p_losses(x1, x2, t, params_dict, *args, **kwargs) |
| return loss |
|
|
|
|
|
|
| class Trainer(nn.Module): |
| def __init__( |
| self, |
| diffusion_model, |
| folder, |
| mode, |
| *, |
| norm="mean_std", |
| ema_decay=0.995, |
| image_size=128, |
| train_batch_size=32, |
| train_lr=2e-5, |
| train_num_steps=700000, |
| gradient_accumulate_every=2, |
| fp16=False, |
| step_start_ema=2000, |
| update_ema_every=100, |
| save_and_sample_every=1000, |
| results_folder='./results', |
| load_path=None, |
| dataset=None, |
| shuffle=True, |
| domain=None, |
| aux_modality=None, |
| num_channels=1, |
| debug=False, |
| ): |
| super().__init__() |
|
|
| self.mode = mode |
| self.model = diffusion_model |
| self.ema = EMA(ema_decay) |
| self.ema_model = copy.deepcopy(self.model) |
| self.update_ema_every = update_ema_every |
|
|
| self.step_start_ema = step_start_ema |
| self.save_and_sample_every = save_and_sample_every if not debug else 10 |
|
|
| self.batch_size = train_batch_size |
| self.image_size = diffusion_model.module.image_size |
| self.gradient_accumulate_every = gradient_accumulate_every |
| self.train_num_steps = train_num_steps |
| self.input_normalize = norm |
|
|
|
|
| if dataset == 'train': |
| print(dataset, "DA used") |
| self.ds = Dataset_Aug1(folder, image_size) |
|
|
| elif dataset.lower() == 'brain': |
| print(dataset, "Brain DA used", "mode=", mode) |
| |
| self.ds = BrainDataset(mode, folder, image_size, 4, |
| debug=debug, |
| domains=domain, |
| num_channels=num_channels, |
| aux_modality=aux_modality) |
|
|
|
|
|
|
| elif dataset.lower() == 'fsm_brain': |
| from dataset.BRATS_dataloader import Hybrid, ToTensor, RandomPadCrop, AddNoise |
| from torchvision import transforms |
| train_data_path = folder |
|
|
| self.ds = Hybrid(split=mode, SNR=0, |
| |
| transform=transforms.Compose([RandomPadCrop(), ToTensor()]), |
| base_dir=train_data_path, input_normalize=norm, |
| image_size=image_size, debug=debug) |
|
|
|
|
| self.test_ds = Hybrid(split="test", SNR=0, |
| |
| transform=transforms.Compose([RandomPadCrop(), ToTensor()]), |
| base_dir=train_data_path, input_normalize=norm, |
| image_size=image_size, debug=debug) |
|
|
| elif dataset.lower() == 'fsm_knee': |
| from dataset.fastmri import SliceDataset |
| root_path = folder |
| transforms = None |
|
|
| |
| self.ds = SliceDataset(root_path, transforms, 'singlecoil', |
| image_size=image_size, input_normalize=norm, |
| sample_rate=1, mode=mode, debug=debug) |
|
|
| self.test_ds = SliceDataset(root_path, transforms, 'singlecoil', |
| image_size=image_size, input_normalize=norm, |
| sample_rate=1, mode="test", debug=debug) |
|
|
| elif dataset.lower() == 'fsm_m4raw': |
| pass |
|
|
|
|
|
|
| else: |
| print(dataset) |
| self.ds = Dataset(folder, image_size) |
|
|
| self.train_batch_size = train_batch_size |
| self.batch_size = train_batch_size if mode == 'train' else 1 |
|
|
| self.dl = cycle( |
| data.DataLoader(self.ds, |
| batch_size=train_batch_size if mode == 'train' else 1, |
| shuffle=(mode == "train"), |
| pin_memory=True, |
| num_workers=16, |
| drop_last=True)) |
|
|
| self.test_dl = cycle( |
| data.DataLoader(self.ds, |
| batch_size=train_batch_size, |
| shuffle="test", |
| pin_memory=True, |
| num_workers=16, |
| drop_last=False)) |
|
|
| self.opt = AdamW(list(self.model.module.restore_fn.parameters()), |
| lr=train_lr, |
| betas=(0.9, 0.999), |
| weight_decay=1e-4) |
|
|
| self.scheduler = lr_scheduler.StepLR(self.opt, step_size=10000, gamma=0.5) |
| self.step = 0 |
|
|
| |
|
|
| self.fp16 = fp16 |
| os.makedirs(results_folder, exist_ok=True) |
| self.results_folder = Path(results_folder) |
|
|
| np.save(str(self.results_folder / "kspace_kernels.npy"), self.model.module.kspace_kernels.cpu()) |
|
|
| self.lpips = LPIPS().eval().cuda() |
|
|
| self.reset_parameters() |
|
|
| if load_path is not None: |
| self.load(load_path) |
| kspace_npy = load_path.replace('model.pt', 'kspace_kernels.npy') |
| self.model.module.kspace_kernels = torch.from_numpy(np.load(kspace_npy)).to(self.model.module.kspace_kernels.device) |
| self.ema_model.module.kspace_kernels = self.model.module.kspace_kernels |
|
|
| def reset_parameters(self): |
| self.ema_model.load_state_dict(self.model.state_dict()) |
|
|
| def step_ema(self): |
| if self.step < self.step_start_ema: |
| self.reset_parameters() |
| return |
| self.ema.update_model_average(self.ema_model, self.model) |
|
|
| def save(self): |
| model_data = { |
| 'step': self.step, |
| 'model': self.model.state_dict(), |
| 'ema': self.ema_model.state_dict() |
| } |
| save_name = str(self.results_folder / f'model.pt') |
| print("Save_name=", save_name) |
| torch.save(model_data, save_name) |
|
|
| def load(self, load_path): |
| print("Loading : ", load_path) |
| model_data = torch.load(load_path) |
|
|
| self.step = model_data['step'] |
| self.model.load_state_dict(model_data['model']) |
| self.ema_model.load_state_dict(model_data['ema']) |
| print("Loading complete") |
|
|
| @staticmethod |
| def add_title(path, title): |
| img1 = cv2.imread(path) |
|
|
| black = [0, 0, 0] |
| constant = cv2.copyMakeBorder(img1, 10, 10, 10, 10, cv2.BORDER_CONSTANT, value=black) |
| height = 20 |
| violet = np.zeros((height, constant.shape[1], 3), np.uint8) |
| violet[:] = (255, 0, 180) |
|
|
| vcat = cv2.vconcat((violet, constant)) |
|
|
| font = cv2.FONT_HERSHEY_SIMPLEX |
|
|
| cv2.putText(vcat, str(title), (violet.shape[1] // 2, height - 2), font, 0.5, (0, 0, 0), 1, 0) |
| cv2.imwrite(path, vcat) |
|
|
| |
| def calculate_metrics(self, all_images, og_img): |
| img_ = all_images.cpu() |
| og_img_ = og_img.cpu() |
|
|
| |
|
|
| |
| |
| ssims_ = [] |
| for (img, og_img) in zip(img_, og_img_): |
| img_np = img.squeeze().numpy() |
| og_img_np = og_img.squeeze().numpy() |
|
|
| |
| ssim_ = structural_similarity(og_img_np, img_np) |
| ssims_.append(ssim_) |
|
|
| ssim_ = np.mean(ssims_) |
|
|
| psnr_ = peak_signal_noise_ratio(og_img_.numpy(), img_.numpy()).mean() |
| nmse_ = nmse(og_img_.numpy(), img_.numpy()).mean() |
|
|
|
|
| return ssim_, psnr_, nmse_ |
|
|
| |
|
|
| |
| def calculate_metrics_3d(self, all_images, og_img): |
| all_images = torch.clamp(all_images, 1e-6, 1) |
| og_img = torch.clamp(og_img, 1e-6, 1) |
|
|
| |
| all_images_new = all_images.unsqueeze(1).repeat(1, 3, 1, 1) |
| og_img_new = og_img.unsqueeze(1).repeat(1, 3, 1, 1) |
|
|
| cal_fid = False |
| if cal_fid: |
| |
| fid_value = calculate_fid(all_images_new.cpu().numpy(), og_img_new.cpu().numpy(), |
| use_multiprocessing=False, batch_size=og_img.shape[-1]) |
| |
| fid_value_3d = calculate_fid_3d(all_images_new.cpu().numpy(), og_img_new.cpu().numpy(), |
| use_multiprocessing=False, batch_size=og_img.shape[-1]) |
| else: |
| fid_value = 0 |
| fid_value_3d = 0 |
|
|
| |
| lpips = self.lpips(all_images_new.cuda(), og_img_new.cuda()).mean().item() |
|
|
| |
| img_ = all_images.cpu().unsqueeze(0) |
| og_img_ = og_img.cpu().unsqueeze(0) |
|
|
| |
| ssim = StructuralSimilarityIndexMeasure(data_range=1.0) |
| ssim_ = ssim(og_img_, img_).mean() |
| psnr_ = psnr(og_img_.numpy(), img_.numpy(), data_range=1.0).mean() |
|
|
| return ssim_, psnr_, fid_value, fid_value_3d, lpips |
|
|
| |
| def test_data_dict(self, tag, data_dict, batches, routine): |
| print("=== Test tag: ", tag) |
|
|
| og_img = data_dict['img'].cuda() |
| aux = data_dict['aux'].cuda() |
| img_mean = data_dict['img_mean'].cuda() |
| img_std = data_dict['img_std'].cuda() |
| aux_mean = data_dict['aux_mean'].cuda() |
| aux_std = data_dict['aux_std'].cuda() |
| params_dict = {'img_mean': img_mean, 'img_std': img_std, 'aux_mean': aux_mean, 'aux_std': aux_std} |
|
|
|
|
| |
| |
|
|
| xt, direct_recons, all_images, return_k, all_recons, all_recons_fre, all_masks = ( |
| self.ema_model.module.sample( |
| batch_size=batches, |
| faded_recon_sample=og_img, |
| aux=aux, params_dict=params_dict, |
| sample_routine=routine)) |
|
|
| img_std = img_std.view(-1, 1, 1, 1) |
| img_mean = img_mean.view(-1, 1, 1, 1) |
| aux_std = aux_std.view(-1, 1, 1, 1) |
| aux_mean = aux_mean.view(-1, 1, 1, 1) |
|
|
| og_img = og_img * img_std + img_mean |
|
|
| _min = og_img.min() |
| _max = og_img.max() |
|
|
| og_img = (og_img - _min) / (_max - _min) |
| all_images = all_images * img_std + img_mean |
| all_images = (all_images - _min) / (_max - _min) |
| all_recons = all_recons * img_std + img_mean |
| all_recons = (all_recons - _min) / (_max - _min) |
| all_images = all_recons[-1] |
|
|
| direct_recons = direct_recons * img_std + img_mean |
| direct_recons = (direct_recons - _min) / (_max - _min) |
| xt = xt * img_std + img_mean |
| xt = (xt - _min) / (_max - _min) |
|
|
| aux = aux * aux_std + aux_mean |
| aux = (aux - aux.min()) / (aux.max() - aux.min()) |
|
|
| |
| |
| |
| |
| |
|
|
| all_recons = torch.clamp(all_recons, 1e-6, 1).mul(255).to(torch.int8) |
| direct_recons = torch.clamp(direct_recons, 1e-6, 1).mul(255).to(torch.int8) |
| all_images = torch.clamp(all_images, 1e-6, 1).mul(255).to(torch.int8) |
| og_img = torch.clamp(og_img, 1e-6, 1).mul(255).to(torch.int8) |
|
|
|
|
|
|
| |
| |
| ssims = [] |
| psnrs = [] |
|
|
| ssim_, psnr_, nmse_ = self.calculate_metrics(og_img, direct_recons) |
| lpips = self.lpips(direct_recons.float()/255, og_img.float()/255).mean().item() |
| print(f"=== first step Metrics {routine}: SSIM: ", ssim_, " PSNR: ", psnr_, |
| " LPIPS: ", lpips, " NMSE: ", nmse_) |
|
|
| for im in all_recons: |
| im = torch.clamp(im, 1e-6, 1).mul(255).to(torch.int8) |
| ssim_, psnr_, nmse_ = self.calculate_metrics(og_img, im) |
| ssims.append(ssim_) |
| psnrs.append(psnr_) |
|
|
|
|
| ssim_, psnr_, nmse_ = self.calculate_metrics(og_img, all_images) |
| lpips = self.lpips(all_images.float()/255, og_img.float()/255).mean().item() |
|
|
| print(f"=== Final Metrics {routine}: SSIM: ", ssim_, " PSNR: ", psnr_, |
| " LPIPS: ", lpips, " NMSE: ", nmse_) |
|
|
| os.makedirs(self.results_folder, exist_ok=True) |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| import matplotlib.pyplot as plt |
| fig, axs = plt.subplots(2, figsize=(8, 6), dpi=100, sharex=True) |
| axs[0].plot(ssims, marker='o', linestyle='-', color='blue', label='SSIM') |
| axs[0].set_title('Structural Similarity Index (SSIM)', fontsize=14) |
| axs[0].set_ylabel('SSIM', fontsize=12) |
| axs[0].grid(True, linestyle='--', alpha=0.7) |
| axs[0].legend() |
|
|
| |
| axs[1].plot(psnrs, marker='o', linestyle='-', color='green', label='PSNR') |
| axs[1].set_title('Peak Signal-to-Noise Ratio (PSNR)', fontsize=14) |
| axs[1].set_xlabel('Iterations', fontsize=12) |
| axs[1].set_ylabel('PSNR (dB)', fontsize=12) |
| axs[1].grid(True, linestyle='--', alpha=0.7) |
| axs[1].legend() |
|
|
| fig.tight_layout() |
|
|
| plt.savefig(str(self.results_folder / f'{self.step}-metrics-{routine}.png')) |
|
|
| return_k = return_k.cuda() |
|
|
|
|
| combine = torch.cat((return_k, |
| xt, |
| direct_recons, |
| all_images, |
| all_recons[-1].to(all_images.device), |
| og_img, aux), 2) |
|
|
| utils.save_image(combine, str(self.results_folder / f'{self.step}-combine-{routine}.png'), nrow=6) |
|
|
| |
| |
|
|
| all_recons = torch.cat(list(all_recons), dim=-1).cpu() |
| all_masks = all_masks.cpu() |
| |
|
|
| s = all_recons.shape[-2] |
| repeats = all_recons.shape[3] // og_img.shape[3] |
| |
| og_img = og_img.cpu() |
| all_recons_residual = all_recons - og_img.repeat(1, 1, 1, repeats) |
| all_recons_residual = (all_recons_residual - all_recons_residual.min()) / (all_recons_residual.max() - all_recons_residual.min()) |
|
|
| |
| all_recons_residual_2 = all_recons[:, :, :, s:] - all_recons[:, :, :, :-s] |
| all_recons_residual_2 = (all_recons_residual_2 - all_recons_residual_2.min()) / (all_recons_residual_2.max() - all_recons_residual_2.min()) |
| padding = torch.zeros_like(all_recons_residual[:, :, :, :s // 2]) |
| all_recons_residual_2 = torch.cat([padding, all_recons_residual_2, padding], dim=-1) |
|
|
| all_recons = torch.cat([all_recons, all_recons_residual_2, all_recons_residual], dim=-2) |
|
|
| utils.save_image(all_recons, str(self.results_folder / f'{self.step}-all_recons-{routine}.png'), |
| nrow=1) |
| |
| |
|
|
| |
| print(f'Mean of last {self.step}: save to :', str(self.results_folder / f'{self.step}-combine.png')) |
|
|
| |
|
|
|
|
| def train(self): |
| backwards = partial(loss_backwards, self.fp16) |
| |
|
|
| acc_loss = 0 |
| start_time = time.time() |
|
|
| while self.step < self.train_num_steps: |
| d_time = time.time() |
| self.opt.zero_grad() |
| u_loss = 0 |
|
|
| if (self.step + 1 )% 20000 == 0: |
| self.ds.update_chunk() |
| self.dl = cycle( |
| data.DataLoader(self.ds, |
| batch_size=self.train_batch_size, |
| shuffle=True, |
| pin_memory=True, |
| num_workers=16, |
| drop_last=True)) |
| self.test_dl = cycle( |
| data.DataLoader(self.test_ds, |
| batch_size=self.train_batch_size, |
| shuffle=False, |
| pin_memory=True, |
| num_workers=16, |
| drop_last=False)) |
|
|
| self.debug_time = False |
|
|
| for i in range(self.gradient_accumulate_every): |
|
|
| last_model_state = self.model.state_dict() |
| optimizer_state = self.opt.state_dict() |
|
|
| data_dict = next(self.dl) |
| if self.debug_time: |
| print("Data loading time=", time.time() - d_time) |
|
|
| img = data_dict['img'].cuda() |
| aux = data_dict['aux'].cuda() |
| img_mean = data_dict['img_mean'].cuda() |
| img_std = data_dict['img_std'].cuda() |
| aux_mean = data_dict['aux_mean'].cuda() |
| aux_std = data_dict['aux_std'].cuda() |
| params_dict = {'img_mean': img_mean, 'img_std': img_std, 'aux_mean': aux_mean, 'aux_std': aux_std} |
|
|
|
|
| loss = torch.mean(self.model(img, aux, params_dict)) |
| if self.debug_time: |
| print("Model iter=", self.step, " time=", time.time() - d_time) |
|
|
| if torch.isnan(loss).any(): |
| print(f"NaN encountered in step {self.step}. Reverting model.") |
| self.model.load_state_dict(last_model_state) |
| self.opt.load_state_dict(optimizer_state) |
| continue |
| if self.debug_time: |
| print("before loss=", time.time() - d_time) |
| u_loss += loss |
| loss.backward() |
| |
| if self.debug_time: |
| print("after loss=", time.time() - d_time) |
|
|
| del img, aux, img_mean, img_std, aux_mean, aux_std, params_dict |
|
|
|
|
|
|
| if (self.step + 1) % (min(self.train_num_steps // 100 + 1, 100)) == 0: |
| print('iteration %d [%.2f sec]: learning rate : %f loss : %f ' % ( |
| self.step + 1, time.time() - start_time, self.scheduler.get_lr()[0], loss.item())) |
|
|
| |
| acc_loss = acc_loss + (u_loss.item() / self.gradient_accumulate_every) |
|
|
| max_norm = 0.01 |
| torch.nn.utils.clip_grad_norm_(self.model.module.restore_fn.parameters(), max_norm) |
| if self.debug_time: |
| print("Before Optimization time=", time.time() - d_time) |
|
|
| if self.fp16: |
| scaler.step(self.opt) |
| self.scheduler.step() |
| scaler.update() |
|
|
| else: |
| self.opt.step() |
| self.scheduler.step() |
|
|
| if self.debug_time: |
| print("Optimization time=", time.time() - d_time) |
|
|
| if self.step % self.update_ema_every == 0: |
| self.step_ema() |
|
|
|
|
| |
| if self.step != 0 and (self.step + 1) % self.save_and_sample_every == 0: |
| batches = self.batch_size |
| data_dict = next(self.test_dl) |
| train_dict = next(self.dl) |
|
|
| |
| for routine in ['x0_step_down_fre']: |
| self.test_data_dict("Train", train_dict, batches, routine) |
| self.test_data_dict("Test", data_dict, batches, routine) |
|
|
| self.ema_model.module.restore_fn.train() |
| self.save() |
|
|
| self.step += 1 |
| clean_start = time.time() |
|
|
| del data_dict |
| del u_loss |
| torch.cuda.empty_cache() |
| gc.collect() |
| if self.debug_time: |
| print("Clean time=", time.time() - clean_start) |
|
|
| print("Iter time = ", time.time() - d_time, "total time = ", time.time() - start_time) |
|
|
| print('training completed') |
|
|
| def test_loader(self, sampling_routine): |
| """ |
| Computes patient-wise 3D for a dataset of MRI slices. |
| |
| """ |
| print("Starting testing with sampling routine: ", sampling_routine) |
|
|
| model = self.model |
| model.eval() |
|
|
| |
| num_timesteps = model.module.num_timesteps |
| patient_wise_pred = {} |
| patient_wise_gt = [] |
| count = 1 |
|
|
| while True: |
| batches = 1 |
| data_dict = next(self.dl) |
|
|
|
|
| |
| og_img = data_dict['img'].cuda() |
| aux = data_dict['aux'].cuda() |
| img_mean = data_dict['img_mean'].cuda() |
| img_std = data_dict['img_std'].cuda() |
| aux_mean = data_dict['aux_mean'].cuda() |
| aux_std = data_dict['aux_std'].cuda() |
| params_dict = {'img_mean': img_mean, 'img_std': img_std, 'aux_mean': aux_mean, 'aux_std': aux_std} |
|
|
| xt, direct_recons, all_images, return_k, all_recons, all_recons_fre, all_masks = ( |
| model.module.sample( |
| batch_size=batches, |
| faded_recon_sample=og_img, |
| aux=aux, params_dict=params_dict, |
| sample_routine=sampling_routine)) |
|
|
| img_std = img_std.view(-1, 1, 1, 1) |
| img_mean = img_mean.view(-1, 1, 1, 1) |
| aux_std = aux_std.view(-1, 1, 1, 1) |
| aux_mean = aux_mean.view(-1, 1, 1, 1) |
|
|
|
|
| print("before or_img shape: ", og_img.shape, og_img.min(), og_img.max()) |
| print("before direct_recons shape: ", direct_recons.shape, direct_recons.min(), direct_recons.max()) |
| |
| direct_recons_norm = (direct_recons - direct_recons.mean()) / (direct_recons.std()) |
| all_images_norm = (all_images - all_images.mean()) / (all_images.std()) |
|
|
| og_img = og_img * img_std + img_mean |
| _min = og_img.min() |
| _max = og_img.max() |
|
|
| |
| all_images = all_images * img_std + img_mean |
| |
| all_images_norm = all_images_norm * img_std + img_mean |
| |
|
|
| all_recons = all_recons * img_std + img_mean |
| |
| direct_recons = direct_recons * img_std + img_mean |
| |
|
|
| direct_recons_norm = direct_recons_norm * img_std + img_mean |
| |
|
|
| xt = xt * img_std + img_mean |
| |
|
|
| aux = aux * aux_std + aux_mean |
| |
| all_recons = all_recons.cpu() |
|
|
| all_recons = torch.clamp(all_recons, 1e-6, 1).mul(255).to(torch.int8) |
| direct_recons = torch.clamp(direct_recons, 1e-6, 1).mul(255).to(torch.int8) |
| all_images = torch.clamp(all_images, 1e-6, 1).mul(255).to(torch.int8) |
| og_img = torch.clamp(og_img, 1e-6, 1).mul(255).to(torch.int8) |
|
|
| |
| |
| ssims = [] |
| psnrs = [] |
|
|
|
|
|
|
| ssim_, psnr_, nmse_ = self.calculate_metrics(og_img, direct_recons) |
| lpips = self.lpips(direct_recons.float(), og_img.float()).mean().item() |
| print(f"=== first step Metrics {sampling_routine}: SSIM: ", ssim_, " PSNR: ", psnr_, " LPIPS: ", lpips, " NMSE: ", nmse_) |
|
|
| for im in all_recons: |
| ssim_, psnr_, nmse_ = self.calculate_metrics(og_img, im) |
| ssims.append(ssim_) |
| psnrs.append(psnr_) |
|
|
| ssim_, psnr_, nmse_ = self.calculate_metrics(og_img, all_images) |
| lpips = self.lpips(all_images.float(), og_img.float()).mean().item() |
|
|
| print(f"=== Final Metrics {sampling_routine}: SSIM: ", ssim_, " PSNR: ", psnr_, " LPIPS: ", lpips, " NMSE: ", nmse_) |
|
|
|
|
| os.makedirs(self.results_folder, exist_ok=True) |
|
|
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| import matplotlib.pyplot as plt |
| fig, axs = plt.subplots(2, figsize=(8, 6), dpi=100, sharex=True) |
| axs[0].plot(ssims, marker='o', linestyle='-', color='blue', label='SSIM') |
| axs[0].set_title('Structural Similarity Index (SSIM)', fontsize=14) |
| axs[0].set_ylabel('SSIM', fontsize=12) |
| axs[0].grid(True, linestyle='--', alpha=0.7) |
| axs[0].legend() |
|
|
| |
| axs[1].plot(psnrs, marker='o', linestyle='-', color='green', label='PSNR') |
| axs[1].set_title('Peak Signal-to-Noise Ratio (PSNR)', fontsize=14) |
| axs[1].set_xlabel('Iterations', fontsize=12) |
| axs[1].set_ylabel('PSNR (dB)', fontsize=12) |
| axs[1].grid(True, linestyle='--', alpha=0.7) |
| axs[1].legend() |
|
|
| fig.tight_layout() |
|
|
| plt.savefig(str(self.results_folder / f'{count}-metrics-{sampling_routine}.png')) |
|
|
| return_k = return_k.cuda() |
|
|
| combine = torch.cat((return_k, |
| xt, |
| all_images, direct_recons, og_img, aux), 2) |
|
|
| |
|
|
| |
| |
|
|
| all_recons = torch.cat(list(all_recons), dim=-1) |
| all_masks = all_masks.cpu() |
| all_masks = torch.cat(list(all_masks), dim=-1) |
|
|
| s = all_recons.shape[-2] |
| repeats = all_recons.shape[3] // og_img.shape[3] |
| |
| og_img = og_img.cpu() |
| all_recons_residual = all_recons - og_img.repeat(1, 1, 1, repeats) |
| |
| all_recons_residual_2 = all_recons[:, :, :, s:] - all_recons[:, :, :, :-s] |
| padding = torch.zeros_like(all_recons_residual[:, :, :, :s // 2]) |
| all_recons_residual_2 = torch.cat([padding, all_recons_residual_2, padding], dim=-1) |
|
|
| all_recons = torch.cat([all_recons, all_recons_residual_2, all_recons_residual], dim=-2) |
|
|
| |
| |
| count += 1 |
|
|
|
|
| def test_loader_3d(self, sampling_routine): |
| """ |
| Computes patient-wise 3D for a dataset of MRI slices. |
| |
| """ |
| print("Starting testing with sampling routine: ", sampling_routine) |
|
|
| model = self.model |
| model.eval() |
|
|
| |
| num_timesteps = model.module.num_timesteps |
| patient_wise_pred = {} |
| patient_wise_gt = [] |
|
|
| while True: |
| batches = 1 |
| data_dict = next(self.dl) |
| if data_dict['is_start']: |
| patient_wise_pred = {} |
| for i in range(num_timesteps): |
| patient_wise_pred[i] = [] |
| patient_wise_gt = [] |
|
|
| |
| og_img = data_dict['img'].cuda() |
| aux = data_dict['aux'].cuda() |
| img_mean = data_dict['img_mean'].cuda() |
| img_std = data_dict['img_std'].cuda() |
| aux_mean = data_dict['aux_mean'].cuda() |
| aux_std = data_dict['aux_std'].cuda() |
| params_dict = {'img_mean': img_mean, 'img_std': img_std, |
| 'aux_mean': aux_mean, 'aux_std': aux_std} |
|
|
| |
| xt, direct_recons, all_images, return_k, all_recons, all_recons_fre, all_masks = ( |
| model.module.sample( |
| batch_size=batches, |
| faded_recon_sample=og_img, |
| aux=aux, params_dict=params_dict, |
| sample_routine=sampling_routine)) |
|
|
| |
| img_std = img_std.view(-1, 1, 1, 1) |
| img_mean = img_mean.view(-1, 1, 1, 1) |
| aux_std = aux_std.view(-1, 1, 1, 1) |
| aux_mean = aux_mean.view(-1, 1, 1, 1) |
|
|
| og_img = og_img * img_std + img_mean |
| _min = og_img.min() |
| _max = og_img.max() |
|
|
| og_img = (og_img - _min) / (_max - _min) |
| all_images = all_images * img_std + img_mean |
| all_images = (all_images - _min) / (_max - _min) |
|
|
| all_recons = all_recons * img_std + img_mean |
| all_recons = (all_recons - _min) / (_max - _min) |
| direct_recons = direct_recons * img_std + img_mean |
| direct_recons = (direct_recons - _min) / (_max - _min) |
| xt = xt * img_std + img_mean |
| xt = (xt - _min) / (_max - _min) |
|
|
| aux = aux * aux_std + aux_mean |
| aux = (aux - aux.min()) / (aux.max() - aux.min()) |
| all_recons = all_recons.cpu() |
|
|
| |
| patient_wise_gt.append(og_img) |
| for i in range(num_timesteps): |
| patient_wise_pred[i].append(all_recons[i]) |
|
|
| if data_dict['is_end']: |
| |
| |
| patient_gt = torch.cat(patient_wise_gt, dim=0).squeeze() |
|
|
| ssims, psnrs, lpips, fid, fid_3d = [], [], [], [], [] |
|
|
| for i in range(num_timesteps): |
| patient_pred = torch.cat(patient_wise_pred[i], dim=0).squeeze() |
|
|
| ssim_, psnr_, fid_, fid3d_, lpips_ = self.calculate_metrics_3d(patient_gt, patient_pred) |
| print("time step: ", i, "SSIM: ", ssim_, " PSNR: ", psnr_, " LPIPS: ", lpips_, " FID: ", fid_, " FID_3D: ", fid3d_) |
|
|
| lpips.append(lpips_) |
| fid.append(fid_) |
| fid_3d.append(fid3d_) |
| ssims.append(ssim_) |
| psnrs.append(psnr_) |
|
|
| print(f"=== Final Metrics {sampling_routine}: SSIM: ", ssim_, " PSNR: ", psnr_, " LPIPS: ", lpips_) |
|
|
|
|
| file_id = data_dict['file_id'][0].split("/")[-1] |
|
|
| os.makedirs(self.results_folder, exist_ok=True) |
|
|
| import matplotlib.pyplot as plt |
| fig, axs = plt.subplots(5, figsize=(8, 6), dpi=100, sharex=True) |
| axs[0].plot(ssims, marker='o', linestyle='-', color='blue', label='SSIM') |
| axs[0].set_title('Structural Similarity Index (SSIM)', fontsize=14) |
| axs[0].set_ylabel('SSIM', fontsize=12) |
| axs[0].grid(True, linestyle='--', alpha=0.7) |
| axs[0].legend() |
|
|
| |
| axs[1].plot(psnrs, marker='o', linestyle='-', color='green', label='PSNR') |
| axs[1].set_title('Peak Signal-to-Noise Ratio (PSNR)', fontsize=14) |
| axs[1].set_xlabel('Iterations', fontsize=12) |
| axs[1].set_ylabel('PSNR (dB)', fontsize=12) |
| axs[1].grid(True, linestyle='--', alpha=0.7) |
| axs[1].legend() |
|
|
| |
| axs[2].plot(lpips, marker='o', linestyle='-', color='red', label='LPIPS') |
| axs[2].set_title('LPIPS', fontsize=14) |
| axs[2].set_xlabel('Iterations', fontsize=12) |
| axs[2].set_ylabel('LPIPS', fontsize=12) |
| axs[2].grid(True, linestyle='--', alpha=0.7) |
| axs[2].legend() |
|
|
| |
| axs[3].plot(fid, marker='o', linestyle='-', color='orange', label='FID') |
| axs[3].set_title('FID', fontsize=14) |
| axs[3].set_xlabel('Iterations', fontsize=12) |
| axs[3].set_ylabel('FID', fontsize=12) |
| axs[3].grid(True, linestyle='--', alpha=0.7) |
| axs[3].legend() |
|
|
| axs[4].plot(fid_3d, marker='o', linestyle='-', color='purple', label='FID_3D') |
| axs[4].set_title('FID_3D', fontsize=14) |
| axs[4].set_xlabel('Iterations', fontsize=12) |
| axs[4].set_ylabel('FID_3D', fontsize=12) |
| axs[4].grid(True, linestyle='--', alpha=0.7) |
| axs[4].legend() |
|
|
| fig.tight_layout() |
| save = str(self.results_folder / f'{file_id}-metrics-{sampling_routine}.png') |
| plt.savefig(save) |
|
|
| print("Save metrics to ", save) |
|
|
|
|
|
|
|
|
| def test_from_data(self, extra_path, s_times=None): |
| batches = self.batch_size |
| og_img = next(self.dl).cuda() |
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=batches, faded_recon_sample=og_img, times=s_times) |
|
|
| og_img = (og_img + 1) * 0.5 |
| utils.save_image(og_img, str(self.results_folder / f'og-{extra_path}.png'), nrow=6) |
|
|
| frames_t = [] |
| frames_0 = [] |
|
|
| for i in range(len(x0_list)): |
| print(i) |
|
|
| x_0 = x0_list[i] |
| x_0 = (x_0 + 1) * 0.5 |
| utils.save_image(x_0, str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'), nrow=6) |
| self.add_title(str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'), str(i)) |
| frames_0.append(imageio.imread(str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'))) |
|
|
| x_t = xt_list[i] |
| all_images = (x_t + 1) * 0.5 |
| utils.save_image(all_images, str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'), nrow=6) |
| self.add_title(str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'), str(i)) |
| frames_t.append(imageio.imread(str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'))) |
|
|
| imageio.mimsave(str(self.results_folder / f'Gif-{extra_path}-x0.gif'), frames_0) |
| imageio.mimsave(str(self.results_folder / f'Gif-{extra_path}-xt.gif'), frames_t) |
|
|
| def test_with_mixup(self, extra_path): |
| batches = self.batch_size |
| og_img_1 = next(self.dl).cuda() |
| og_img_2 = next(self.dl).cuda() |
| og_img = (og_img_1 + og_img_2) / 2 |
|
|
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=batches, faded_recon_sample=og_img) |
|
|
| og_img_1 = (og_img_1 + 1) * 0.5 |
| utils.save_image(og_img_1, str(self.results_folder / f'og1-{extra_path}.png'), nrow=6) |
|
|
| og_img_2 = (og_img_2 + 1) * 0.5 |
| utils.save_image(og_img_2, str(self.results_folder / f'og2-{extra_path}.png'), nrow=6) |
|
|
| og_img = (og_img + 1) * 0.5 |
| utils.save_image(og_img, str(self.results_folder / f'og-{extra_path}.png'), nrow=6) |
|
|
| frames_t = [] |
| frames_0 = [] |
|
|
| for i in range(len(x0_list)): |
| print(i) |
| x_0 = x0_list[i] |
| x_0 = (x_0 + 1) * 0.5 |
| utils.save_image(x_0, str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'), nrow=6) |
| self.add_title(str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'), str(i)) |
| frames_0.append(Image.open(str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'))) |
|
|
| x_t = xt_list[i] |
| all_images = (x_t + 1) * 0.5 |
| utils.save_image(all_images, str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'), nrow=6) |
| self.add_title(str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'), str(i)) |
| frames_t.append(Image.open(str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'))) |
|
|
| frame_one = frames_0[0] |
| frame_one.save(str(self.results_folder / f'Gif-{extra_path}-x0.gif'), format="GIF", append_images=frames_0, |
| save_all=True, duration=100, loop=0) |
|
|
| frame_one = frames_t[0] |
| frame_one.save(str(self.results_folder / f'Gif-{extra_path}-xt.gif'), format="GIF", append_images=frames_t, |
| save_all=True, duration=100, loop=0) |
|
|
| def test_from_random(self, extra_path): |
| batches = self.batch_size |
| og_img = next(self.dl).cuda() |
| og_img = og_img * 0.9 |
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=batches, faded_recon_sample=og_img) |
|
|
| og_img = (og_img + 1) * 0.5 |
| utils.save_image(og_img, str(self.results_folder / f'og-{extra_path}.png'), nrow=6) |
|
|
| frames_t_names = [] |
| frames_0_names = [] |
|
|
| for i in range(len(x0_list)): |
| print(i) |
|
|
| x_0 = x0_list[i] |
| x_0 = (x_0 + 1) * 0.5 |
| utils.save_image(x_0, str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'), nrow=6) |
| self.add_title(str(self.results_folder / f'sample-{i}-{extra_path}-x0.png'), str(i)) |
| frames_0_names.append(str(self.results_folder / f'sample-{i}-{extra_path}-x0.png')) |
|
|
| x_t = xt_list[i] |
| all_images = (x_t + 1) * 0.5 |
| utils.save_image(all_images, str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'), nrow=6) |
| self.add_title(str(self.results_folder / f'sample-{i}-{extra_path}-xt.png'), str(i)) |
| frames_t_names.append(str(self.results_folder / f'sample-{i}-{extra_path}-xt.png')) |
|
|
| frames_0 = [] |
| frames_t = [] |
| for i in range(len(x0_list)): |
| print(i) |
| frames_0.append(imageio.imread(frames_0_names[i])) |
| frames_t.append(imageio.imread(frames_t_names[i])) |
|
|
| imageio.mimsave(str(self.results_folder / f'Gif-{extra_path}-x0.gif'), frames_0) |
| imageio.mimsave(str(self.results_folder / f'Gif-{extra_path}-xt.gif'), frames_t) |
|
|
| def controlled_direct_reconstruct(self, extra_path): |
| batches = self.batch_size |
| torch.manual_seed(0) |
| og_img = next(self.dl).cuda() |
| xt, direct_recons, all_images = self.ema_model.module.sample(batch_size=batches, faded_recon_sample=og_img) |
|
|
| og_img = (og_img + 1) * 0.5 |
| utils.save_image(og_img, str(self.results_folder / f'sample-og-{extra_path}.png'), nrow=6) |
|
|
| all_images = (all_images + 1) * 0.5 |
| utils.save_image(all_images, str(self.results_folder / f'sample-recon-{extra_path}.png'), nrow=6) |
|
|
| direct_recons = (direct_recons + 1) * 0.5 |
| utils.save_image(direct_recons, str(self.results_folder / f'sample-direct_recons-{extra_path}.png'), nrow=6) |
|
|
| xt = (xt + 1) * 0.5 |
| utils.save_image(xt, str(self.results_folder / f'sample-xt-{extra_path}.png'), |
| nrow=6) |
|
|
| self.save() |
|
|
| def fid_distance_decrease_from_manifold(self, fid_func, start=0, end=1000): |
|
|
| all_samples = [] |
| dataset = self.ds |
|
|
| print(len(dataset)) |
| for idx in range(len(dataset)): |
| img = dataset[idx] |
| img = torch.unsqueeze(img, 0).cuda() |
| if idx > start: |
| all_samples.append(img[0]) |
| if idx % 1000 == 0: |
| print(idx) |
| if end is not None: |
| if idx == end: |
| print(idx) |
| break |
|
|
| all_samples = torch.stack(all_samples) |
| blurred_samples = None |
| original_sample = None |
| deblurred_samples = None |
| direct_deblurred_samples = None |
|
|
| sanity_check = blurred_samples |
|
|
| cnt = 0 |
| while cnt < all_samples.shape[0]: |
| og_x = all_samples[cnt: cnt + 50] |
| og_x = og_x.cuda() |
| og_x = og_x.type(torch.cuda.FloatTensor) |
| og_img = og_x |
| print(og_img.shape) |
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=og_img.shape[0], |
| faded_recon_sample=og_img, |
| times=None) |
|
|
| og_img = og_img.to('cpu') |
| blurry_imgs = xt_list[0].to('cpu') |
| deblurry_imgs = x0_list[-1].to('cpu') |
| direct_deblurry_imgs = x0_list[0].to('cpu') |
|
|
| og_img = og_img.repeat(1, 3 // og_img.shape[1], 1, 1) |
| blurry_imgs = blurry_imgs.repeat(1, 3 // blurry_imgs.shape[1], 1, 1) |
| deblurry_imgs = deblurry_imgs.repeat(1, 3 // deblurry_imgs.shape[1], 1, 1) |
| direct_deblurry_imgs = direct_deblurry_imgs.repeat(1, 3 // direct_deblurry_imgs.shape[1], 1, 1) |
|
|
| og_img = (og_img + 1) * 0.5 |
| blurry_imgs = (blurry_imgs + 1) * 0.5 |
| deblurry_imgs = (deblurry_imgs + 1) * 0.5 |
| direct_deblurry_imgs = (direct_deblurry_imgs + 1) * 0.5 |
|
|
| if cnt == 0: |
| print(og_img.shape) |
| print(blurry_imgs.shape) |
| print(deblurry_imgs.shape) |
| print(direct_deblurry_imgs.shape) |
|
|
| if sanity_check: |
| folder = './sanity_check/' |
| create_folder(folder) |
|
|
| san_imgs = og_img[0: 32] |
| utils.save_image(san_imgs, str(folder + f'sample-og.png'), nrow=6) |
|
|
| san_imgs = blurry_imgs[0: 32] |
| utils.save_image(san_imgs, str(folder + f'sample-xt.png'), nrow=6) |
|
|
| san_imgs = deblurry_imgs[0: 32] |
| utils.save_image(san_imgs, str(folder + f'sample-recons.png'), nrow=6) |
|
|
| san_imgs = direct_deblurry_imgs[0: 32] |
| utils.save_image(san_imgs, str(folder + f'sample-direct-recons.png'), nrow=6) |
|
|
| if blurred_samples is None: |
| blurred_samples = blurry_imgs |
| else: |
| blurred_samples = torch.cat((blurred_samples, blurry_imgs), dim=0) |
|
|
| if original_sample is None: |
| original_sample = og_img |
| else: |
| original_sample = torch.cat((original_sample, og_img), dim=0) |
|
|
| if deblurred_samples is None: |
| deblurred_samples = deblurry_imgs |
| else: |
| deblurred_samples = torch.cat((deblurred_samples, deblurry_imgs), dim=0) |
|
|
| if direct_deblurred_samples is None: |
| direct_deblurred_samples = direct_deblurry_imgs |
| else: |
| direct_deblurred_samples = torch.cat((direct_deblurred_samples, direct_deblurry_imgs), dim=0) |
|
|
| cnt += og_img.shape[0] |
|
|
| print(blurred_samples.shape) |
| print(original_sample.shape) |
| print(deblurred_samples.shape) |
| print(direct_deblurred_samples.shape) |
|
|
| fid_blur = fid_func(samples=[original_sample, blurred_samples]) |
| rmse_blur = torch.sqrt(torch.mean((original_sample - blurred_samples) ** 2)) |
| ssim_blur = ssim(original_sample, blurred_samples, data_range=1, size_average=True) |
| print(f'The FID of blurry images with original image is {fid_blur}') |
| print(f'The RMSE of blurry images with original image is {rmse_blur}') |
| print(f'The SSIM of blurry images with original image is {ssim_blur}') |
|
|
| fid_deblur = fid_func(samples=[original_sample, deblurred_samples]) |
| rmse_deblur = torch.sqrt(torch.mean((original_sample - deblurred_samples) ** 2)) |
| ssim_deblur = ssim(original_sample, deblurred_samples, data_range=1, size_average=True) |
| print(f'The FID of deblurred images with original image is {fid_deblur}') |
| print(f'The RMSE of deblurred images with original image is {rmse_deblur}') |
| print(f'The SSIM of deblurred images with original image is {ssim_deblur}') |
|
|
| print(f'Hence the improvement in FID using sampling is {fid_blur - fid_deblur}') |
|
|
| fid_direct_deblur = fid_func(samples=[original_sample, direct_deblurred_samples]) |
| rmse_direct_deblur = torch.sqrt(torch.mean((original_sample - direct_deblurred_samples) ** 2)) |
| ssim_direct_deblur = ssim(original_sample, direct_deblurred_samples, data_range=1, size_average=True) |
| print(f'The FID of direct deblurred images with original image is {fid_direct_deblur}') |
| print(f'The RMSE of direct deblurred images with original image is {rmse_direct_deblur}') |
| print(f'The SSIM of direct deblurred images with original image is {ssim_direct_deblur}') |
|
|
| print(f'Hence the improvement in FID using direct sampling is {fid_blur - fid_direct_deblur}') |
|
|
| def paper_invert_section_images(self, s_times=None): |
|
|
| cnt = 0 |
| for i in range(50): |
| batches = self.batch_size |
| og_img = next(self.dl).cuda() |
| print(og_img.shape) |
|
|
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=batches, |
| faded_recon_sample=og_img, |
| times=s_times) |
| og_img = (og_img + 1) * 0.5 |
|
|
| for j in range(og_img.shape[0]//3): |
| original = og_img[j: j + 1] |
| utils.save_image(original, str(self.results_folder / f'original_{cnt}.png'), nrow=3) |
|
|
| direct_recons = x0_list[0][j: j + 1] |
| direct_recons = (direct_recons + 1) * 0.5 |
| utils.save_image(direct_recons, str(self.results_folder / f'direct_recons_{cnt}.png'), nrow=3) |
|
|
| sampling_recons = x0_list[-1][j: j + 1] |
| sampling_recons = (sampling_recons + 1) * 0.5 |
| utils.save_image(sampling_recons, str(self.results_folder / f'sampling_recons_{cnt}.png'), nrow=3) |
|
|
| blurry_image = xt_list[0][j: j + 1] |
| blurry_image = (blurry_image + 1) * 0.5 |
| utils.save_image(blurry_image, str(self.results_folder / f'blurry_image_{cnt}.png'), nrow=3) |
|
|
| blurry_image = cv2.imread(f'{self.results_folder}/blurry_image_{cnt}.png') |
| direct_recons = cv2.imread(f'{self.results_folder}/direct_recons_{cnt}.png') |
| sampling_recons = cv2.imread(f'{self.results_folder}/sampling_recons_{cnt}.png') |
| original = cv2.imread(f'{self.results_folder}/original_{cnt}.png') |
|
|
| black = [0, 0, 0] |
| blurry_image = cv2.copyMakeBorder(blurry_image, 10, 10, 10, 10, cv2.BORDER_CONSTANT, value=black) |
| direct_recons = cv2.copyMakeBorder(direct_recons, 10, 10, 10, 10, cv2.BORDER_CONSTANT, value=black) |
| sampling_recons = cv2.copyMakeBorder(sampling_recons, 10, 10, 10, 10, cv2.BORDER_CONSTANT, value=black) |
| original = cv2.copyMakeBorder(original, 10, 10, 10, 10, cv2.BORDER_CONSTANT, value=black) |
|
|
| im_h = cv2.hconcat([blurry_image, direct_recons, sampling_recons, original]) |
| cv2.imwrite(f'{self.results_folder}/all_{cnt}.png', im_h) |
|
|
| cnt += 1 |
|
|
| def paper_showing_diffusion_images(self, s_times=None): |
|
|
| cnt = 0 |
| to_show = [0, 1, 2, 4, 8, 16, 32, 64, 90, 91, 92, 93, 94, 95, 96, 97, 98, 99, 100] |
|
|
| for i in range(100): |
| batches = self.batch_size |
| og_img = next(self.dl).cuda() |
| print(og_img.shape) |
|
|
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=batches, faded_recon_sample=og_img, times=s_times) |
|
|
| for k in range(xt_list[0].shape[0]): |
| lst = [] |
|
|
| for j in range(len(xt_list)): |
| x_t = xt_list[j][k] |
| x_t = (x_t + 1) * 0.5 |
| utils.save_image(x_t, str(self.results_folder / f'x_{len(xt_list)-j}_{cnt}.png'), nrow=1) |
| x_t = cv2.imread(f'{self.results_folder}/x_{len(xt_list)-j}_{cnt}.png') |
| if j in to_show: |
| lst.append(x_t) |
|
|
| x_0 = x0_list[-1][k] |
| x_0 = (x_0 + 1) * 0.5 |
| utils.save_image(x_0, str(self.results_folder / f'x_best_{cnt}.png'), nrow=1) |
| x_0 = cv2.imread(f'{self.results_folder}/x_best_{cnt}.png') |
| lst.append(x_0) |
| im_h = cv2.hconcat(lst) |
| cv2.imwrite(f'{self.results_folder}/all_{cnt}.png', im_h) |
| cnt += 1 |
|
|
| def test_from_data_save_results(self): |
| batch_size = 100 |
| dl = data.DataLoader(self.ds, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=16, |
| drop_last=True) |
|
|
| all_samples = None |
|
|
| for i, img in enumerate(dl, 0): |
| print(i) |
| print(img.shape) |
| if all_samples is None: |
| all_samples = img |
| else: |
| all_samples = torch.cat((all_samples, img), dim=0) |
|
|
| |
|
|
| |
| blurred_samples = None |
| original_sample = None |
| deblurred_samples = None |
| direct_deblurred_samples = None |
|
|
| sanity_check = 1 |
|
|
| orig_folder = f'{self.results_folder}_orig/' |
| create_folder(orig_folder) |
|
|
| blur_folder = f'{self.results_folder}_blur/' |
| create_folder(blur_folder) |
|
|
| d_deblur_folder = f'{self.results_folder}_d_deblur/' |
| create_folder(d_deblur_folder) |
|
|
| deblur_folder = f'{self.results_folder}_deblur/' |
| create_folder(deblur_folder) |
|
|
| cnt = 0 |
| while cnt < all_samples.shape[0]: |
| print(cnt) |
| og_x = all_samples[cnt: cnt + 32] |
| og_x = og_x.cuda() |
| og_x = og_x.type(torch.cuda.FloatTensor) |
| og_img = og_x |
| x0_list, xt_list = self.ema_model.module.all_sample(batch_size=og_img.shape[0], faded_recon_sample=og_img, times=None) |
|
|
| og_img = og_img.to('cpu') |
| blurry_imgs = xt_list[0].to('cpu') |
| deblurry_imgs = x0_list[-1].to('cpu') |
| direct_deblurry_imgs = x0_list[0].to('cpu') |
|
|
| og_img = og_img.repeat(1, 3 // og_img.shape[1], 1, 1) |
| blurry_imgs = blurry_imgs.repeat(1, 3 // blurry_imgs.shape[1], 1, 1) |
| deblurry_imgs = deblurry_imgs.repeat(1, 3 // deblurry_imgs.shape[1], 1, 1) |
| direct_deblurry_imgs = direct_deblurry_imgs.repeat(1, 3 // direct_deblurry_imgs.shape[1], 1, 1) |
|
|
| og_img = (og_img + 1) * 0.5 |
| blurry_imgs = (blurry_imgs + 1) * 0.5 |
| deblurry_imgs = (deblurry_imgs + 1) * 0.5 |
| direct_deblurry_imgs = (direct_deblurry_imgs + 1) * 0.5 |
|
|
| if cnt == 0: |
| print(og_img.shape) |
| print(blurry_imgs.shape) |
| print(deblurry_imgs.shape) |
| print(direct_deblurry_imgs.shape) |
|
|
| if blurred_samples is None: |
| blurred_samples = blurry_imgs |
| else: |
| blurred_samples = torch.cat((blurred_samples, blurry_imgs), dim=0) |
|
|
| if original_sample is None: |
| original_sample = og_img |
| else: |
| original_sample = torch.cat((original_sample, og_img), dim=0) |
|
|
| if deblurred_samples is None: |
| deblurred_samples = deblurry_imgs |
| else: |
| deblurred_samples = torch.cat((deblurred_samples, deblurry_imgs), dim=0) |
|
|
| if direct_deblurred_samples is None: |
| direct_deblurred_samples = direct_deblurry_imgs |
| else: |
| direct_deblurred_samples = torch.cat((direct_deblurred_samples, direct_deblurry_imgs), dim=0) |
|
|
| cnt += og_img.shape[0] |
|
|
| print(blurred_samples.shape) |
| print(original_sample.shape) |
| print(deblurred_samples.shape) |
| print(direct_deblurred_samples.shape) |
|
|
| for i in range(blurred_samples.shape[0]): |
| utils.save_image(original_sample[i], f'{orig_folder}{i}.png', nrow=1) |
| utils.save_image(blurred_samples[i], f'{blur_folder}{i}.png', nrow=1) |
| utils.save_image(deblurred_samples[i], f'{deblur_folder}{i}.png', nrow=1) |
| utils.save_image(direct_deblurred_samples[i], f'{d_deblur_folder}{i}.png', nrow=1) |
|
|