OSCR

Dual-conditioned diffusion model with anatomical guidance for geometric distortion correction in prostate MRI.

Code ↔ Paper

5 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 5 matches
  1. [1] § Methods › DeDistortNet model architecture and training ↔ source/train_DeDistortNet.py, lines 958–996 · score 0.71 · forward diffusion process, noisy latent, predicted, timestep, noise, training
  2. [2] § Methods › Evaluation metrics ↔ source/train_DeDistortNet.py, lines 1–51 · score 0.58 · Peak Signal, Noise Ratio, PSNR, SSIM
  3. [3] § Methods › DeDistortNet model architecture and training ↔ source/train_DeDistortNet.py, lines 53–137 · score 0.56 · CLIP image encoder, Stable Diffusion, prompts, pretrained, rotation, weights
  4. [4] § Methods › Comparison models ↔ source/train_DeDistortNet.py, lines 53–137 · score 0.54 · CLIP image encoder, Stable Diffusion, prompt, guidance, models, training
  5. [5] § Methods › DeDistortNet model architecture and training ↔ source/train_DeDistortNet.py, lines 787–866 · score 0.50 · AdamW, optimizer, channels, batch, Training, weighted

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,060 lines · 44 KB · Apache-2.0 · 5 matches

  1. import os
  2. import random
  3. import argparse
  4. from pathlib import Path
  5. import json
  6. import itertools
  7. import time
  8. import logging
  9. import math
  10. import elasticdeform
  11. import accelerate
  12. import numpy as np
  13. from tqdm.auto import tqdm
  14. import matplotlib.pyplot as plt
  15. import torch
  16. import torch.nn.functional as F
  17. from torchvision import transforms
  18. import SimpleITK as sitk
  19. from transformers import CLIPImageProcessor
  20. from accelerate import Accelerator
  21. from accelerate.logging import get_logger
  22. from accelerate.utils import ProjectConfiguration, set_seed
  23. from packaging import version
  24. from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel, DDIMScheduler
  25. from diffusers import ControlNetModel
  26. from diffusers.utils.torch_utils import is_compiled_module
  27. from transformers import CLIPVisionModelWithProjection
  28. from diffusers import StableDiffusionControlNetPipeline
  29. from torchmetrics.image import StructuralSimilarityIndexMeasure
  30. from torchmetrics.image import PeakSignalNoiseRatio
  31. from scipy.ndimage import zoom
  32. from skimage import transform
  33. PSNR = PeakSignalNoiseRatio(data_range=1.0)
  34. SSIM = StructuralSimilarityIndexMeasure(data_range=1.0)
  35. print("PyTorch version:", torch.__version__)
  36. print("Is CUDA available:", torch.cuda.is_available())
  37. print("cuDNN version:", torch.backends.cudnn.version())
  38. print("Is cuDNN enabled:", torch.backends.cudnn.enabled)
  39. import wandb
  40. logger = get_logger(__name__)
  41. def log_validation(
  42. args, accelerator, weight_dtype, step, checkpoint_path
  43. ):
  44. weight_dtype = torch.float16
  45. logger.info("Running validation... ")
  46. val_dataset = MyDataset(args.val_data_json_file, size=args.resolution, displacement_rate=args.displacement_rate, random_y_squeeze_rate=args.random_y_squeeze_rate, random_rotation_degree=args.random_rotation_degree, random_shift_range=args.random_shift_range, image_root_path=args.data_root_path, clip_image_processor_path=args.image_encoder_path)
  47. val_dataloader = torch.utils.data.DataLoader(
  48. val_dataset,
  49. shuffle=False,
  50. collate_fn=collate_fn,
  51. batch_size=args.train_batch_size,
  52. num_workers=args.dataloader_num_workers,
  53. pin_memory=True
  54. )
  55. controlnet = ControlNetModel.from_pretrained(os.path.join(checkpoint_path, "controlnet"))
  56. controlnet = controlnet.to(accelerator.device, dtype=weight_dtype)
  57. vae = AutoencoderKL.from_pretrained(args.pretrained_model_name_or_path, subfolder="vae").to(accelerator.device, dtype=weight_dtype)
  58. noise_scheduler = DDIMScheduler(
  59. num_train_timesteps=1000,
  60. beta_start=0.00085,
  61. beta_end=0.012,
  62. beta_schedule="scaled_linear",
  63. clip_sample=False,
  64. set_alpha_to_one=False,
  65. steps_offset=1,
  66. )
  67. image_encoder = CLIPVisionModelWithProjection.from_pretrained(os.path.join(checkpoint_path, 'image_encoder')).to(accelerator.device, dtype=weight_dtype)
  68. unet = UNet2DConditionModel.from_pretrained(checkpoint_path, subfolder="unet").to(accelerator.device, dtype=weight_dtype)
  69. pipeline = StableDiffusionControlNetPipeline.from_pretrained(
  70. args.pretrained_model_name_or_path,
  71. vae=vae,
  72. unet=unet,
  73. controlnet=controlnet,
  74. safety_checker=None,
  75. feature_extractor=None,
  76. torch_dtype=weight_dtype,
  77. scheduler=noise_scheduler,
  78. )
  79. pipeline = pipeline.to(accelerator.device, dtype=weight_dtype)
  80. pipeline.set_progress_bar_config(disable=True)
  81. if args.seed is None:
  82. generator = None
  83. else:
  84. generator = torch.Generator(device=accelerator.device).manual_seed(args.seed)
  85. val_psnr = 0.0
  86. val_ssim = 0.0
  87. val_psnr_cropped = 0.0
  88. val_ssim_cropped = 0.0
  89. for v_i, val_example in enumerate(tqdm(val_dataloader)):
  90. image = val_example["images"]
  91. control_image = val_example["control_images"]
  92. clip_image = val_example["clip_images"].to(accelerator.device, dtype=weight_dtype)
  93. with torch.no_grad():
  94. encoder_hidden_states = image_encoder(clip_image).last_hidden_state
  95. with torch.autocast("cuda"):
  96. result_image = pipeline(
  97. prompt_embeds=encoder_hidden_states,
  98. image=control_image,
  99. num_inference_steps=100,
  100. guidance_scale=0.0,
  101. generator=generator,
  102. output_type="np",
  103. ).images
  104. image = image * 0.5 + 0.5
  105. image = image.cpu().float()
  106. control_image = control_image * 0.5 + 0.5
  107. control_image = control_image.cpu().float()
  108. result_image = np.transpose(result_image, (0, 3, 1, 2))
  109. result_image = torch.from_numpy(result_image).float()
  110. val_psnr += PSNR(result_image, image)
  111. val_ssim += SSIM(result_image, image)
  112. for i in range(image.shape[0]):
  113. # Find non-zero pixel indices in the first channel
  114. non_zero_indices = torch.nonzero(control_image[i, 0])
  115. # Get bounding box of non-zero region
  116. min_row, min_col = torch.min(non_zero_indices, dim=0)[0]
  117. max_row, max_col = torch.max(non_zero_indices, dim=0)[0]
  118. result_image_cropped = result_image[i:i+1].clone()
  119. result_image_cropped = result_image_cropped[:, :, min_row:max_row+1, min_col:max_col+1]
  120. image_cropped = image[i:i+1, :, min_row:max_row+1, min_col:max_col+1]
  121. val_psnr_cropped += PSNR(result_image_cropped, image_cropped)
  122. val_ssim_cropped += SSIM(result_image_cropped, image_cropped)
  123. val_psnr /= len(val_dataloader)
  124. val_ssim /= len(val_dataloader)
  125. val_psnr_cropped /= len(val_dataset)
  126. val_ssim_cropped /= len(val_dataset)
  127. logger.info(f"Validation PSNR: {val_psnr}, Cropped: {val_psnr_cropped}")
  128. logger.info(f"Validation SSIM: {val_ssim}, Cropped: {val_ssim_cropped}")
  129. wandb.log({"val_psnr": val_psnr, "val_ssim": val_ssim, "step": step, "cropped_val_psnr": val_psnr_cropped, "cropped_val_ssim": val_ssim_cropped})
  130. logger.info("Run plot validation...")
  131. val_dataset = ValDataset(args.plot_data_json_file, size=args.resolution, image_root_path=args.data_root_path, clip_image_processor_path=args.image_encoder_path)
  132. fig, axs = plt.subplots(len(val_dataset), 8, figsize=(40, 5 * len(val_dataset)))
  133. for v_i, val_example in enumerate(tqdm(val_dataset)):
  134. image = val_example["image"]
  135. control_image = val_example["control_image"]
  136. gt_seg_image = val_example["gt_seg_image"]
  137. control_image = control_image.unsqueeze(0)
  138. clip_image = val_example["clip_image"].to(accelerator.device, dtype=weight_dtype)
  139. with torch.no_grad():
  140. encoder_hidden_states = image_encoder(clip_image).last_hidden_state
  141. with torch.autocast("cuda"):
  142. result_image = pipeline(
  143. prompt_embeds=encoder_hidden_states,
  144. image=control_image,
  145. num_inference_steps=100,
  146. guidance_scale=0.0,
  147. generator=generator,
  148. output_type="np",
  149. ).images[0]
  150. channel_names = val_example["channel_names"]
  151. # Plot context images and generated images for each channel
  152. for ch in range(3):
  153. vmax = max(image[:,:,ch].max(), result_image[:,:,ch].max())
  154. axs[v_i, ch].imshow(image[:,:,ch], cmap="gray", vmin=0, vmax=vmax)
  155. axs[v_i, ch].set_title("Context Image " + channel_names[ch])
  156. axs[v_i, ch].axis("off")
  157. axs[v_i, ch+4].imshow(result_image[:,:,ch], cmap="gray", vmin=0, vmax=vmax)
  158. axs[v_i, ch+4].set_title("Generated Image " + channel_names[ch])
  159. axs[v_i, ch+4].axis("off")
  160. axs[v_i, 3].imshow(control_image[0,0], cmap="gray", vmin=-1, vmax=1)
  161. axs[v_i, 3].set_title("Structure Image")
  162. axs[v_i, 3].axis("off")
  163. axs[v_i, 7].imshow(gt_seg_image, cmap="viridis", vmin=0, vmax=3)
  164. axs[v_i, 7].set_title("Ground Truth Segmentation")
  165. axs[v_i, 7].axis("off")
  166. os.makedirs(os.path.join(args.output_dir, "validation_plots"), exist_ok=True)
  167. plt.tight_layout()
  168. plt.savefig(os.path.join(args.output_dir, "validation_plots", f"validation_{step}.png"), dpi=300)
  169. plt.close()
  170. del pipeline
  171. del controlnet
  172. return val_psnr, val_ssim, val_psnr_cropped, val_ssim_cropped
  173. class ValDataset(torch.utils.data.Dataset):
  174. def __init__(self, json_file, size=512, image_root_path="", clip_image_processor_path=None):
  175. super().__init__()
  176. self.size = size
  177. self.image_root_path = image_root_path
  178. train_min_max_file = os.path.join(self.image_root_path, "train_min_max.json")
  179. with open(train_min_max_file, "r") as f:
  180. train_min_max = json.load(f)
  181. self.b50_min = train_min_max["B50"]["prostate_min"]
  182. self.b50_max = train_min_max["B50"]["prostate_max"]
  183. self.b400_min = train_min_max["B400"]["prostate_min"]
  184. self.b400_max = train_min_max["B400"]["prostate_max"]
  185. self.b800_min = train_min_max["B800"]["prostate_min"]
  186. self.b800_max = train_min_max["B800"]["prostate_max"]
  187. self.t2_min = train_min_max["T2"]["prostate_min"]
  188. self.t2_max = train_min_max["T2"]["prostate_max"]
  189. self.data = []
  190. with open(json_file, "r") as f:
  191. for line in f:
  192. self.data.append(json.loads(line))
  193. self.transform = transforms.Compose([
  194. transforms.ToTensor(),
  195. transforms.Normalize([0.5], [0.5]),
  196. ])
  197. if clip_image_processor_path is None:
  198. self.clip_image_processor = CLIPImageProcessor()
  199. else:
  200. self.clip_image_processor = CLIPImageProcessor.from_pretrained(clip_image_processor_path)
  201. def __getitem__(self, idx):
  202. item = self.data[idx]
  203. image_file = item["image"]
  204. crop_coor = item["crop_coor"]
  205. b50_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'B50', image_file)))
  206. b400_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'B400', image_file)))
  207. b800_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'B800', image_file)))
  208. control_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'T2', image_file))).astype(np.float32)
  209. gt_seg_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'PZTZ', image_file)))
  210. # global min-max normalization
  211. b50_image = (b50_image - self.b50_min) / (self.b50_max - self.b50_min)
  212. b400_image = (b400_image - self.b400_min) / (self.b400_max - self.b400_min)
  213. b800_image = (b800_image - self.b800_min) / (self.b800_max - self.b800_min)
  214. # clip 0.0 ~ 1.0
  215. b50_image = np.clip(b50_image, 0, 1)
  216. b400_image = np.clip(b400_image, 0, 1)
  217. b800_image = np.clip(b800_image, 0, 1)
  218. image = np.stack([b50_image, b400_image, b800_image], axis=-1)
  219. channel_names = ["B50", "B400", "B800"]
  220. image = image[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  221. control_image = control_image[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  222. gt_seg_image = gt_seg_image[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  223. # Define target size
  224. target_size = (self.size, self.size)
  225. # Create an empty array of zeros with the target size
  226. image_padded = np.zeros((target_size[0], target_size[1], image.shape[2]), dtype=np.float32)
  227. control_image_padded = np.zeros((target_size[0], target_size[1]), dtype=np.float32)
  228. gt_seg_image_padded = np.zeros((target_size[0], target_size[1]), dtype=np.int32)
  229. # Calculate padding offsets
  230. y_offset = (target_size[0] - image.shape[0]) // 2
  231. x_offset = (target_size[1] - image.shape[1]) // 2
  232. # Place the cropped image in the center of the padded image
  233. image_padded[y_offset:y_offset+image.shape[0], x_offset:x_offset+image.shape[1], :] = image
  234. control_image_padded[y_offset:y_offset+control_image.shape[0], x_offset:x_offset+control_image.shape[1]] = control_image
  235. gt_seg_image_padded[y_offset:y_offset+gt_seg_image.shape[0], x_offset:x_offset+gt_seg_image.shape[1]] = gt_seg_image
  236. # Apply min-max normalization to each channel
  237. control_image = (control_image_padded - control_image_padded.min()) / (control_image_padded.max() - control_image_padded.min())
  238. image = image_padded
  239. control_image = control_image[..., np.newaxis]
  240. control_image = self.transform(control_image)
  241. clip_image = self.clip_image_processor(images=image, return_tensors="pt", do_rescale=False).pixel_values
  242. return {
  243. "channel_names": channel_names,
  244. "image": image,
  245. "control_image": control_image,
  246. "clip_image": clip_image,
  247. "gt_seg_image": gt_seg_image_padded
  248. }
  249. def __len__(self):
  250. return len(self.data)
  251. # Dataset
  252. class MyDataset(torch.utils.data.Dataset):
  253. def __init__(self, json_file, size=512, displacement_rate=32, random_y_squeeze_rate=0.05, random_rotation_degree=5, random_shift_range=2, i_drop_rate=0.0, image_root_path="", clip_image_processor_path=None, isTrain=False):
  254. super().__init__()
  255. self.size = size
  256. self.displacement_rate = displacement_rate
  257. self.random_y_squeeze_rate = random_y_squeeze_rate
  258. self.random_rotation_degree = random_rotation_degree
  259. self.random_shift_range = random_shift_range
  260. self.i_drop_rate = i_drop_rate
  261. self.image_root_path = image_root_path
  262. self.isTrain = isTrain
  263. train_min_max_file = os.path.join(self.image_root_path, "train_min_max.json")
  264. with open(train_min_max_file, "r") as f:
  265. train_min_max = json.load(f)
  266. self.b50_min = train_min_max["B50"]["prostate_min"]
  267. self.b50_max = train_min_max["B50"]["prostate_max"]
  268. self.b400_min = train_min_max["B400"]["prostate_min"]
  269. self.b400_max = train_min_max["B400"]["prostate_max"]
  270. self.b800_min = train_min_max["B800"]["prostate_min"]
  271. self.b800_max = train_min_max["B800"]["prostate_max"]
  272. self.t2_min = train_min_max["T2"]["prostate_min"]
  273. self.t2_max = train_min_max["T2"]["prostate_max"]
  274. self.data = []
  275. with open(json_file, "r") as f:
  276. for line in f:
  277. self.data.append(json.loads(line))
  278. self.transform = transforms.Compose([
  279. transforms.ToTensor(),
  280. transforms.Normalize([0.5], [0.5]),
  281. ])
  282. if clip_image_processor_path is None:
  283. self.clip_image_processor = CLIPImageProcessor()
  284. else:
  285. self.clip_image_processor = CLIPImageProcessor.from_pretrained(clip_image_processor_path)
  286. def __getitem__(self, idx):
  287. item = self.data[idx]
  288. image_file = item["image"]
  289. distortion_pivot = item["distortion_pivot"]
  290. crop_coor = item["crop_coor"]
  291. # read image
  292. raw_b50_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'B50', image_file))).astype(np.float32)
  293. raw_b400_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'B400', image_file))).astype(np.float32)
  294. raw_b800_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'B800', image_file))).astype(np.float32)
  295. control_raw_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'T2', image_file))).astype(np.float32)
  296. gt_seg_image = sitk.GetArrayFromImage(sitk.ReadImage(os.path.join(self.image_root_path, 'PZTZ', image_file)))
  297. # global min-max normalization
  298. raw_b50_image = (raw_b50_image - self.b50_min) / (self.b50_max - self.b50_min)
  299. raw_b400_image = (raw_b400_image - self.b400_min) / (self.b400_max - self.b400_min)
  300. raw_b800_image = (raw_b800_image - self.b800_min) / (self.b800_max - self.b800_min)
  301. # clip 0.0 ~ 1.0
  302. raw_b50_image = np.clip(raw_b50_image, 0, 1)
  303. raw_b400_image = np.clip(raw_b400_image, 0, 1)
  304. raw_b800_image = np.clip(raw_b800_image, 0, 1)
  305. image = np.stack([raw_b50_image, raw_b400_image, raw_b800_image], axis=-1)
  306. image = image[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  307. control_image = control_raw_image[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  308. gt_seg_image = gt_seg_image[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  309. # Define target size
  310. target_size = (self.size, self.size)
  311. # Create an empty array of zeros with the target size
  312. image_padded = np.zeros((target_size[0], target_size[1], image.shape[2]), dtype=np.float32)
  313. control_image_padded = np.zeros((target_size[0], target_size[1]), dtype=np.float32)
  314. gt_seg_image_padded = np.zeros((target_size[0], target_size[1]), dtype=np.int32)
  315. # Calculate padding offsets
  316. y_offset = (target_size[0] - image.shape[0]) // 2
  317. x_offset = (target_size[1] - image.shape[1]) // 2
  318. # Place the cropped image in the center of the padded image
  319. image_padded[y_offset:y_offset+image.shape[0], x_offset:x_offset+image.shape[1], :] = image
  320. control_image_padded[y_offset:y_offset+control_image.shape[0], x_offset:x_offset+control_image.shape[1]] = control_image
  321. gt_seg_image_padded[y_offset:y_offset+gt_seg_image.shape[0], x_offset:x_offset+gt_seg_image.shape[1]] = gt_seg_image
  322. control_image = (control_image_padded - control_image_padded.min()) / (control_image_padded.max() - control_image_padded.min())
  323. image = image_padded
  324. control_image = control_image[..., np.newaxis]
  325. image = self.transform(image)
  326. control_image = self.transform(control_image)
  327. gt_seg_image = torch.from_numpy(gt_seg_image_padded)
  328. if np.random.rand() > 0.5:
  329. displacement = self.get_random_displacement(distortion_pivot, raw_b50_image)
  330. b50_deformed = elasticdeform.deform_grid(raw_b50_image, displacement, order=3)
  331. b400_deformed = elasticdeform.deform_grid(raw_b400_image, displacement, order=3)
  332. b800_deformed = elasticdeform.deform_grid(raw_b800_image, displacement, order=3)
  333. upscaled_displacement = self.upscale_displacement(displacement, raw_b50_image.shape)
  334. jacobian_det = self.compute_jacobian(upscaled_displacement)
  335. b50_adjusted = self.adjust_pixel_values(raw_b50_image, b50_deformed, jacobian_det)
  336. b400_adjusted = self.adjust_pixel_values(raw_b400_image, b400_deformed, jacobian_det)
  337. b800_adjusted = self.adjust_pixel_values(raw_b800_image, b800_deformed, jacobian_det)
  338. image_deformed = np.stack([b50_adjusted, b400_adjusted, b800_adjusted], axis=-1)
  339. image_deformed = np.clip(image_deformed, 0, 1)
  340. image_deformed = self.random_squeeze_y_axis(image_deformed)
  341. image_deformed = self.random_rotation(image_deformed)
  342. x_shift = random.randint(-self.random_shift_range, self.random_shift_range)
  343. y_shift = random.randint(-self.random_shift_range, self.random_shift_range)
  344. image_deformed = np.roll(image_deformed, (x_shift, y_shift), axis=(1, 0))
  345. else:
  346. image_deformed = np.stack([raw_b50_image, raw_b400_image, raw_b800_image], axis=-1)
  347. image_deformed = image_deformed[crop_coor['y_start']:crop_coor['y_start']+crop_coor['y_size'], crop_coor['x_start']:crop_coor['x_start']+crop_coor['x_size']]
  348. image_deformed_padded = np.zeros((target_size[0], target_size[1], image_deformed.shape[2]), dtype=image_deformed.dtype)
  349. image_deformed_padded[y_offset:y_offset+image_deformed.shape[0], x_offset:x_offset+image_deformed.shape[1], :] = image_deformed
  350. image_deformed = image_deformed_padded
  351. clip_image = self.clip_image_processor(images=image_deformed, return_tensors="pt", do_rescale=False).pixel_values
  352. return {
  353. "image": image,
  354. "clip_image": clip_image,
  355. "control_image": control_image,
  356. "gt_seg_image": gt_seg_image,
  357. }
  358. def __len__(self):
  359. return len(self.data)
  360. def downsample_point(self, point, arr):
  361. original_size = np.array(arr.shape)
  362. target_size = original_size // self.displacement_rate
  363. return [int(round(p * t / o)) for p, o, t in zip(point, original_size, target_size)]
  364. def get_random_displacement(self, points, arr):
  365. points_downsampled = []
  366. for point in points:
  367. points_downsampled.append(self.downsample_point(point, arr))
  368. points_downsampled = np.array(points_downsampled)
  369. center_point = points_downsampled[4]
  370. min_x = points_downsampled[0][1]
  371. max_x = points_downsampled[2][1]
  372. min_y = points_downsampled[0][0]
  373. max_y = points_downsampled[6][0]
  374. displacement_size_y = arr.shape[0] // self.displacement_rate
  375. displacement_size_x = arr.shape[1] // self.displacement_rate
  376. displacement = np.zeros((2, displacement_size_y, displacement_size_x))
  377. for i in range(displacement_size_y):
  378. for j in range(displacement_size_x):
  379. if (i,j) == (center_point[0], center_point[1]):
  380. displacement[0, i, j] = np.random.randn() * 10 # randn: expansion+compression, rand: expansion only
  381. displacement[1, i, j] = np.random.randn() * 10
  382. elif (min_y < i < max_y) and (min_x < j < max_x):
  383. displacement[0, i, j] = np.random.randn() * 10
  384. displacement[1, i, j] = np.random.randn() * 10
  385. return displacement
  386. def upscale_displacement(self, disp, target_shape, order=3):
  387. factor = target_shape[0] / disp.shape[1] # assumes square image and displacement grid
  388. upscaled_disp = np.zeros((2, *target_shape))
  389. upscaled_disp[0] = zoom(disp[0], factor, order=order)
  390. upscaled_disp[1] = zoom(disp[1], factor, order=order)
  391. return upscaled_disp
  392. def compute_jacobian(self, disp):
  393. # Compute spatial gradients of displacement field
  394. grad_y_x = np.gradient(disp[0], axis=1) # ∂(y-disp)/∂x
  395. grad_y_y = np.gradient(disp[0], axis=0) # ∂(y-disp)/∂y
  396. grad_x_x = np.gradient(disp[1], axis=1) # ∂(x-disp)/∂x
  397. grad_x_y = np.gradient(disp[1], axis=0) # ∂(x-disp)/∂y
  398. # Jacobian determinant: J = (1 + ∂y/∂y)(1 + ∂x/∂x) - (∂y/∂x)(∂x/∂y)
  399. jacobian = (1 + grad_y_y) * (1 + grad_x_x) - (grad_y_x * grad_x_y)
  400. return jacobian
  401. def adjust_pixel_values(self, original_image, deformed_image, jacobian_det):
  402. return deformed_image * jacobian_det
  403. def random_squeeze_y_axis(self, image):
  404. """
  405. Randomly scales the image along the Y-axis while keeping the X and Z axes unchanged.
  406. Parameters:
  407. image (numpy.ndarray): 3D image array with shape (height, width, channels).
  408. scale_range (tuple): Tuple with min and max scale factors for the Y-axis.
  409. Returns:
  410. numpy.ndarray: Scaled image.
  411. """
  412. # Ensure image is a 3D numpy array
  413. if image.ndim != 3:
  414. raise ValueError("Input image must be a 3D numpy array.")
  415. # Randomly choose a scale factor for the Y-axis
  416. scale_y = np.random.uniform(1-self.random_y_squeeze_rate, 1+self.random_y_squeeze_rate)
  417. # Define scale factors for X, Y, Z axes
  418. scale_factors = [scale_y, 1, 1]
  419. # Rescale the image
  420. image_rescaled = transform.rescale(image, scale_factors, mode='reflect', anti_aliasing=True)
  421. return image_rescaled
  422. def random_rotation(self, image):
  423. # Randomly choose an angle for rotation
  424. max_angle = self.random_rotation_degree
  425. angle = np.random.uniform(-max_angle, max_angle)
  426. # Rotate the image
  427. image_rotated = transform.rotate(image, angle, mode='reflect')
  428. return image_rotated
  429. def collate_fn(data):
  430. images = torch.stack([example["image"] for example in data])
  431. clip_images = torch.cat([example["clip_image"] for example in data], dim=0)
  432. control_images = torch.stack([example["control_image"] for example in data])
  433. gt_seg_images = torch.stack([example["gt_seg_image"] for example in data])
  434. return {
  435. "images": images,
  436. "clip_images": clip_images,
  437. "control_images": control_images,
  438. "gt_seg_images": gt_seg_images
  439. }
  440. def parse_args():
  441. parser = argparse.ArgumentParser(description="Simple example of a training script.")
  442. parser.add_argument(
  443. "--pretrained_model_name_or_path",
  444. type=str,
  445. default=None,
  446. required=True,
  447. help="Path to pretrained model or model identifier from huggingface.co/models.",
  448. )
  449. parser.add_argument(
  450. "--data_json_file",
  451. type=str,
  452. default=None,
  453. required=True,
  454. help="Training data",
  455. )
  456. parser.add_argument(
  457. "--val_data_json_file",
  458. type=str,
  459. default=None,
  460. required=True,
  461. help="Validation data",
  462. )
  463. parser.add_argument(
  464. "--plot_data_json_file",
  465. type=str,
  466. default=None,
  467. required=True,
  468. help="Plot data",
  469. )
  470. parser.add_argument(
  471. "--data_root_path",
  472. type=str,
  473. default="",
  474. required=True,
  475. help="Training data root path",
  476. )
  477. parser.add_argument(
  478. "--image_encoder_path",
  479. type=str,
  480. default=None,
  481. required=True,
  482. help="Path to CLIP image encoder",
  483. )
  484. parser.add_argument(
  485. "--output_dir",
  486. type=str,
  487. default="sd-ipadapter-dec+controlnet",
  488. help="The output directory where the model predictions and checkpoints will be written.",
  489. )
  490. parser.add_argument(
  491. "--logging_dir",
  492. type=str,
  493. default="logs",
  494. help=(
  495. "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
  496. " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
  497. ),
  498. )
  499. parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
  500. parser.add_argument(
  501. "--resolution",
  502. type=int,
  503. default=512,
  504. help=(
  505. "The resolution for input images"
  506. ),
  507. )
  508. parser.add_argument(
  509. "--displacement_rate",
  510. type=int,
  511. default=32,
  512. help=(
  513. "The rate for displacement"
  514. ),
  515. )
  516. parser.add_argument(
  517. "--random_y_squeeze_rate",
  518. type=float,
  519. default=0.05,
  520. help=(
  521. "The rate for random y squeeze"
  522. ),
  523. )
  524. parser.add_argument(
  525. "--random_rotation_degree",
  526. type=int,
  527. default=5,
  528. help=(
  529. "The degree for random rotation"
  530. ),
  531. )
  532. parser.add_argument(
  533. "--random_shift_range",
  534. type=int,
  535. default=2,
  536. help=(
  537. "The range for random shift"
  538. ),
  539. )
  540. parser.add_argument(
  541. "--image_condition_drop_rate",
  542. type=float,
  543. default=0.1,
  544. help=(
  545. "The rate for image condition drop"
  546. ),
  547. )
  548. parser.add_argument(
  549. "--learning_rate",
  550. type=float,
  551. default=1e-4,
  552. help="Learning rate to use.",
  553. )
  554. parser.add_argument("--weight_decay", type=float, default=1e-2, help="Weight decay to use.")
  555. parser.add_argument("--num_train_epochs", type=int, default=100)
  556. parser.add_argument(
  557. "--max_train_steps",
  558. type=int,
  559. default=None,
  560. help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
  561. )
  562. parser.add_argument(
  563. "--gradient_accumulation_steps",
  564. type=int,
  565. default=1,
  566. help="Number of updates steps to accumulate before performing a backward/update pass.",
  567. )
  568. parser.add_argument(
  569. "--scale_lr",
  570. action="store_true",
  571. default=False,
  572. help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
  573. )
  574. parser.add_argument(
  575. "--train_batch_size", type=int, default=8, help="Batch size (per device) for the training dataloader."
  576. )
  577. parser.add_argument(
  578. "--dataloader_num_workers",
  579. type=int,
  580. default=0,
  581. help=(
  582. "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process."
  583. ),
  584. )
  585. parser.add_argument(
  586. "--checkpointing_steps",
  587. type=int,
  588. default=500,
  589. help=(
  590. "Save a checkpoint of the training state every X updates. Checkpoints can be used for resuming training via `--resume_from_checkpoint`. "
  591. "In the case that the checkpoint is better than the final trained model, the checkpoint can also be used for inference."
  592. "Using a checkpoint for inference requires separate loading of the original pipeline and the individual checkpointed model components."
  593. "See https://huggingface.co/docs/diffusers/main/en/training/dreambooth#performing-inference-using-a-saved-checkpoint for step by step"
  594. "instructions."
  595. ),
  596. )
  597. parser.add_argument(
  598. "--resume_from_checkpoint",
  599. type=str,
  600. default=None,
  601. help=(
  602. "Whether training should be resumed from a previous checkpoint. Use a path saved by"
  603. ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
  604. ),
  605. )
  606. parser.add_argument(
  607. "--mixed_precision",
  608. type=str,
  609. default=None,
  610. choices=["no", "fp16", "bf16"],
  611. help=(
  612. "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
  613. " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
  614. " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
  615. ),
  616. )
  617. parser.add_argument(
  618. "--report_to",
  619. type=str,
  620. default="tensorboard",
  621. help=(
  622. 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`'
  623. ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.'
  624. ),
  625. )
  626. parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank")
  627. args = parser.parse_args()
  628. env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
  629. if env_local_rank != -1 and env_local_rank != args.local_rank:
  630. args.local_rank = env_local_rank
  631. return args
  632. def main():
  633. args = parse_args()
  634. logging_dir = Path(args.output_dir, args.logging_dir)
  635. accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir)
  636. accelerator = Accelerator(
  637. gradient_accumulation_steps=args.gradient_accumulation_steps,
  638. mixed_precision=args.mixed_precision,
  639. log_with=args.report_to,
  640. project_config=accelerator_project_config,
  641. )
  642. # Make one log on every process with the configuration for debugging.
  643. logging.basicConfig(
  644. format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
  645. datefmt="%m/%d/%Y %H:%M:%S",
  646. level=logging.INFO,
  647. )
  648. logger.info(accelerator.state, main_process_only=False)
  649. wandb.init(project="DWIDistortion")
  650. if args.seed is not None:
  651. set_seed(args.seed)
  652. if accelerator.is_main_process:
  653. if args.output_dir is not None:
  654. os.makedirs(args.output_dir, exist_ok=True)
  655. # Load scheduler and models.
  656. noise_scheduler = DDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler")
  657. vae = AutoencoderKL.from_pretrained(args.pretrained_model_name_or_path, subfolder="vae")
  658. unet = UNet2DConditionModel.from_pretrained(args.pretrained_model_name_or_path, subfolder="unet", cross_attention_dim=1280, ignore_mismatched_sizes=True, low_cpu_mem_usage=False)
  659. image_encoder = CLIPVisionModelWithProjection.from_pretrained(args.image_encoder_path)
  660. # freeze parameters of models to save more memory
  661. unet.requires_grad_(False)
  662. unet.conv_out.requires_grad_(True)
  663. pretrained_state_dict = UNet2DConditionModel.from_pretrained(
  664. args.pretrained_model_name_or_path, subfolder="unet"
  665. ).state_dict()
  666. # Enable gradients for randomly initialized layers
  667. # (layers whose shapes changed due to ignore_mismatched_sizes=True)
  668. random_initialized_layers = []
  669. for name, param in unet.named_parameters():
  670. if param.shape != pretrained_state_dict[name].shape:
  671. param.requires_grad = True
  672. random_initialized_layers.append(name)
  673. print("Random initialized layers: ", random_initialized_layers)
  674. # Free up memory by deleting pretrained_state_dict
  675. del pretrained_state_dict
  676. vae.requires_grad_(False)
  677. image_encoder.requires_grad_(True)
  678. controlnet = ControlNetModel.from_unet(unet, conditioning_channels=1)
  679. controlnet.train()
  680. unet.train()
  681. vae.eval()
  682. image_encoder.train()
  683. # Taken from [Sayak Paul's Diffusers PR #6511](https://github.com/huggingface/diffusers/pull/6511/files)
  684. def unwrap_model(model):
  685. model = accelerator.unwrap_model(model)
  686. model = model._orig_mod if is_compiled_module(model) else model
  687. return model
  688. # `accelerate` 0.16.0 will have better support for customized saving
  689. if version.parse(accelerate.__version__) >= version.parse("0.16.0"):
  690. # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
  691. def save_model_hook(models, weights, output_dir):
  692. if accelerator.is_main_process:
  693. models[0].save_pretrained(os.path.join(output_dir, "unet"))
  694. weights.pop()
  695. models[1].save_pretrained(os.path.join(output_dir, "controlnet"))
  696. weights.pop()
  697. models[2].save_pretrained(os.path.join(output_dir, "image_encoder"))
  698. weights.pop()
  699. def load_model_hook(models, input_dir):
  700. model = models.pop()
  701. load_image_encoder = CLIPVisionModelWithProjection.from_pretrained(os.path.join(input_dir, "image_encoder"))
  702. model.load_state_dict(load_image_encoder.state_dict())
  703. del load_image_encoder
  704. model = models.pop()
  705. load_controlnet = ControlNetModel.from_pretrained(input_dir, subfolder="controlnet")
  706. model.register_to_config(**load_controlnet.config)
  707. model.load_state_dict(load_controlnet.state_dict())
  708. del load_controlnet
  709. model = models.pop()
  710. load_unet = UNet2DConditionModel.from_pretrained(input_dir, subfolder="unet")
  711. model.load_state_dict(load_unet.state_dict())
  712. del load_unet
  713. accelerator.register_save_state_pre_hook(save_model_hook)
  714. accelerator.register_load_state_pre_hook(load_model_hook)
  715. if args.scale_lr:
  716. args.learning_rate = (
  717. args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes
  718. )
  719. weight_dtype = torch.float32
  720. if accelerator.mixed_precision == "fp16":
  721. weight_dtype = torch.float16
  722. elif accelerator.mixed_precision == "bf16":
  723. weight_dtype = torch.bfloat16
  724. vae.to(accelerator.device, dtype=weight_dtype)
  725. params_to_opt = itertools.chain(
  726. image_encoder.parameters(),
  727. unet.conv_out.parameters(),
  728. (param for name, param in unet.named_parameters() if name in random_initialized_layers), # randomly initialized layers
  729. controlnet.parameters()
  730. )
  731. optimizer = torch.optim.AdamW(params_to_opt, lr=args.learning_rate, weight_decay=args.weight_decay)
  732. # dataloader
  733. train_dataset = MyDataset(args.data_json_file, size=args.resolution, displacement_rate=args.displacement_rate, random_y_squeeze_rate=args.random_y_squeeze_rate, random_rotation_degree=args.random_rotation_degree, random_shift_range=args.random_shift_range, image_root_path=args.data_root_path, clip_image_processor_path=args.image_encoder_path, isTrain=True)
  734. train_dataloader = torch.utils.data.DataLoader(
  735. train_dataset,
  736. shuffle=True,
  737. collate_fn=collate_fn,
  738. batch_size=args.train_batch_size,
  739. num_workers=args.dataloader_num_workers,
  740. pin_memory=True,
  741. drop_last=True,
  742. )
  743. # Prepare everything with our `accelerator`.
  744. unet, controlnet, image_encoder, optimizer, train_dataloader = accelerator.prepare(unet, controlnet, image_encoder, optimizer, train_dataloader)
  745. # Scheduler and math around the number of training steps.
  746. overrode_max_train_steps = False
  747. num_update_steps_per_epoch = math.ceil(
  748. len(train_dataloader) / args.gradient_accumulation_steps
  749. )
  750. if args.max_train_steps is None:
  751. args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
  752. overrode_max_train_steps = True
  753. # Train!
  754. total_batch_size = (
  755. args.train_batch_size
  756. * accelerator.num_processes
  757. * args.gradient_accumulation_steps
  758. )
  759. logger.info("***** Running training *****")
  760. logger.info(f" Num examples = {len(train_dataset)}")
  761. logger.info(f" Num Epochs = {args.num_train_epochs}")
  762. logger.info(f" Instantaneous batch size per device = {args.train_batch_size}")
  763. logger.info(
  764. f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}"
  765. )
  766. logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
  767. logger.info(f" Total optimization steps = {args.max_train_steps}")
  768. global_step = 0
  769. first_epoch = 0
  770. initial_global_step = 0
  771. train_losses = []
  772. train_noise_losses = []
  773. val_psnrs = []
  774. val_ssims = []
  775. progress_bar = tqdm(
  776. range(0, args.max_train_steps),
  777. initial=initial_global_step,
  778. desc="Steps",
  779. # Only show the progress bar once on each machine.
  780. disable=not accelerator.is_local_main_process,
  781. )
  782. for epoch in range(first_epoch, args.num_train_epochs):
  783. begin = time.perf_counter()
  784. epoch_train_loss = 0.0
  785. epoch_train_noise_loss = 0.0
  786. for step, batch in enumerate(train_dataloader):
  787. if accelerator.sync_gradients:
  788. if accelerator.is_main_process:
  789. if global_step % args.checkpointing_steps == 0:
  790. save_path = os.path.join(args.output_dir, f"checkpoint-last")
  791. accelerator.save_state(save_path)
  792. logger.info(f"Saved state to {save_path}")
  793. if args.val_data_json_file is not None and args.plot_data_json_file is not None:
  794. val_psnr, val_ssim, val_psnr_cropped, val_ssim_cropped = log_validation(args, accelerator, weight_dtype, global_step, save_path)
  795. if global_step!=0 and (val_psnr_cropped+val_ssim_cropped) > (max(val_psnrs)+max(val_ssims)):
  796. save_path = os.path.join(args.output_dir, f"checkpoint-best-{global_step}")
  797. accelerator.save_state(save_path)
  798. logger.info(f"Saved state to {save_path}")
  799. val_psnrs.append(float(val_psnr_cropped))
  800. val_ssims.append(float(val_ssim_cropped))
  801. load_data_time = time.perf_counter() - begin
  802. with accelerator.accumulate(unet, controlnet, image_encoder):
  803. # Convert images to latent space
  804. with torch.no_grad():
  805. latents = vae.encode(batch["images"].to(accelerator.device, dtype=weight_dtype)).latent_dist.sample()
  806. latents = latents * vae.config.scaling_factor
  807. # Sample noise that we'll add to the latents
  808. noise = torch.randn_like(latents)
  809. bsz = latents.shape[0]
  810. # Sample a random timestep for each image
  811. timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (bsz,), device=latents.device)
  812. timesteps = timesteps.long().to(accelerator.device)
  813. # Add noise to the latents according to the noise magnitude at each timestep
  814. # (this is the forward diffusion process)
  815. noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps).to(accelerator.device, dtype=weight_dtype)
  816. encoder_hidden_states = image_encoder(batch["clip_images"].to(accelerator.device, dtype=weight_dtype)).last_hidden_state
  817. controlnet_image = batch["control_images"].to(accelerator.device, dtype=weight_dtype)
  818. down_block_res_samples, mid_block_res_sample = controlnet(
  819. noisy_latents,
  820. timesteps,
  821. encoder_hidden_states,
  822. controlnet_cond=controlnet_image,
  823. return_dict=False,
  824. )
  825. noise_pred = unet(
  826. noisy_latents,
  827. timesteps,
  828. encoder_hidden_states,
  829. down_block_additional_residuals=[
  830. sample.to(dtype=noisy_latents.dtype) for sample in down_block_res_samples
  831. ],
  832. mid_block_additional_residual=mid_block_res_sample.to(dtype=noisy_latents.dtype),
  833. return_dict=False,
  834. )[0]
  835. noise_loss = F.mse_loss(noise_pred.float(), noise.float(), reduction="mean")
  836. timesteps[timesteps == 0] = 1
  837. loss = noise_loss
  838. # Gather the losses across all processes for logging (if we use distributed training).
  839. avg_loss = accelerator.gather(loss.repeat(args.train_batch_size)).mean().item()
  840. avg_noise_loss = accelerator.gather(noise_loss.repeat(args.train_batch_size)).mean().item()
  841. # accumulate the average loss for each batch
  842. epoch_train_loss += avg_loss
  843. epoch_train_noise_loss += avg_noise_loss
  844. # Backward pass and optimization for main model
  845. optimizer.zero_grad()
  846. accelerator.backward(loss, retain_graph=True)
  847. optimizer.step()
  848. # Checks if the accelerator has performed an optimization step behind the scenes
  849. if accelerator.sync_gradients:
  850. progress_bar.update(1)
  851. progress_bar.set_postfix_str(f"Loss: {avg_loss:.3f}")
  852. global_step += 1
  853. if global_step >= args.max_train_steps:
  854. break
  855. begin = time.perf_counter()
  856. epoch_train_loss /= len(train_dataloader)
  857. epoch_train_noise_loss /= len(train_dataloader)
  858. wandb.log({
  859. "train_loss": epoch_train_loss,
  860. "train_noise_loss": epoch_train_noise_loss,
  861. "step": global_step,
  862. "learning_rate": optimizer.param_groups[0]["lr"],
  863. })
  864. train_losses.append(epoch_train_loss)
  865. train_noise_losses.append(epoch_train_noise_loss)
  866. losses_data = {
  867. "train_losses": train_losses,
  868. "train_noise_losses": train_noise_losses,
  869. "val_psnrs": val_psnrs,
  870. "val_ssims": val_ssims,
  871. }
  872. with open(os.path.join(args.output_dir, "losses.json"), "w") as f:
  873. json.dump(losses_data, f)
  874. # Create the pipeline using using the trained modules and save it.
  875. accelerator.wait_for_everyone()
  876. if accelerator.is_main_process:
  877. save_path = os.path.join(args.output_dir, "final")
  878. if not os.path.isdir(save_path):
  879. os.mkdir(save_path)
  880. unet = unwrap_model(unet)
  881. unet.save_pretrained(os.path.join(save_path, "unet"))
  882. controlnet = unwrap_model(controlnet)
  883. controlnet.save_pretrained(os.path.join(save_path, "controlnet"))
  884. image_encoder = unwrap_model(image_encoder)
  885. image_encoder.save_pretrained(os.path.join(save_path, "image_encoder"))
  886. logger.info(f"Saved final model to {save_path}")
  887. accelerator.end_training()
  888. if __name__ == "__main__":
  889. main()

train_DeDistortNet.py at commit 5534062, under Apache-2.0 · at the source

Overview

Authors: Inye Na1, Qi Miao2, Jonghun Kim1, Kyunghyun Sung2, Hyunjin Park1
ORCID iDs: Hyunjin Park
  1. Department of Electrical and Computer Engineering, Sungkyunkwan University,Suwon, Republic of Korea
  2. Department of Radiological Sciences, David Geffen School of Medicine, University of California, Los Angeles,Los Angeles, CA USA
Institutions: Sungkyunkwan University (South Korea); University of California, Los Angeles (United States)
Journal: European radiology experimental, volume 10, issue 1, article 66
Dates: received 5 November 2025; accepted 17 April 2026; published online 13 May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1186/s41747-026-00735-w · PMID 42126718 · PMCID PMC13172242 · OpenAlex W7161056552
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism)
Methods: Connectivity
Keywords: Artifacts, Artificial intelligence, Deep learning, Diffusion magnetic resonance imaging, Prostate
MeSH: Diffusion Magnetic Resonance Imaging*, Image Interpretation, Computer-Assisted*, Image Processing, Computer-Assisted*, Prostate*, Prostatic Neoplasms*, Artifacts, Humans, Male (* major topic)
Topic: Prostate Cancer Diagnosis and Treatment (Pulmonary and Respiratory Medicine, Medicine), according to OpenAlex
Funding: National Research Foundation of Korea (RS-2024-00408040)
Citations: not cited yet (Europe PMC); 39 references in the paper

Abstract

Objective: Geometric distortion from susceptibility artifacts in diffusion-weighted imaging (DWI) degrades anatomical fidelity and complicates clinical prostate magnetic resonance imaging (MRI) interpretation. We developed and evaluated DeDistortNet, a generative dual-conditioned diffusion model for correcting geometric distortions in prostate DWI without experimentally acquired paired distorted-undistorted data.

Materials and methods: We utilized the public PROSTATEx dataset, divided into training (n = 135, 1,893 slices), validation (n = 4, 64 slices), and test (n = 189, 2,623 slices) cohorts stratified by distortion severity. Model training used only undistorted DWIs, with simulated distortions to generate paired examples. DeDistortNet combines contextual guidance from distorted DWIs with structural guidance from T2-weighted images to synthesize distortion-free DWIs. Performance was evaluated using quantitative analysis on simulated distortions and indirect validation on clinically distorted DWIs through anatomical concordance with T2-derived prostate masks.

Results: The cohort included prostate MRI exams from 328 male subjects. The DWIs from 139 exams were undistorted for training/validation, while 189 were distorted for testing. In simulated data, DeDistortNet restored image quality across distortion severities, improving peak signal-to-noise ratio by 36% and structural similarity by 55% in the peripheral zone (PZ) under extreme distortion. In clinically distorted data, concordance with T2-weighted references (PZ Dice similarity) improved by 37% under severe distortion and by 72% under extreme distortion. Radiologist assessment further supported improved geometric fidelity and prostate boundary delineation after correction.

Conclusion: DeDistortNet, trained on undistorted DWIs with simulated distortions, effectively corrected distortions in prostate DWI and restored anatomical fidelity, particularly in the PZ, without requiring additional acquisitions or specialized imaging protocols.

Relevance statement: By correcting severe distortions without additional imaging sequences, DeDistortNet restores anatomical fidelity in prostate diffusion-weighted imaging, particularly in the peripheral zone, enabling more reliable image interpretation and reducing the need for repeat scans.

Key Points: DeDistortNet corrects geometric distortion in prostate diffusion MRI. Our method integrates anatomical guidance from T2-weighted images for distortion-free reconstruction. Our method was trained on simulated distortions without experimentally acquired paired distorted-undistorted data or additional MRI sequences.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 5 matches between paragraphs and lines of code.

nainye/DeDistortNet

License: Apache-2.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 553406232f87111a737a1582f96b1978f10f64e4, 10 June 2026
Languages: Jupyter (3), Shell (1), Python (1)
Size: 16 files, 5 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (scripts/requirements.txt), 3 notebooks
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (4 files), SimpleITK (4 files), Matplotlib (3 files), SciPy (2 files), pandas (1 file), pydicom (1 file), PyTorch (1 file), scikit-image (1 file), Hugging Face Transformers (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
7 files

Code availability

The derived slice-level distortion severity labels and model implementation codes developed for this study are publicly available in a dedicated GitHub repository at: https://github.com/nainye/DeDistortNet.

Reproduced under the paper's license (CC BY), from the paper cited above.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 5 scripts, each with its path and the digest of its content;
  • 5 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data availability

The dataset analyzed during this study is publicly available in The Cancer Imaging Archive (TCIA) at https://www.cancerimagingarchive.net/collection/prostatex/.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 5 keywords, 8 MeSH terms, 1 funder, 31 references.

Cite

This paper

Na, I., Miao, Q., Kim, J., Sung, K., & Park, H. (2026). Dual-conditioned diffusion model with anatomical guidance for geometric distortion correction in prostate MRI. European radiology experimental, 10(1), 66. https://doi.org/10.1186/s41747-026-00735-w

BibTeX

@article{na2026dual,
author = {Na, Inye and Miao, Qi and Kim, Jonghun and Sung, Kyunghyun and Park, Hyunjin},
title = {{Dual-conditioned diffusion model with anatomical guidance for geometric distortion correction in prostate MRI}},
journal = {European radiology experimental},
year = {2026},
month = may,
volume = {10},
number = {1},
pages = {66},
publisher = {Springer},
issn = {2509-9280},
doi = {10.1186/s41747-026-00735-w},
url = {https://doi.org/10.1186/s41747-026-00735-w},
pmid = {42126718},
pmcid = {PMC13172242}
}

RIS

TY - JOUR
AU - Na, Inye
AU - Miao, Qi
AU - Kim, Jonghun
AU - Sung, Kyunghyun
AU - Park, Hyunjin
TI - Dual-conditioned diffusion model with anatomical guidance for geometric distortion correction in prostate MRI
T2 - European radiology experimental
J2 - Eur Radiol Exp
PY - 2026
DA - 2026/05/13
VL - 10
IS - 1
SP - 66
SN - 2509-9280
PB - Springer
DO - 10.1186/s41747-026-00735-w
UR - https://doi.org/10.1186/s41747-026-00735-w
LA - en
ER -

CSL-JSON

{
"id": "10.1186/s41747-026-00735-w",
"type": "article-journal",
"title": "Dual-conditioned diffusion model with anatomical guidance for geometric distortion correction in prostate MRI",
"container-title": "European radiology experimental",
"author": [
{
"family": "Na",
"given": "Inye"
},
{
"family": "Miao",
"given": "Qi"
},
{
"family": "Kim",
"given": "Jonghun"
},
{
"family": "Sung",
"given": "Kyunghyun"
},
{
"family": "Park",
"given": "Hyunjin"
}
],
"container-title-short": "Eur Radiol Exp",
"volume": "10",
"issue": "1",
"page": "66",
"DOI": "10.1186/s41747-026-00735-w",
"PMID": "42126718",
"PMCID": "PMC13172242",
"ISSN": "2509-9280",
"publisher": "Springer",
"URL": "https://doi.org/10.1186/s41747-026-00735-w",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
13
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1111/joa.70203 [code]
Two-step workflow integrating automatic registration and manual refinement for the accurate alignment of serial histological sections in 3D reconstruction.
Journal: Journal of anatomy
In common: pydicom, Hugging Face Transformers, SimpleITK, 6 other tools
[2] doi:10.1162/imag.a.1326 [code]
RAVEN: Robust, generalizable, multi-resolution structural MRI upsampling using autoencoders.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Hugging Face Transformers, SimpleITK, scikit-image, 5 other tools, structural MRI / diffusion, 1 reference
[3] doi:10.1371/journal.pcbi.1014555 [code]
Body surface potential driven personalisation of electrophysiological digital twins in hypertrophic cardiomyopathy.
Journal: PLoS computational biology
In common: pydicom, SimpleITK, scikit-image, 5 other tools, structural MRI / diffusion
[4] doi:10.1080/07853890.2026.2685416 [code]
Pulmonary and cerebral damage in COVID-19 survivors: is there any association?
Journal: Annals of medicine
In common: pydicom, SimpleITK, scikit-image, 5 other tools, structural MRI / diffusion
[5] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: pydicom, SimpleITK, scikit-image, 5 other tools, structural MRI / diffusion
[6] doi:10.21037/qims-2026-0792 [code]
An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation.
Journal: Quantitative imaging in medicine and surgery
In common: SimpleITK, scikit-image, PyTorch, 4 other tools, structural MRI / diffusion, 2 references
[7] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: pydicom, Hugging Face Transformers, scikit-image, 5 other tools
[8] doi:10.1038/s41598-026-48496-1 [code]
A unified FLAIR hyperintensity segmentation model for various CNS tumor types and acquisition time points.
Journal: Scientific reports
In common: SimpleITK, scikit-image, pandas, 3 other tools, structural MRI / diffusion, 2 references
[9] doi:10.1002/hipo.70124 [code]
Association Between Anterior Hippocampal Gyrification and Episodic Memory Performance in Neurotypical Young Adults.
Journal: Hippocampus
In common: SimpleITK, scikit-image, PyTorch, 4 other tools, structural MRI / diffusion, 1 reference
[10] doi:10.1186/s12880-026-02335-x [code]
Automatic lateral ventricle and choroid plexus segmentation method in infant brain MR images.
Journal: BMC medical imaging
In common: SimpleITK, scikit-image, PyTorch, 4 other tools, structural MRI / diffusion, 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.