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