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 # from einops import rearrange import os import errno from PIL import Image # from pytorch_msssim import ssim import cv2 import numpy as np import imageio from skimage.metrics import mean_squared_error, peak_signal_noise_ratio, structural_similarity # from torch.utils.tensorboard import SummaryWriter 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 structural_similarity as ssim 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() # try: # from apex import amp # # APEX_AVAILABLE = True # except: # APEX_AVAILABLE = False 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.backbone = diffusion_type.split('_')[0] 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 # Frequency Loss if self.backbone == 'twobranch' or self.backbone == 'twounet': self.amploss = AMPLoss() # .to(self.device, non_blocking=True) self.kl_loss = torch.nn.KLDivLoss(reduction='sum') # sum , log_target=True self.lpips = LPIPS().eval().cuda() # .to(self.device, non_blocking=True) self.ssim = StructuralSimilarityIndexMeasure() self.use_fre_loss = True self.use_ssim = False self.use_lpips = True # 141 -> 144s self.clamp_every_sample = False # Stride 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) # print("=== self.fade_kernels shape = ", self.fade_kernels.shape) # [5, 256, 256] elif self.degradation_type == "kspace": self.get_new_kspace(is_training=True) # print("=== self.kspace_kernels shape = ", self.kspace_kernels.shape) # [5, 256, 256] else: raise NotImplementedError() def get_new_kspace(self, is_training=False): # LinearSamplingRate, LogSamplingRate 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): # Test 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 # print("self.kspace_kernels = ", self.kspace_kernels.shape) # self.kspace_kernels = torch.Size([5, 320, 320]) # print("faded_recon_sample = ", faded_recon_sample.shape) # faded_recon_sample = torch.Size([16, 3, 320, 320]) # for i in range(t): with torch.no_grad(): k = torch.stack([self.kspace_kernels[[t - 1]]], 1) # print("k = ", k.shape) # k = torch.Size([1, 1, 320, 320]) 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 # print("faded_recon_sample = ", faded_recon_sample.shape) 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': # faded_recon_sample = recon_sample 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: # 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 = recon_sample #faded_recon_sample - recon_sample + recon_sample_sub_1 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() # last one 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) # print('k_full = ', k_full.shape) with torch.no_grad(): kt_sub_1 = self.get_kspace_kernels(t - 1).cuda() kt = self.get_kspace_kernels(t - 0).cuda() # last one 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 # * (1-k_residual) 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 #.cpu() 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() # last one 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 # new faded_recon_sample_fre = faded_recon_sample_fre * k_full + fre_amend # * (1-k_residual) # faded_recon_sample_fre = faded_recon_sample_fre * kt_sub_1 # + recon_sample * (1 - kt_sub_1) faded_recon_sample = apply_to_spatial(faded_recon_sample_fre, params_dict) k_known_mask += k_residual # .cpu() 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 # print("recon_sample = ", recon_sample.shape) 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): # TODO 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: # print(f"kspace randkeynel k={rand_kernels[:, i].shape}, x={x.shape}") k = torch.stack([rand_kernels[:, i]], 1) faded_recon_sample = apply_ksu_kernel(faded_recon_sample, k, params_dict) else: # print(f"kspace k={self.kspace_kernels[i].shape}, x={x.shape}") 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 #// 2 + recon_fre // 2 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': # faded_recon_sample = recon_sample 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 # Train 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) # self.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).""" # Create Gaussian kernel kernel = self.gaussian_kernel(kernel_size, sigma).unsqueeze(0).unsqueeze(0) kernel = kernel.expand(input_tensor.size(1), 1, kernel_size, kernel_size) # For each channel # Pad the input tensor to avoid size reduction padding = kernel_size // 2 input_tensor = F.pad(input_tensor, (padding, padding, padding, padding), mode='reflect') # Apply convolution 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") # Perform FFT and compute magnitudes 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): # Flatten the elements 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) # minus max value 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]) # 4B pred_all = F.log_softmax(pred_all, dim=1) # consistency_loss # get probability 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) # 2 B 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 # + kl_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)) # negative mask 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) # print("mag loss=", mag_kl_loss, "pha loss=", pha_kl_loss) 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) # gaussian blur for x_mix # if np.random.rand() > 0.5: # x_mix = self.gaussian_blur( # x_mix, # kernel_size=int(torch.randint(1, 9, (1,)).item() * 2 + 1), # Ensure odd kernel size # sigma=torch.abs(torch.randn(1) * 3.0).item() # Ensure sigma is positive # ) # Add gaussian noise # sigma = 0.1 * torch.abs(torch.rand(1)).item() # Standard Deviation # x_mix = x_mix + torch.randn_like(x_mix) * sigma # aux = aux + torch.randn_like(x_mix) * sigma 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) # 0.02s ~ 0.03s 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: # NAN 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 # LPIPS 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: # NAN 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) # img_mean = params_dict['img_mean'].cuda().view(-1, 1, 1, 1) # img_std = params_dict['img_std'].cuda().view(-1, 1, 1, 1) # x_start_golden = x_start_golden * img_std + img_mean # 0 - 1 # x_recon = x_recon * img_std + img_mean # x_recon_fre = x_recon_fre * img_std + img_mean 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) # LPIPS 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: # NAN fft_weight = 0.01 fre_loss = fft_weight * self.amploss(x_recon_fre, x_start_golden, k) loss += fre_loss # fre_loss = fft_weight * self.amploss(x_recon, 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(), # "lpips_loss:", lpips_loss.item(), # "ssim_loss", ssim_loss.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) # mode, base_dir, image_size, nclass, domains, aux_modality, self.ds = BrainDataset(mode, folder, image_size, 4, debug=debug, domains=domain, num_channels=num_channels, aux_modality=aux_modality) # mode, base_dir, domains: 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(), AddNoise()]), 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(), AddNoise()]), 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 = build_dataset(args, mode='train') 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 # assert not fp16 or fp16 and APEX_AVAILABLE, 'Apex must be installed for mixed precision training on' 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() #.permute(0, 2, 3, 1).numpy()[..., 0] og_img_ = og_img.cpu() #.permute(0, 2, 3, 1).numpy()[..., 0] # print("img_=", img_.shape, "og_img_=", og_img_.shape) # img_= torch.Size([4, 1, 240, 240]) og_img_= torch.Size([4, 1, 240, 240]) # B, C, H, W # ssim = StructuralSimilarityIndexMeasure(data_range=255) ssims_ = [] for (img, og_img) in zip(img_, og_img_): img_np = img.squeeze().numpy() # Convert to 2D og_img_np = og_img.squeeze().numpy() # Convert to 2D # Compute SSIM for each pair of images 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_ # pip install pytorch-fid # (N, H, W, 3), dtype: np.float32 between 0-1 or np.uint8 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) # img_ = torch.Size([5, 1, 64, 64]) torch.Size([5, 1, 64, 64] 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: # (N, 3, 299, 299) fid_value = calculate_fid(all_images_new.cpu().numpy(), og_img_new.cpu().numpy(), use_multiprocessing=False, batch_size=og_img.shape[-1]) # (N, 3, C, 256, 256) 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 # B, C, H, W lpips = self.lpips(all_images_new.cuda(), og_img_new.cuda()).mean().item() # H, W, C img_ = all_images.cpu().unsqueeze(0) # .permute(1, 2, 0).numpy() og_img_ = og_img.cpu().unsqueeze(0) #.permute(1, 2, 0).numpy() #.numpy() # 0-1, H, W, C 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 # Evaluate all 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} # print("input shape minmax, og_img:", og_img.min().item(), og_img.max().item(), # "aux:", aux.min().item(), aux.max().item()) 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()) # print("----------------------------------------\n" # "all_recons:", all_recons.min().item(), all_recons.max().item(), # "all_images:", all_images.min().item(), all_images.max().item(), # "og_img:", og_img.min().item(), og_img.max().item(), # "direct_recons:", direct_recons.min().item(), direct_recons.max().item()) 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) # 24, 1, 128, 128 # Calculate SSIM and PSNR, LPIPS 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) # utils.save_image(xt, str(self.results_folder / f'{self.step}-xt-Noise.png'), nrow=6) # utils.save_image(all_images, str(self.results_folder / f'{self.step}-full_recons.png'), # nrow=6) # utils.save_image(direct_recons, # str(self.results_folder / f'{self.step}-sample-direct_recons.png'), nrow=6) # utils.save_image(og_img, str(self.results_folder / f'{self.step}-img.png'), nrow=6) # utils.save_image(aux, str(self.results_folder / f'{self.step}-aux.png'), nrow=6) # plot ssim and psnr in two sub plots same canvas parallel 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() # PSNR plot 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_recon = all_recons[:, 0] # 50, 1, 128, 128 # Ensure all_recons is on the CPU all_recons = torch.cat(list(all_recons), dim=-1).cpu() 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] # Calculate repeat factor # tensor_small = tensor_small.repeat(1, 1, 1, repeats) 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()) # before and after residual 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) # utils.save_image(all_masks, str(self.results_folder / f'{self.step}-all_masks-{routine}.png'), # nrow=1) # acc_loss = acc_loss / (self.save_and_sample_every + 1) print(f'Mean of last {self.step}: save to :', str(self.results_folder / f'{self.step}-combine.png')) # acc_loss = 0 def train(self): backwards = partial(loss_backwards, self.fp16) # writer = SummaryWriter() 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) # Very slow, 0.5 s second on bask, 0.001 on local 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) # Revert model self.opt.load_state_dict(optimizer_state) # Revert optimizer continue # Skip the rest of this training step if self.debug_time: print("before loss=", time.time() - d_time) u_loss += loss loss.backward() # backwards(loss / self.gradient_accumulate_every, self.opt) 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())) # writer.add_scalar("Loss/train", loss.item(), self.step) acc_loss = acc_loss + (u_loss.item() / self.gradient_accumulate_every) max_norm = 0.01 # Maximum norm for gradients 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() # TEST and SAVE if self.step != 0 and (self.step + 1) % self.save_and_sample_every == 0: batches = self.batch_size data_dict = next(self.test_dl) # .cuda() train_dict = next(self.dl) # .cuda() # 'default', "fre_progressive", "x0_step_down" 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 # self.ema_model # or self.model model.eval() # Set model to evaluation mode # sampling_routine = ['default', 'x0_step_down', 'x0_step_down_fre', "fre_progressive"]: num_timesteps = model.module.num_timesteps patient_wise_pred = {} # num_timesteps patient_wise_gt = [] count = 1 while True: batches = 1 #self.batch_size data_dict = next(self.dl) # .cuda() # Prediction 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()) # into a Normalized Image 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() # og_img = (og_img - _min) / (_max - _min) # 0 - 1 all_images = all_images * img_std + img_mean # all_images = (all_images - _min) / (_max - _min) all_images_norm = all_images_norm * img_std + img_mean # all_images_norm = (all_images_norm - _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) direct_recons_norm = direct_recons_norm * img_std + img_mean # direct_recons_norm = (direct_recons_norm - _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() 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) # 24, 1, 128, 128 # Calculate SSIM and PSNR, LPIPS 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) # utils.save_image(xt, str(self.results_folder / f'{self.step}-xt-Noise.png'), nrow=6) # utils.save_image(all_images, str(self.results_folder / f'{self.step}-full_recons.png'), # nrow=6) # utils.save_image(direct_recons, # str(self.results_folder / f'{self.step}-sample-direct_recons.png'), nrow=6) # utils.save_image(og_img, str(self.results_folder / f'{self.step}-img.png'), nrow=6) # utils.save_image(aux, str(self.results_folder / f'{self.step}-aux.png'), nrow=6) # plot ssim and psnr in two sub plots same canvas parallel 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() # PSNR plot 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) # utils.save_image(combine, str(self.results_folder / f'{self.step}-combine-{routine}.png'), nrow=6) # all_recon = all_recons[:, 0] # 50, 1, 128, 128 # Ensure all_recons is on the CPU 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] # Calculate repeat factor # tensor_small = tensor_small.repeat(1, 1, 1, repeats) og_img = og_img.cpu() all_recons_residual = all_recons - og_img.repeat(1, 1, 1, repeats) # all_recons[:, :, :, s:] 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) # utils.save_image(all_recons, str(self.results_folder / f'{self.step}-all_recons-{routine}.png'), # nrow=1) 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 # self.ema_model # or self.model model.eval() # Set model to evaluation mode # sampling_routine = ['default', 'x0_step_down', 'x0_step_down_fre', "fre_progressive"]: num_timesteps = model.module.num_timesteps patient_wise_pred = {} # num_timesteps patient_wise_gt = [] while True: batches = 1 #self.batch_size data_dict = next(self.dl) # .cuda() if data_dict['is_start']: patient_wise_pred = {} # num_timesteps of list for i in range(num_timesteps): patient_wise_pred[i] = [] patient_wise_gt = [] # Original Input 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} # Prediction 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)) # Handling the output 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() # Save to list 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']: # or len(patient_wise_gt) >=10: # TODO # 24, 1, 128, 128 # Calculate SSIM and PSNR, LPIPS patient_gt = torch.cat(patient_wise_gt, dim=0).squeeze() # C, H, W 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() # PSNR plot 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() # LPIPS plot 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() # FID plot 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) # break # create_folder(f'{self.results_folder}/') 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)