OSCR

Accurate estimation of intravoxel incoherent motion parameters based on implicit neural representation.

Code ↔ Paper

4 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 4 matches
  1. [1] § METHODS › Overview of the two‐stage IVIM‐INR framework ↔ train.py, lines 56–95 · score 0.76 · Xavier uniform initialization, hidden layers, ReLU, MLP, Weights, linear
  2. [2] § METHODS › Implementation details ↔ train.py, lines 450–490 · score 0.67 · cosine annealing learning, rate scheduling, Adam, GPU, Training, optimizer
  3. [3] § METHODS › Overview of the two‐stage IVIM‐INR framework ↔ train.py, lines 1–35 · score 0.60 · implicit neural representation, ReLU, IVIM INR, S2, global, S1
  4. [4] § METHODS › Overview of the two‐stage IVIM‐INR framework ↔ train.py, lines 56–95 · score 0.59 · layer MLP, hidden layers, ReLU, IVIM INR, SIREN, denoised

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,690 lines · 71 KB · no license · 4 matches

  1. #!/usr/bin/env python3
  2. """
  3. IVIM-INR: Two-stage implicit neural representation for IVIM parameter estimation.
  4. Stage 1 (S1): 4-layer ReLU MLP (512 hidden units per layer) that denoises each
  5. b-value image, taking as input spatial coordinates plus 3D patches
  6. from the other (non-target) b-value images.
  7. Stage 2 (S2): 4-layer SIREN (512 hidden units per layer) that maps spatial
  8. coordinates (plus a patch from the denoised b=0 image) to the
  9. IVIM parameters (D_p, D_t, F_p) and a global offset delta.
  10. Usage
  11. -----
  12. # Default: noisy SNR=50 brain phantom in ./data_1/SNR_50
  13. python ivim_fitting.py --data_dir data_1/SNR_50
  14. # Ground-truth (noise-free) fit
  15. python ivim_fitting.py --data_dir data_1 --no-use_noisy_data
  16. """
  17. import os
  18. import sys
  19. import time
  20. import json
  21. import argparse
  22. from typing import List, Tuple
  23. import numpy as np
  24. import nibabel as nib
  25. import torch
  26. import torch.nn as nn
  27. import torch.nn.functional as F
  28. from scipy import ndimage
  29. import pydicom
  30. from datetime import datetime
  31. # Configuration classes defined inline (previously from config_loader)
  32. # Default IVIM parameter ranges
  33. IVIM_PARAMETER_RANGES = {
  34. 'Dp': {'min': 0.001, 'max': 0.02}, # mm²/s
  35. 'Dt': {'min': 0.00001, 'max': 0.003}, # mm²/s
  36. 'Fp': {'min': 0.0, 'max': 0.5} # fraction
  37. }
  38. # ==================== Denoising Model Components ====================
  39. def fc_block(in_size, out_size, dropout, *args, **kwargs):
  40. """Fully connected block with ReLU and dropout."""
  41. return nn.Sequential(
  42. nn.Linear(in_size, out_size, *args, **kwargs),
  43. nn.ReLU(),
  44. nn.Dropout(dropout),
  45. )
  46. class MLPv1(nn.Module):
  47. """ReLU MLP used for DWI denoising (Stage 1 of IVIM-INR).
  48. num_layers counts TOTAL linear layers (input + hidden + output),
  49. matching the SirenNet convention. With the default num_layers=4 and
  50. hidden_size=512, this builds a 4-layer MLP where every hidden layer
  51. has 512 units (one input projection, two 512->512 hidden projections,
  52. and one output projection).
  53. """
  54. def __init__(self, input_size=3, hidden_size=512, output_size=1, dropout=0, num_layers=4):
  55. super(MLPv1, self).__init__()
  56. assert num_layers >= 2, "num_layers must be >= 2 (input + output)"
  57. self.input_size = input_size
  58. self.hidden_size = hidden_size
  59. self.output_size = output_size
  60. self.dropout = dropout
  61. self.num_layers = num_layers
  62. layers = [nn.Linear(input_size, hidden_size), nn.ReLU()]
  63. if dropout > 0:
  64. layers.append(nn.Dropout(dropout))
  65. for _ in range(num_layers - 2):
  66. layers.append(nn.Linear(hidden_size, hidden_size))
  67. layers.append(nn.ReLU())
  68. if dropout > 0:
  69. layers.append(nn.Dropout(dropout))
  70. layers.append(nn.Linear(hidden_size, output_size))
  71. # Xavier uniform initialization for all Linear layers
  72. for m in layers:
  73. if isinstance(m, nn.Linear):
  74. nn.init.xavier_uniform_(m.weight)
  75. if m.bias is not None:
  76. nn.init.zeros_(m.bias)
  77. self.net = nn.Sequential(*layers)
  78. def forward(self, x):
  79. x = x.view(-1, self.input_size)
  80. return self.net(x)
  81. class IntegratedPositionalEncoding(nn.Module):
  82. """Integrated Positional Encoding from MDINR."""
  83. def __init__(self, num_freqs: int, max_freq: float):
  84. super().__init__()
  85. self.freqs = torch.linspace(1.0, max_freq, steps=num_freqs)
  86. self.omega_max = max_freq
  87. def forward(self, x: torch.Tensor) -> torch.Tensor:
  88. """
  89. x: Tensor of shape [N, 3] (3D coordinates, normalized to [-1,1])
  90. Returns: Tensor of shape [N, 3 * F * 2] (concatenated sin/cos encoding)
  91. """
  92. N, D = x.shape
  93. device = x.device
  94. # Shape freqs to [1, 1, F]
  95. f = self.freqs.to(device).view(1, 1, -1)
  96. mu = x.view(N, D, 1)
  97. # Calculate sinc with d = π / ω_max
  98. d = torch.pi / self.omega_max
  99. sinc_arg = f * d
  100. sinc = torch.where(
  101. sinc_arg == 0,
  102. torch.ones_like(sinc_arg),
  103. torch.sin(sinc_arg) / sinc_arg
  104. )
  105. # Calculate sin and cos terms
  106. x_freqed = mu * f
  107. sin_term = torch.sin(x_freqed) * sinc
  108. cos_term = torch.cos(x_freqed) * sinc
  109. # Flatten and concatenate
  110. sin_term = sin_term.view(N, -1)
  111. cos_term = cos_term.view(N, -1)
  112. return torch.cat([sin_term, cos_term], dim=-1)
  113. # ==================== Scanner Type Detection ====================
  114. def detect_scanner_type(patient_id: str, date: str, dicom_base_dir: str) -> str:
  115. """
  116. Detect if the scanner is Unity (MR Linac) or MR Sim by checking DICOM header.
  117. Args:
  118. patient_id: Patient identifier
  119. date: Scan date
  120. dicom_base_dir: Base directory containing DICOM files
  121. Returns:
  122. 'unity' if Unity scanner, 'sim' if MR Sim scanner
  123. """
  124. # Find DWI directory
  125. patient_dir = os.path.join(dicom_base_dir, patient_id, date)
  126. if not os.path.exists(patient_dir):
  127. print(f" ⚠ Patient directory not found: {patient_dir}")
  128. return 'sim' # Default to sim
  129. # Look for DWI directories
  130. dwi_dirs = []
  131. for dirname in os.listdir(patient_dir):
  132. if 'DWI' in dirname.upper() and os.path.isdir(os.path.join(patient_dir, dirname)):
  133. dwi_dirs.append(os.path.join(patient_dir, dirname))
  134. if not dwi_dirs:
  135. print(f" ⚠ No DWI directories found in {patient_dir}")
  136. return 'sim' # Default to sim
  137. # Check first DICOM file in first DWI directory
  138. dwi_dir = dwi_dirs[0]
  139. dcm_files = [f for f in os.listdir(dwi_dir) if f.endswith('.dcm') or f.endswith('.DCM')]
  140. if not dcm_files:
  141. print(f" ⚠ No DICOM files found in {dwi_dir}")
  142. return 'sim' # Default to sim
  143. # Read first DICOM file
  144. dcm_path = os.path.join(dwi_dir, dcm_files[0])
  145. try:
  146. ds = pydicom.dcmread(dcm_path)
  147. # Check for (0008,1090) - ManufacturerModelName tag
  148. # Unity scanners have ManufacturerModelName = "Marlin"
  149. if hasattr(ds, 'ManufacturerModelName') or (0x0008, 0x1090) in ds:
  150. model_name = str(getattr(ds, 'ManufacturerModelName', '')).strip()
  151. print(f" ManufacturerModelName (0008,1090): {model_name}")
  152. if model_name.lower() == 'marlin':
  153. print(f" → Unity (MR Linac) detected")
  154. return 'unity'
  155. else:
  156. print(f" → MR Sim detected")
  157. return 'sim'
  158. else:
  159. print(f" No ManufacturerModelName tag (0008,1090) found - MR Sim detected")
  160. return 'sim'
  161. except Exception as e:
  162. print(f" ⚠ Error reading DICOM file {dcm_path}: {e}")
  163. return 'sim' # Default to sim
  164. # ==================== Noise Floor Detection ====================
  165. def detect_noise_floor_mask(b0_data: np.ndarray) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
  166. """
  167. Detect noise floor regions in Unity data based on intensity thresholds.
  168. Args:
  169. b0_data: B0 image data
  170. Returns:
  171. background_mask: Pure background (intensity < 50)
  172. noise_floor_mask: Noise floor regions (50 <= intensity <= 300)
  173. tissue_mask: Signal regions (intensity > 300)
  174. """
  175. # Create masks based on thresholds
  176. background_mask = b0_data < 50
  177. noise_floor_mask = (b0_data >= 50) & (b0_data <= 300)
  178. tissue_mask = b0_data > 300
  179. # Get largest connected component of noise floor
  180. if np.any(noise_floor_mask):
  181. labeled, num_features = ndimage.label(noise_floor_mask)
  182. if num_features > 0:
  183. # Find largest component
  184. component_sizes = np.bincount(labeled.ravel())
  185. component_sizes[0] = 0 # Ignore background
  186. largest_component = component_sizes.argmax()
  187. # Keep only largest component as noise floor
  188. noise_floor_mask_clean = labeled == largest_component
  189. # Reassign small disconnected regions to tissue
  190. small_regions = noise_floor_mask & ~noise_floor_mask_clean
  191. tissue_mask = tissue_mask | small_regions
  192. noise_floor_mask = noise_floor_mask_clean
  193. return background_mask, noise_floor_mask, tissue_mask
  194. def rician_noise_correction(dwi_data: np.ndarray, noise_floor_mask: np.ndarray,
  195. method: str = "rayleigh_mean") -> Tuple[np.ndarray, float]:
  196. """
  197. Apply Rician noise floor correction to DWI data using noise floor regions.
  198. Args:
  199. dwi_data: DWI data array [X, Y, Z, B]
  200. noise_floor_mask: Binary mask indicating noise floor regions
  201. method: Correction method ('rayleigh_mean' or 'simple_subtraction')
  202. Returns:
  203. corrected_dwi: Corrected DWI data
  204. median_sigma: Estimated noise standard deviation
  205. """
  206. _, _, _, num_b = dwi_data.shape
  207. corrected_dwi = np.zeros_like(dwi_data)
  208. # Estimate noise from noise floor regions for each b-value
  209. sigma_values = []
  210. for b_idx in range(num_b):
  211. # Extract noise floor values
  212. noise_values = dwi_data[:, :, :, b_idx][noise_floor_mask]
  213. if len(noise_values) > 0:
  214. # For Rayleigh distribution: σ = mean / sqrt(π/2)
  215. mean_noise = np.mean(noise_values)
  216. sigma = mean_noise / np.sqrt(np.pi / 2)
  217. sigma_values.append(sigma)
  218. print(f" b={b_idx}: Noise floor mean={mean_noise:.3f}, σ={sigma:.3f}")
  219. if method == "rayleigh_mean":
  220. # Rician correction: S_corrected = sqrt(max(S^2 - 2σ^2, 0))
  221. S_squared = dwi_data[:, :, :, b_idx] ** 2
  222. S_corrected_squared = np.maximum(S_squared - 2 * sigma**2, 0)
  223. corrected_dwi[:, :, :, b_idx] = np.sqrt(S_corrected_squared)
  224. elif method == "simple_subtraction":
  225. # Simple subtraction method
  226. corrected_dwi[:, :, :, b_idx] = np.maximum(dwi_data[:, :, :, b_idx] - sigma, 0)
  227. else:
  228. print(f" b={b_idx}: No noise floor regions available, skipping correction")
  229. corrected_dwi[:, :, :, b_idx] = dwi_data[:, :, :, b_idx]
  230. sigma_values.append(0)
  231. # Use median sigma across b-values
  232. median_sigma = np.median(sigma_values) if sigma_values else 0
  233. print(f" Median σ across b-values: {median_sigma:.3f}")
  234. return corrected_dwi, median_sigma
  235. # ==================== Denoising Functions ====================
  236. def precompute_patch_features(noisy_dwi_padded: List[torch.Tensor], patch_coords: torch.Tensor,
  237. X_dim: int, Y_dim: int, Z_dim: int, num_b: int,
  238. device: torch.device, mask: np.ndarray = None) -> Tuple[torch.Tensor, np.ndarray]:
  239. """
  240. Precompute patch features for masked voxels and all b-values.
  241. Args:
  242. noisy_dwi_padded: List of padded DWI volumes for each b-value
  243. patch_coords: Patch coordinate offsets
  244. X_dim, Y_dim, Z_dim: Volume dimensions
  245. num_b: Number of b-values
  246. device: Device to use
  247. mask: Binary mask to select voxels (if None, use all voxels)
  248. Returns:
  249. all_patch_features: Tensor of shape [num_masked_voxels, num_b, patch_feature_size]
  250. voxel_coords: Array of voxel coordinates for masked voxels
  251. """
  252. print(" Precomputing patch features for masked voxels...")
  253. # Create coordinate grids
  254. all_indices = np.indices((X_dim, Y_dim, Z_dim))
  255. # Apply mask if provided
  256. if mask is not None:
  257. mask_indices = np.where(mask)
  258. voxel_coords = np.stack([mask_indices[0], mask_indices[1], mask_indices[2]], axis=1)
  259. print(f" Using mask: {len(voxel_coords):,} voxels out of {X_dim*Y_dim*Z_dim:,}")
  260. else:
  261. voxel_coords = np.stack([all_indices[0].ravel(),
  262. all_indices[1].ravel(),
  263. all_indices[2].ravel()], axis=1)
  264. num_voxels = len(voxel_coords)
  265. patch_feature_size = len(patch_coords)
  266. # Preallocate the feature tensor
  267. all_patch_features = torch.zeros((num_voxels, num_b, patch_feature_size),
  268. dtype=torch.float32, device=device)
  269. # Process in batches to manage memory
  270. batch_size = 5000
  271. for start_idx in range(0, num_voxels, batch_size):
  272. end_idx = min(start_idx + batch_size, num_voxels)
  273. batch_coords = voxel_coords[start_idx:end_idx]
  274. batch_coords_tensor = torch.tensor(batch_coords, dtype=torch.float32, device=device)
  275. # Extract patches for this batch
  276. # batch_coords are in original volume space, need to add padding offset
  277. # then add patch offsets (which already include padding)
  278. batch_coords_padded = batch_coords_tensor + 1 # Add padding offset to center coords
  279. coords_patch = batch_coords_padded[:, None, :] + patch_coords[None, :, :] - 1 # Subtract 1 because patch_coords already has padding
  280. x_idx = coords_patch[..., 0].reshape(-1).long()
  281. y_idx = coords_patch[..., 1].reshape(-1).long()
  282. z_idx = coords_patch[..., 2].reshape(-1).long()
  283. # Extract features for each b-value
  284. for b_idx in range(num_b):
  285. features = noisy_dwi_padded[b_idx][0, x_idx, y_idx, z_idx]
  286. features = features.reshape(len(batch_coords), -1)
  287. all_patch_features[start_idx:end_idx, b_idx, :] = features
  288. if (start_idx + batch_size) % 100000 == 0:
  289. print(f" Processed {end_idx}/{num_voxels} voxels")
  290. print(f" ✔ Precomputed patch features: shape {all_patch_features.shape}")
  291. print(f" Memory usage: {all_patch_features.element_size() * all_patch_features.nelement() / (1024**3):.2f} GB")
  292. return all_patch_features, voxel_coords
  293. def denoise_dwi_mdinr_optimized(dwi_data: np.ndarray, mask: np.ndarray, b_values: np.ndarray,
  294. args, device: torch.device, affine=None, save_intermediate=False,
  295. output_dir=None) -> np.ndarray:
  296. """
  297. Denoise DWI using MDINR INR model (Optimized version without DataLoader).
  298. Args:
  299. dwi_data: DWI data array [X, Y, Z, B]
  300. mask: Binary mask
  301. b_values: Array of b-values
  302. args: Command line arguments
  303. device: PyTorch device
  304. affine: Affine matrix for saving intermediate results (optional)
  305. save_intermediate: Whether to save each b-value after denoising (optional)
  306. output_dir: Directory for saving intermediate results (optional)
  307. Returns:
  308. denoised_dwi: Denoised DWI data
  309. """
  310. print(f"\n{'='*60}")
  311. print(f"DWI Denoising (MDINR INR - Optimized)")
  312. print(f"{'='*60}")
  313. X_dim, Y_dim, Z_dim, num_b = dwi_data.shape
  314. # Normalize DWI data by percentile
  315. percentiles = [np.percentile(dwi_data[:, :, :, b][mask], 99) for b in range(num_b)]
  316. noisy_dwi_norm = np.zeros_like(dwi_data)
  317. for b in range(num_b):
  318. noisy_dwi_norm[:, :, :, b] = dwi_data[:, :, :, b] / (percentiles[b] + 1e-6)
  319. # Setup positional encoding
  320. input_mapper = IntegratedPositionalEncoding(
  321. num_freqs=args.denoise_npe,
  322. max_freq=args.denoise_maxf
  323. ).to(device)
  324. # Convert to PyTorch tensors
  325. noisy_dwi_torch = [torch.from_numpy(noisy_dwi_norm[:, :, :, b]).float().unsqueeze(0).unsqueeze(0)
  326. for b in range(num_b)]
  327. # Pad volumes for patch extraction
  328. # For 3x3x3 patch, we need padding of 1 on each side
  329. pad_size = 1
  330. noisy_dwi_padded = [F.pad(img, (pad_size, pad_size, pad_size, pad_size, pad_size, pad_size),
  331. mode='constant', value=0) for img in noisy_dwi_torch]
  332. noisy_dwi_padded = [img.to(device)[0] for img in noisy_dwi_padded]
  333. # Create patch coordinates for 3x3x3 patch
  334. # Offsets should be [-1, 0, 1] in each dimension
  335. patch_offsets = [-1, 0, 1]
  336. patch_coords = np.array([(x, y, z) for x in patch_offsets for y in patch_offsets for z in patch_offsets])
  337. # Adjust for padding: add pad_size to make coordinates valid in padded volume
  338. patch_coords = patch_coords + pad_size
  339. patch_coords = torch.tensor(patch_coords, dtype=torch.float32).to(device)
  340. # Precompute all patch features (only for masked voxels)
  341. all_patch_features, masked_voxel_coords = precompute_patch_features(
  342. noisy_dwi_padded, patch_coords, X_dim, Y_dim, Z_dim, num_b, device, mask
  343. )
  344. # Process each b-value
  345. denoised_dwi = np.zeros_like(dwi_data)
  346. for b_idx in range(num_b):
  347. print(f"\nDenoising b-value {b_idx} (b={b_values[b_idx]})...")
  348. # Input size calculation
  349. # For 3x3x3 patch: 27 features per b-value, times (num_b - 1) excluding current b
  350. patch_feature_size = 27 * (num_b - 1)
  351. if args.denoise_use_coords:
  352. coord_encoding_size = args.denoise_npe * 2 * 3
  353. input_size = coord_encoding_size + patch_feature_size
  354. else:
  355. input_size = patch_feature_size
  356. # Create MLPv1 model
  357. model = MLPv1(
  358. input_size=input_size,
  359. output_size=1,
  360. hidden_size=args.denoise_hidden_size,
  361. num_layers=args.denoise_num_layers,
  362. dropout=args.denoise_dropout
  363. ).to(device)
  364. # Use masked voxel coordinates
  365. points = masked_voxel_coords.astype(np.float32)
  366. # Normalize coordinates
  367. points_normalized = np.empty_like(points, dtype=np.float32)
  368. points_normalized[:, 0] = 2.0 * points[:, 0] / (X_dim - 1) - 1.0
  369. points_normalized[:, 1] = 2.0 * points[:, 1] / (Y_dim - 1) - 1.0
  370. points_normalized[:, 2] = 2.0 * points[:, 2] / (Z_dim - 1) - 1.0
  371. # Move all data to GPU at once
  372. coords_normal = torch.tensor(points_normalized, dtype=torch.float32).to(device)
  373. coords_prior = torch.tensor(points, dtype=torch.float32).to(device)
  374. # Extract labels for masked voxels
  375. labels_masked = noisy_dwi_norm[:, :, :, b_idx][mask]
  376. labels = torch.tensor(labels_masked.reshape(-1, 1), dtype=torch.float32).to(device)
  377. if len(labels) == 0:
  378. continue
  379. # Get precomputed features for current b-value (excluding current b)
  380. current_b_features = torch.cat([
  381. all_patch_features[:, other_b, :]
  382. for other_b in range(num_b) if other_b != b_idx
  383. ], dim=1)
  384. # Training
  385. optimizer = torch.optim.Adam(model.parameters(), lr=args.denoise_lr)
  386. scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.denoise_epochs)
  387. criterion = nn.MSELoss()
  388. model.train()
  389. num_samples = len(coords_normal)
  390. batch_size = args.denoise_batch_size
  391. for epoch in range(args.denoise_epochs):
  392. epoch_loss = 0.0
  393. epoch_start_time = time.time()
  394. # Shuffle indices for this epoch
  395. indices = torch.randperm(num_samples, device=device)
  396. # Manual batching
  397. for start_idx in range(0, num_samples, batch_size):
  398. batch_start_time = time.time()
  399. end_idx = min(start_idx + batch_size, num_samples)
  400. batch_indices = indices[start_idx:end_idx]
  401. # Get batch data - already on GPU
  402. coords_normal_batch = coords_normal[batch_indices]
  403. # coords_prior_batch = coords_prior[batch_indices] # Not used in training
  404. label_batch = labels[batch_indices]
  405. patch_features_batch = current_b_features[batch_indices]
  406. # Build model input
  407. if args.denoise_use_coords:
  408. coords_encoded = input_mapper(coords_normal_batch)
  409. model_input = torch.cat([coords_encoded, patch_features_batch], dim=1)
  410. else:
  411. model_input = patch_features_batch
  412. # Forward pass
  413. pred = model(model_input)
  414. loss = criterion(pred, label_batch)
  415. # Backward pass
  416. optimizer.zero_grad()
  417. loss.backward()
  418. optimizer.step()
  419. epoch_loss += loss.item() * batch_indices.size(0)
  420. epoch_loss /= num_samples
  421. scheduler.step()
  422. if (epoch + 1) % 10 == 0:
  423. epoch_time = time.time() - epoch_start_time
  424. print(f" Epoch {epoch+1}/{args.denoise_epochs}, Loss: {epoch_loss:.6f}, Time: {epoch_time:.2f}s")
  425. # Inference
  426. model.eval()
  427. predictions = []
  428. with torch.no_grad():
  429. # Process in batches for inference
  430. for start_idx in range(0, num_samples, batch_size * 2): # Larger batch for inference
  431. end_idx = min(start_idx + batch_size * 2, num_samples)
  432. coords_norm_inf = coords_normal[start_idx:end_idx]
  433. patch_features_inf = current_b_features[start_idx:end_idx]
  434. # Build model input
  435. if args.denoise_use_coords:
  436. coords_encoded = input_mapper(coords_norm_inf)
  437. model_input = torch.cat([coords_encoded, patch_features_inf], dim=1)
  438. else:
  439. model_input = patch_features_inf
  440. pred = model(model_input).cpu().numpy()
  441. predictions.append(pred)
  442. # Reshape predictions - place them in the correct locations
  443. predictions = np.concatenate(predictions, axis=0)
  444. # Create full volume and place predictions at masked locations
  445. denoised_volume = np.zeros((X_dim, Y_dim, Z_dim), dtype=np.float32)
  446. for i, (x, y, z) in enumerate(masked_voxel_coords):
  447. denoised_volume[x, y, z] = predictions[i, 0] * percentiles[b_idx]
  448. denoised_dwi[:, :, :, b_idx] = denoised_volume
  449. # Save intermediate result if requested
  450. if save_intermediate and affine is not None and output_dir is not None:
  451. intermediate_path = os.path.join(output_dir, f'dwi_ivim_b{int(b_values[b_idx])}_denoised.nii.gz')
  452. if not os.path.exists(intermediate_path):
  453. intermediate_img = nib.Nifti1Image(denoised_volume, affine)
  454. nib.save(intermediate_img, intermediate_path)
  455. print(f" ✔ Saved intermediate denoised b={int(b_values[b_idx])} to: {intermediate_path}")
  456. else:
  457. print(f" ⏭ Skipping b={int(b_values[b_idx])} (already exists): {intermediate_path}")
  458. print("\n✔ Denoising completed")
  459. return denoised_dwi
  460. # ==================== Neural Network Components ====================
  461. class SirenLayer(nn.Module):
  462. """SIREN layer with sine activation."""
  463. def __init__(self, in_features: int, out_features: int,
  464. bias: bool = True, is_first: bool = False, omega_0: float = 30.0):
  465. super().__init__()
  466. self.omega_0 = omega_0
  467. self.is_first = is_first
  468. self.in_features = in_features
  469. self.linear = nn.Linear(in_features, out_features, bias=bias)
  470. self.init_weights()
  471. def init_weights(self):
  472. with torch.no_grad():
  473. if self.is_first:
  474. self.linear.weight.uniform_(-1 / self.in_features, 1 / self.in_features)
  475. else:
  476. bound = np.sqrt(6 / self.in_features) / self.omega_0
  477. self.linear.weight.uniform_(-bound, bound)
  478. def forward(self, x):
  479. return torch.sin(self.omega_0 * self.linear(x))
  480. class SirenNet(nn.Module):
  481. """Standard SIREN network (Sitzmann et al., NeurIPS 2020).
  482. num_layers counts TOTAL layers (first SIREN + hidden SIREN + final linear).
  483. """
  484. def __init__(self, input_size: int, output_size: int = 1,
  485. hidden_size: int = 512, num_layers: int = 4,
  486. first_omega_0: float = 30.0, hidden_omega_0: float = 30.0):
  487. super().__init__()
  488. assert num_layers >= 2, "num_layers must be >= 2"
  489. layers = []
  490. layers.append(SirenLayer(input_size, hidden_size,
  491. is_first=True, omega_0=first_omega_0))
  492. for _ in range(num_layers - 2):
  493. layers.append(SirenLayer(hidden_size, hidden_size,
  494. omega_0=hidden_omega_0))
  495. final_linear = nn.Linear(hidden_size, output_size)
  496. with torch.no_grad():
  497. bound = np.sqrt(6 / hidden_size) / hidden_omega_0
  498. final_linear.weight.uniform_(-bound, bound)
  499. layers.append(final_linear)
  500. self.net = nn.Sequential(*layers)
  501. def forward(self, x):
  502. return self.net(x)
  503. class IVIMNet(nn.Module):
  504. """Network for IVIM parameter estimation."""
  505. def __init__(self, backbone):
  506. super().__init__()
  507. self.backbone = backbone
  508. # Extract the output size from the last linear layer
  509. last_layer = backbone.net[-1]
  510. in_features = last_layer.in_features
  511. self.head = nn.Linear(in_features, 3)
  512. def forward(self, x):
  513. features = self.backbone.net[:-1](x) # All layers except last
  514. params = self.head(features)
  515. # Apply constraints to IVIM parameters using hardcoded ranges
  516. Dp = torch.sigmoid(params[:, 0:1]) * (IVIM_PARAMETER_RANGES['Dp']['max'] - IVIM_PARAMETER_RANGES['Dp']['min']) + IVIM_PARAMETER_RANGES['Dp']['min']
  517. Dt = torch.sigmoid(params[:, 1:2]) * (IVIM_PARAMETER_RANGES['Dt']['max'] - IVIM_PARAMETER_RANGES['Dt']['min']) + IVIM_PARAMETER_RANGES['Dt']['min']
  518. Fp = torch.sigmoid(params[:, 2:3]) * (IVIM_PARAMETER_RANGES['Fp']['max'] - IVIM_PARAMETER_RANGES['Fp']['min']) + IVIM_PARAMETER_RANGES['Fp']['min']
  519. return Dp, Dt, Fp
  520. class IVIMNetSoftConstraint(nn.Module):
  521. """Network for IVIM parameter estimation with soft b0 constraint."""
  522. def __init__(self, backbone, delta_min=-0.1, delta_max=0.1):
  523. super().__init__()
  524. self.backbone = backbone
  525. self.delta_min = delta_min
  526. self.delta_max = delta_max
  527. # Extract the output size from the last linear layer
  528. last_layer = backbone.net[-1]
  529. in_features = last_layer.in_features
  530. self.head = nn.Linear(in_features, 4) # Output 4 parameters: Dp, Dt, Fp, delta
  531. def forward(self, x):
  532. features = self.backbone.net[:-1](x) # All layers except last
  533. params = self.head(features)
  534. # Apply constraints to IVIM parameters using hardcoded ranges
  535. Dp = torch.sigmoid(params[:, 0:1]) * (IVIM_PARAMETER_RANGES['Dp']['max'] - IVIM_PARAMETER_RANGES['Dp']['min']) + IVIM_PARAMETER_RANGES['Dp']['min']
  536. Dt = torch.sigmoid(params[:, 1:2]) * (IVIM_PARAMETER_RANGES['Dt']['max'] - IVIM_PARAMETER_RANGES['Dt']['min']) + IVIM_PARAMETER_RANGES['Dt']['min']
  537. Fp = torch.sigmoid(params[:, 2:3]) * (IVIM_PARAMETER_RANGES['Fp']['max'] - IVIM_PARAMETER_RANGES['Fp']['min']) + IVIM_PARAMETER_RANGES['Fp']['min']
  538. delta = torch.tanh(params[:, 3:4]) * self.delta_max # tanh to get symmetric range [-delta_max, delta_max]
  539. return Dp, Dt, Fp, delta
  540. # ADCNet class removed - only IVIM fitting is performed
  541. # ==================== Signal Models ====================
  542. def compute_ivim_signal(Dp, Dt, Fp, b_vals):
  543. """Compute IVIM signal: Fp·exp(-b·Dp) + (1-Fp)·exp(-b·Dt)"""
  544. return Fp * torch.exp(-b_vals * Dp) + (1 - Fp) * torch.exp(-b_vals * Dt)
  545. def compute_ivim_signal_soft(Dp, Dt, Fp, delta, b_vals):
  546. """Compute IVIM signal with soft constraint: [Fp·exp(-b·Dp) + (1-Fp)·exp(-b·Dt)] + δ"""
  547. signal = Fp * torch.exp(-b_vals * Dp) + (1 - Fp) * torch.exp(-b_vals * Dt)
  548. return signal + delta
  549. # compute_adc_signal function removed - only IVIM fitting is performed
  550. # ==================== Parameter Fitting Functions ====================
  551. def fit_ivim_parameters(dwi_data: np.ndarray, mask: np.ndarray, b_values: np.ndarray,
  552. args, device: torch.device) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
  553. """
  554. Fit IVIM parameters using spatial SIREN network.
  555. Args:
  556. dwi_data: DWI data array [X, Y, Z, B]
  557. mask: Binary mask for fitting
  558. b_values: Array of b-values
  559. args: Command line arguments
  560. device: PyTorch device
  561. Returns:
  562. Dp_map, Dt_map, Fp_map: IVIM parameter maps
  563. """
  564. print(f"\n{'='*60}")
  565. print(f"IVIM Parameter Fitting")
  566. print(f"{'='*60}")
  567. # Filter b-values for IVIM based on hardcoded thresholds
  568. # Use b <= b_value_threshold and exclude b < min_b_value_threshold (except b=0)
  569. b_value_threshold = 1000 # Default threshold
  570. min_b_value_threshold = 10 # Default minimum threshold
  571. ivim_indices = np.where(
  572. (b_values <= b_value_threshold) &
  573. ((b_values >= min_b_value_threshold) | (b_values == 0))
  574. )[0]
  575. ivim_b_values = b_values[ivim_indices]
  576. ivim_dwi_data = dwi_data[:, :, :, ivim_indices]
  577. print(f" Using b-values <= {b_value_threshold} and >= {min_b_value_threshold} (except b=0) for IVIM fitting")
  578. print(f" Selected b-values: {ivim_b_values}")
  579. print(f" Number of b-values for IVIM: {len(ivim_b_values)}")
  580. X_dim, Y_dim, Z_dim, num_b_ivim = ivim_dwi_data.shape
  581. # Normalize by b0
  582. b0_expanded = ivim_dwi_data[:, :, :, 0:1]
  583. dwi_norm = np.clip(ivim_dwi_data / (b0_expanded + 1e-6), 0, 1)
  584. # Normalize B0 by its own maximum for use as additional input
  585. if args.ivim_use_b0_signal:
  586. b0_max = np.max(ivim_dwi_data[:, :, :, 0])
  587. b0_norm_by_max = ivim_dwi_data[:, :, :, 0] / (b0_max + 1e-6)
  588. print(f" B0 max value: {b0_max:.2f}, will use normalized B0 as additional input")
  589. # Prepare B0 for patch extraction if patch size > 1
  590. if args.ivim_b0_patch_size > 1:
  591. print(f" Using B0 patch size: {args.ivim_b0_patch_size}x{args.ivim_b0_patch_size}x{args.ivim_b0_patch_size}")
  592. # Convert to tensor and pad
  593. b0_tensor = torch.from_numpy(b0_norm_by_max).float().unsqueeze(0).unsqueeze(0)
  594. pad_size = args.ivim_b0_patch_size // 2
  595. b0_padded = F.pad(b0_tensor, (pad_size, pad_size, pad_size, pad_size, pad_size, pad_size),
  596. mode='constant', value=0)
  597. b0_padded = b0_padded.to(device)[0, 0] # Remove batch and channel dims
  598. # Create patch coordinates for B0
  599. patch_offsets = list(range(-pad_size, pad_size + 1))
  600. b0_patch_coords = np.array([(x, y, z) for x in patch_offsets for y in patch_offsets for z in patch_offsets])
  601. b0_patch_coords = b0_patch_coords + pad_size # Adjust for padding
  602. b0_patch_coords = torch.tensor(b0_patch_coords, dtype=torch.float32).to(device)
  603. b0_patch_size = len(b0_patch_coords)
  604. else:
  605. b0_patch_size = 1
  606. # Create coordinate grids
  607. X = np.arange(0, X_dim, 1)
  608. Y = np.arange(0, Y_dim, 1)
  609. Z = np.arange(0, Z_dim, 1)
  610. points = np.meshgrid(X, Y, Z, indexing='ij')
  611. points = np.stack(points).transpose(1, 2, 3, 0).reshape(-1, 3).astype(np.float32)
  612. # Normalize coordinates to [-1, 1]
  613. points_normalized = np.empty_like(points, dtype=np.float32)
  614. points_normalized[:, 0] = 2.0 * points[:, 0] / (X_dim - 1) - 1.0
  615. points_normalized[:, 1] = 2.0 * points[:, 1] / (Y_dim - 1) - 1.0
  616. points_normalized[:, 2] = 2.0 * points[:, 2] / (Z_dim - 1) - 1.0
  617. # Create IVIM model
  618. # Determine input size based on additional inputs
  619. input_size = 3 # Start with 3D coordinates
  620. if args.ivim_use_dwi_signals:
  621. input_size += num_b_ivim # Add normalized DWI signals
  622. print(f" Using DWI signals as additional input")
  623. if args.ivim_use_b0_signal:
  624. if args.ivim_b0_patch_size > 1:
  625. input_size += b0_patch_size # Add B0 patch features
  626. print(f" Using B0 patch features as additional input (size: {b0_patch_size})")
  627. else:
  628. input_size += 1 # Add single B0 value
  629. print(f" Using normalized B0 signal as additional input")
  630. print(f" Total input size: {input_size}")
  631. backbone = SirenNet(
  632. input_size=input_size,
  633. output_size=args.ivim_hidden_size,
  634. hidden_size=args.ivim_hidden_size,
  635. num_layers=args.ivim_num_layers,
  636. first_omega_0=args.ivim_first_omega_0,
  637. hidden_omega_0=args.ivim_hidden_omega_0,
  638. )
  639. model = IVIMNet(backbone).to(device)
  640. # Prepare data and move to GPU
  641. mask_flat = mask.reshape(-1)
  642. print(f" Mask voxels for fitting: {mask_flat.sum():,}")
  643. # Move masked data to GPU directly
  644. coords_tensor = torch.tensor(points_normalized[mask_flat], dtype=torch.float32).to(device)
  645. dwi_signals = torch.tensor(dwi_norm.reshape(-1, num_b_ivim)[mask_flat], dtype=torch.float32).to(device)
  646. if args.ivim_use_b0_signal:
  647. if args.ivim_b0_patch_size > 1:
  648. # Extract B0 patches for masked voxels
  649. print(f" Extracting B0 patches for {mask_flat.sum():,} masked voxels...")
  650. voxel_indices = np.where(mask_flat)[0]
  651. voxel_coords_full = points[voxel_indices]
  652. # Preallocate B0 patch features
  653. b0_patch_features = torch.zeros((len(voxel_coords_full), b0_patch_size), dtype=torch.float32, device=device)
  654. # Extract patches in batches
  655. batch_size = 5000
  656. for start_idx in range(0, len(voxel_coords_full), batch_size):
  657. end_idx = min(start_idx + batch_size, len(voxel_coords_full))
  658. batch_coords = voxel_coords_full[start_idx:end_idx]
  659. batch_coords_tensor = torch.tensor(batch_coords, dtype=torch.float32, device=device)
  660. # Extract B0 patches
  661. batch_coords_padded = batch_coords_tensor + pad_size
  662. coords_patch = batch_coords_padded[:, None, :] + b0_patch_coords[None, :, :] - pad_size
  663. x_idx = coords_patch[..., 0].reshape(-1).long()
  664. y_idx = coords_patch[..., 1].reshape(-1).long()
  665. z_idx = coords_patch[..., 2].reshape(-1).long()
  666. features = b0_padded[x_idx, y_idx, z_idx]
  667. features = features.reshape(len(batch_coords), -1)
  668. b0_patch_features[start_idx:end_idx, :] = features
  669. b0_signals = b0_patch_features
  670. else:
  671. b0_signals = torch.tensor(b0_norm_by_max.reshape(-1)[mask_flat], dtype=torch.float32).to(device).unsqueeze(1)
  672. # Exclude b0 from fitting (only use b>0)
  673. b_vals_t = torch.tensor(ivim_b_values[1:], device=device, dtype=torch.float32)
  674. # Training
  675. optimizer = torch.optim.Adam(model.parameters(), lr=args.ivim_lr)
  676. scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.ivim_epochs)
  677. criterion = nn.MSELoss()
  678. model.train()
  679. print(f"\nTraining for {args.ivim_epochs} epochs...")
  680. num_samples = len(coords_tensor)
  681. batch_size = args.ivim_batch_size
  682. for epoch in range(args.ivim_epochs):
  683. epoch_loss = 0.0
  684. epoch_start_time = time.time()
  685. # Shuffle indices for this epoch
  686. indices = torch.randperm(num_samples, device=device)
  687. # Manual batching
  688. for start_idx in range(0, num_samples, batch_size):
  689. end_idx = min(start_idx + batch_size, num_samples)
  690. batch_indices = indices[start_idx:end_idx]
  691. # Get batch data - already on GPU
  692. coords_batch = coords_tensor[batch_indices]
  693. signals_batch = dwi_signals[batch_indices]
  694. # Prepare model input
  695. model_input = coords_batch
  696. if args.ivim_use_dwi_signals:
  697. # Concatenate DWI signals
  698. model_input = torch.cat([model_input, signals_batch], dim=1)
  699. if args.ivim_use_b0_signal:
  700. # Concatenate B0 signal/patches
  701. b0_batch = b0_signals[batch_indices]
  702. model_input = torch.cat([model_input, b0_batch], dim=1)
  703. # Forward pass
  704. Dp, Dt, Fp = model(model_input)
  705. # Compute predicted signals (excluding b0)
  706. pred_signals = compute_ivim_signal(Dp, Dt, Fp, b_vals_t.unsqueeze(0))
  707. # Loss on normalized signals (excluding b0)
  708. loss = criterion(pred_signals, signals_batch[:, 1:])
  709. # Backward pass
  710. optimizer.zero_grad()
  711. loss.backward()
  712. torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
  713. optimizer.step()
  714. epoch_loss += loss.item() * batch_indices.size(0)
  715. epoch_loss /= num_samples
  716. scheduler.step()
  717. if (epoch + 1) % 10 == 0 or epoch == 0:
  718. epoch_time = time.time() - epoch_start_time
  719. print(f" Epoch {epoch+1}/{args.ivim_epochs}, Loss: {epoch_loss:.6f}, Time: {epoch_time:.2f}s")
  720. # Inference
  721. print("\nRunning inference...")
  722. model.eval()
  723. # Prepare full volume coordinates for inference
  724. coords_full = torch.tensor(points_normalized, dtype=torch.float32).to(device)
  725. nvox = coords_full.shape[0]
  726. step = 50000
  727. all_Dp, all_Dt, all_Fp = [], [], []
  728. with torch.no_grad():
  729. for i in range(0, nvox, step):
  730. coords_batch = coords_full[i:i+step]
  731. # Prepare model input for inference
  732. model_input = coords_batch
  733. if args.ivim_use_dwi_signals:
  734. # Get corresponding DWI signals for this batch
  735. dwi_signals_batch = torch.tensor(dwi_norm.reshape(-1, num_b_ivim)[i:i+step], dtype=torch.float32).to(device)
  736. model_input = torch.cat([model_input, dwi_signals_batch], dim=1)
  737. if args.ivim_use_b0_signal:
  738. if args.ivim_b0_patch_size > 1:
  739. # Extract B0 patches for inference batch
  740. batch_coords = points[i:i+step]
  741. batch_coords_tensor = torch.tensor(batch_coords, dtype=torch.float32, device=device)
  742. # Extract B0 patches
  743. batch_coords_padded = batch_coords_tensor + pad_size
  744. coords_patch = batch_coords_padded[:, None, :] + b0_patch_coords[None, :, :] - pad_size
  745. x_idx = coords_patch[..., 0].reshape(-1).long()
  746. y_idx = coords_patch[..., 1].reshape(-1).long()
  747. z_idx = coords_patch[..., 2].reshape(-1).long()
  748. features = b0_padded[x_idx, y_idx, z_idx]
  749. b0_signals_batch = features.reshape(coords_batch.shape[0], -1)
  750. else:
  751. # Get single B0 values
  752. b0_signals_batch = torch.tensor(b0_norm_by_max.reshape(-1)[i:i+step], dtype=torch.float32).to(device).unsqueeze(1)
  753. model_input = torch.cat([model_input, b0_signals_batch], dim=1)
  754. Dp, Dt, Fp = model(model_input)
  755. all_Dp.append(Dp.cpu().numpy())
  756. all_Dt.append(Dt.cpu().numpy())
  757. all_Fp.append(Fp.cpu().numpy())
  758. # Reshape to image dimensions and apply mask
  759. Dp_map = np.concatenate(all_Dp).reshape(X_dim, Y_dim, Z_dim) * mask
  760. Dt_map = np.concatenate(all_Dt).reshape(X_dim, Y_dim, Z_dim) * mask
  761. Fp_map = np.concatenate(all_Fp).reshape(X_dim, Y_dim, Z_dim) * mask
  762. print("✔ IVIM fitting completed")
  763. return Dp_map, Dt_map, Fp_map
  764. # fit_adc_parameters function removed - only IVIM fitting is performed
  765. def fit_ivim_parameters_soft(dwi_data: np.ndarray, mask: np.ndarray, b_values: np.ndarray,
  766. args, device: torch.device) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
  767. """
  768. Fit IVIM parameters using spatial SIREN network with soft b0 constraint.
  769. Args:
  770. dwi_data: DWI data array [X, Y, Z, B]
  771. mask: Binary mask for fitting
  772. b_values: Array of b-values
  773. args: Command line arguments
  774. device: PyTorch device
  775. fitting_config: Configuration object
  776. Returns:
  777. Dp_map, Dt_map, Fp_map, delta_map: IVIM parameter maps and delta offset map
  778. """
  779. print(f"\n{'='*60}")
  780. print(f"IVIM Parameter Fitting with Soft b0 Constraint")
  781. print(f"{'='*60}")
  782. # Filter b-values for IVIM based on hardcoded thresholds
  783. b_value_threshold = 1000 # Default threshold
  784. min_b_value_threshold = 10 # Default minimum threshold
  785. ivim_indices = np.where(
  786. (b_values <= b_value_threshold) &
  787. ((b_values >= min_b_value_threshold) | (b_values == 0))
  788. )[0]
  789. ivim_b_values = b_values[ivim_indices]
  790. ivim_dwi_data = dwi_data[:, :, :, ivim_indices]
  791. print(f" Using b-values <= {b_value_threshold} and >= {min_b_value_threshold} (except b=0) for IVIM fitting")
  792. print(f" Selected b-values: {ivim_b_values}")
  793. print(f" Number of b-values for IVIM: {len(ivim_b_values)}")
  794. X_dim, Y_dim, Z_dim, num_b_ivim = ivim_dwi_data.shape
  795. # Normalize by b0 (soft constraint approach)
  796. b0_expanded = ivim_dwi_data[:, :, :, 0:1]
  797. dwi_norm = np.clip(ivim_dwi_data / (b0_expanded + 1e-6), 0, 2) # Allow some range above 1
  798. print(f"\n DWI normalization for IVIM fitting:")
  799. print(f" B0 range: [{b0_expanded[mask].min():.2f}, {b0_expanded[mask].max():.2f}]")
  800. for i, b in enumerate(ivim_b_values):
  801. print(f" b={b}: normalized range = [{dwi_norm[:,:,:,i][mask].min():.4f}, {dwi_norm[:,:,:,i][mask].max():.4f}]")
  802. # Normalize B0 by its own maximum for use as additional input
  803. if args.ivim_use_b0_signal:
  804. b0_max = np.max(ivim_dwi_data[:, :, :, 0])
  805. b0_norm_by_max = ivim_dwi_data[:, :, :, 0] / (b0_max + 1e-6)
  806. print(f" B0 max value: {b0_max:.2f}, will use normalized B0 as additional input")
  807. # Prepare B0 for patch extraction if patch size > 1
  808. if args.ivim_b0_patch_size > 1:
  809. print(f" Using B0 patch size: {args.ivim_b0_patch_size}x{args.ivim_b0_patch_size}x{args.ivim_b0_patch_size}")
  810. # Convert to tensor and pad
  811. b0_tensor = torch.from_numpy(b0_norm_by_max).float().unsqueeze(0).unsqueeze(0)
  812. pad_size = args.ivim_b0_patch_size // 2
  813. b0_padded = F.pad(b0_tensor, (pad_size, pad_size, pad_size, pad_size, pad_size, pad_size),
  814. mode='constant', value=0)
  815. b0_padded = b0_padded.to(device)[0, 0] # Remove batch and channel dims
  816. # Create patch coordinates for B0
  817. patch_offsets = list(range(-pad_size, pad_size + 1))
  818. b0_patch_coords = np.array([(x, y, z) for x in patch_offsets for y in patch_offsets for z in patch_offsets])
  819. b0_patch_coords = b0_patch_coords + pad_size # Adjust for padding
  820. b0_patch_coords = torch.tensor(b0_patch_coords, dtype=torch.float32).to(device)
  821. b0_patch_size = len(b0_patch_coords)
  822. else:
  823. b0_patch_size = 1
  824. # Create coordinate grids
  825. X = np.arange(0, X_dim, 1)
  826. Y = np.arange(0, Y_dim, 1)
  827. Z = np.arange(0, Z_dim, 1)
  828. points = np.meshgrid(X, Y, Z, indexing='ij')
  829. points = np.stack(points).transpose(1, 2, 3, 0).reshape(-1, 3).astype(np.float32)
  830. # Normalize coordinates to [-1, 1]
  831. points_normalized = np.empty_like(points, dtype=np.float32)
  832. points_normalized[:, 0] = 2.0 * points[:, 0] / (X_dim - 1) - 1.0
  833. points_normalized[:, 1] = 2.0 * points[:, 1] / (Y_dim - 1) - 1.0
  834. points_normalized[:, 2] = 2.0 * points[:, 2] / (Z_dim - 1) - 1.0
  835. # Create IVIM model with soft constraint
  836. # Determine input size based on additional inputs
  837. input_size = 3 # Start with 3D coordinates
  838. if args.ivim_use_dwi_signals:
  839. input_size += num_b_ivim # Add normalized DWI signals
  840. print(f" Using DWI signals as additional input")
  841. if args.ivim_use_b0_signal:
  842. if args.ivim_b0_patch_size > 1:
  843. input_size += b0_patch_size # Add B0 patch features
  844. print(f" Using B0 patch features as additional input (size: {b0_patch_size})")
  845. else:
  846. input_size += 1 # Add single B0 value
  847. print(f" Using normalized B0 signal as additional input")
  848. print(f" Total input size: {input_size}")
  849. backbone = SirenNet(
  850. input_size=input_size,
  851. output_size=args.ivim_hidden_size,
  852. hidden_size=args.ivim_hidden_size,
  853. num_layers=args.ivim_num_layers,
  854. first_omega_0=args.ivim_first_omega_0,
  855. hidden_omega_0=args.ivim_hidden_omega_0,
  856. )
  857. # Use soft constraint model
  858. model = IVIMNetSoftConstraint(
  859. backbone,
  860. delta_min=args.delta_min,
  861. delta_max=args.delta_max
  862. ).to(device)
  863. print(f"\n Delta (offset) constraint: [{args.delta_min:.2f}, {args.delta_max:.2f}]")
  864. print(f" This allows the curve to deviate from exactly passing through (0,1)")
  865. print(f" Delta regularization weight: {args.delta_regularization}")
  866. # Prepare data and move to GPU
  867. mask_flat = mask.reshape(-1)
  868. print(f" Mask voxels for fitting: {mask_flat.sum():,}")
  869. # Move masked data to GPU directly
  870. coords_tensor = torch.tensor(points_normalized[mask_flat], dtype=torch.float32).to(device)
  871. dwi_signals = torch.tensor(dwi_norm.reshape(-1, num_b_ivim)[mask_flat], dtype=torch.float32).to(device)
  872. if args.ivim_use_b0_signal:
  873. if args.ivim_b0_patch_size > 1:
  874. # Extract B0 patches for masked voxels
  875. print(f" Extracting B0 patches for {mask_flat.sum():,} masked voxels...")
  876. voxel_indices = np.where(mask_flat)[0]
  877. voxel_coords_full = points[voxel_indices]
  878. # Preallocate B0 patch features
  879. b0_patch_features = torch.zeros((len(voxel_coords_full), b0_patch_size), dtype=torch.float32, device=device)
  880. # Extract patches in batches
  881. batch_size = 5000
  882. for start_idx in range(0, len(voxel_coords_full), batch_size):
  883. end_idx = min(start_idx + batch_size, len(voxel_coords_full))
  884. batch_coords = voxel_coords_full[start_idx:end_idx]
  885. batch_coords_tensor = torch.tensor(batch_coords, dtype=torch.float32, device=device)
  886. # Extract B0 patches
  887. batch_coords_padded = batch_coords_tensor + pad_size
  888. coords_patch = batch_coords_padded[:, None, :] + b0_patch_coords[None, :, :] - pad_size
  889. x_idx = coords_patch[..., 0].reshape(-1).long()
  890. y_idx = coords_patch[..., 1].reshape(-1).long()
  891. z_idx = coords_patch[..., 2].reshape(-1).long()
  892. features = b0_padded[x_idx, y_idx, z_idx]
  893. features = features.reshape(len(batch_coords), -1)
  894. b0_patch_features[start_idx:end_idx, :] = features
  895. b0_signals = b0_patch_features
  896. else:
  897. b0_signals = torch.tensor(b0_norm_by_max.reshape(-1)[mask_flat], dtype=torch.float32).to(device).unsqueeze(1)
  898. # Include ALL b-values for fitting with soft constraint (including b0)
  899. b_vals_t = torch.tensor(ivim_b_values, device=device, dtype=torch.float32)
  900. # Training
  901. optimizer = torch.optim.Adam(model.parameters(), lr=args.ivim_lr)
  902. scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.ivim_epochs)
  903. model.train()
  904. print(f"\nTraining for {args.ivim_epochs} epochs...")
  905. num_samples = len(coords_tensor)
  906. batch_size = args.ivim_batch_size
  907. for epoch in range(args.ivim_epochs):
  908. epoch_loss = 0.0
  909. epoch_start_time = time.time()
  910. # Shuffle indices for this epoch
  911. indices = torch.randperm(num_samples, device=device)
  912. # Manual batching
  913. for start_idx in range(0, num_samples, batch_size):
  914. end_idx = min(start_idx + batch_size, num_samples)
  915. batch_indices = indices[start_idx:end_idx]
  916. # Get batch data - already on GPU
  917. coords_batch = coords_tensor[batch_indices]
  918. signals_batch = dwi_signals[batch_indices]
  919. # Prepare model input
  920. model_input = coords_batch
  921. if args.ivim_use_dwi_signals:
  922. # Concatenate DWI signals
  923. model_input = torch.cat([model_input, signals_batch], dim=1)
  924. if args.ivim_use_b0_signal:
  925. # Concatenate B0 signal/patches
  926. b0_batch = b0_signals[batch_indices]
  927. model_input = torch.cat([model_input, b0_batch], dim=1)
  928. # Forward pass
  929. Dp, Dt, Fp, delta = model(model_input)
  930. # Compute predicted signals (including b0 with soft constraint)
  931. pred_signals = compute_ivim_signal_soft(Dp, Dt, Fp, delta, b_vals_t.unsqueeze(0))
  932. # Signal fitting loss
  933. signal_loss = F.mse_loss(pred_signals, signals_batch)
  934. # Delta regularization loss (L2 penalty on offset)
  935. delta_reg_loss = args.delta_regularization * (delta**2).mean()
  936. # Total loss
  937. loss = signal_loss + delta_reg_loss
  938. # Backward pass
  939. optimizer.zero_grad()
  940. loss.backward()
  941. torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
  942. optimizer.step()
  943. epoch_loss += loss.item() * batch_indices.size(0)
  944. epoch_loss /= num_samples
  945. scheduler.step()
  946. if (epoch + 1) % 10 == 0 or epoch == 0:
  947. epoch_time = time.time() - epoch_start_time
  948. print(f" Epoch {epoch+1}/{args.ivim_epochs}, Loss: {epoch_loss:.6f}, Time: {epoch_time:.2f}s")
  949. # Inference
  950. print("\nRunning inference...")
  951. model.eval()
  952. # Prepare full volume coordinates for inference
  953. coords_full = torch.tensor(points_normalized, dtype=torch.float32).to(device)
  954. nvox = coords_full.shape[0]
  955. step = 50000
  956. all_Dp, all_Dt, all_Fp, all_delta = [], [], [], []
  957. with torch.no_grad():
  958. for i in range(0, nvox, step):
  959. coords_batch = coords_full[i:i+step]
  960. # Prepare model input for inference
  961. model_input = coords_batch
  962. if args.ivim_use_dwi_signals:
  963. # Get corresponding DWI signals for this batch
  964. dwi_signals_batch = torch.tensor(dwi_norm.reshape(-1, num_b_ivim)[i:i+step], dtype=torch.float32).to(device)
  965. model_input = torch.cat([model_input, dwi_signals_batch], dim=1)
  966. if args.ivim_use_b0_signal:
  967. if args.ivim_b0_patch_size > 1:
  968. # Extract B0 patches for inference batch
  969. batch_coords = points[i:i+step]
  970. batch_coords_tensor = torch.tensor(batch_coords, dtype=torch.float32, device=device)
  971. # Extract B0 patches
  972. batch_coords_padded = batch_coords_tensor + pad_size
  973. coords_patch = batch_coords_padded[:, None, :] + b0_patch_coords[None, :, :] - pad_size
  974. x_idx = coords_patch[..., 0].reshape(-1).long()
  975. y_idx = coords_patch[..., 1].reshape(-1).long()
  976. z_idx = coords_patch[..., 2].reshape(-1).long()
  977. features = b0_padded[x_idx, y_idx, z_idx]
  978. b0_signals_batch = features.reshape(coords_batch.shape[0], -1)
  979. else:
  980. # Get single B0 values
  981. b0_signals_batch = torch.tensor(b0_norm_by_max.reshape(-1)[i:i+step], dtype=torch.float32).to(device).unsqueeze(1)
  982. model_input = torch.cat([model_input, b0_signals_batch], dim=1)
  983. Dp, Dt, Fp, delta = model(model_input)
  984. all_Dp.append(Dp.cpu().numpy())
  985. all_Dt.append(Dt.cpu().numpy())
  986. all_Fp.append(Fp.cpu().numpy())
  987. all_delta.append(delta.cpu().numpy())
  988. # Reshape to image dimensions and apply mask
  989. Dp_map = np.concatenate(all_Dp).reshape(X_dim, Y_dim, Z_dim) * mask
  990. Dt_map = np.concatenate(all_Dt).reshape(X_dim, Y_dim, Z_dim) * mask
  991. Fp_map = np.concatenate(all_Fp).reshape(X_dim, Y_dim, Z_dim) * mask
  992. delta_map = np.concatenate(all_delta).reshape(X_dim, Y_dim, Z_dim) * mask
  993. print("\n IVIM parameter statistics:")
  994. print(f" Dp range: [{Dp_map[mask].min():.6f}, {Dp_map[mask].max():.6f}] mm²/s")
  995. print(f" Dt range: [{Dt_map[mask].min():.6f}, {Dt_map[mask].max():.6f}] mm²/s")
  996. print(f" Fp range: [{Fp_map[mask].min():.4f}, {Fp_map[mask].max():.4f}]")
  997. print(f" Delta range: [{delta_map[mask].min():.4f}, {delta_map[mask].max():.4f}]")
  998. print(f" Mean delta: {delta_map[mask].mean():.4f} (mean absolute deviation: {np.abs(delta_map[mask]).mean():.4f})")
  999. print("✔ IVIM fitting with soft b0 constraint completed")
  1000. return Dp_map, Dt_map, Fp_map, delta_map
  1001. # ==================== Main Processing Function ====================
  1002. def load_ivim_data(args):
  1003. """Load IVIM data from specified directory."""
  1004. print(f"\n{'#'*80}")
  1005. print(f"Loading IVIM data")
  1006. print(f"{'#'*80}")
  1007. # Build paths for IVIM data files
  1008. data_dir = args.data_dir if args.data_dir else args.input_dir
  1009. # Check if we should load noisy or ground truth data
  1010. if args.use_noisy_data:
  1011. data_subdir = os.path.join(data_dir, 'noisy')
  1012. print(f"Loading noisy IVIM data from: {data_subdir}")
  1013. else:
  1014. data_subdir = os.path.join(data_dir, 'ground_truth')
  1015. print(f"Loading ground truth IVIM data from: {data_subdir}")
  1016. # If subdirectory doesn't exist, fall back to main directory
  1017. if not os.path.exists(data_subdir):
  1018. data_subdir = data_dir
  1019. print(f"Subdirectory not found, using main directory: {data_dir}")
  1020. # Define b-values and corresponding files
  1021. b_values = [0, 50, 100, 200, 400, 600, 800, 1000]
  1022. dwi_files = {}
  1023. for b_val in b_values:
  1024. filename = f'dwi_ivim_b{b_val}.nii.gz'
  1025. filepath = os.path.join(data_subdir, filename)
  1026. if os.path.exists(filepath):
  1027. dwi_files[b_val] = filepath
  1028. else:
  1029. print(f" ⚠ File not found: {filepath}")
  1030. if not dwi_files:
  1031. print(f" ⚠ No IVIM files found in {data_subdir}")
  1032. return None, None, None
  1033. # Sort by b-value
  1034. b_values = sorted(dwi_files.keys())
  1035. print(f"Found {len(b_values)} b-values: {b_values}")
  1036. # Load DWI data
  1037. dwi_data = []
  1038. for b_val in b_values:
  1039. img = nib.load(dwi_files[b_val])
  1040. dwi_data.append(img.get_fdata().astype(np.float32))
  1041. affine = img.affine
  1042. dwi_data = np.stack(dwi_data, axis=-1) # Shape: [X, Y, Z, B]
  1043. b_values_array = np.array(b_values, dtype=np.float32)
  1044. print(f"DWI data shape: {dwi_data.shape}")
  1045. return dwi_data, b_values_array, affine
  1046. def process_ivim_fitting(args):
  1047. """Process IVIM parameter fitting."""
  1048. print(f"\n{'#'*80}")
  1049. print(f"IVIM parameter fitting")
  1050. print(f"{'#'*80}")
  1051. # This function has been replaced by process_ivim
  1052. # ADC fitting removed - only IVIM fitting is performed
  1053. def process_ivim(args):
  1054. """Process IVIM data (both loading and fitting)."""
  1055. # Validate output directory and files before processing
  1056. print(f"\nValidating output directory and file permissions...")
  1057. # Check if output directory exists, create if needed
  1058. try:
  1059. os.makedirs(args.output_dir, exist_ok=True)
  1060. except Exception as e:
  1061. sys.exit(f"Error: Cannot create output directory {args.output_dir}: {e}")
  1062. # Check write permissions
  1063. test_file = os.path.join(args.output_dir, '.test_write_permission')
  1064. try:
  1065. with open(test_file, 'w') as f:
  1066. f.write('test')
  1067. os.remove(test_file)
  1068. except Exception as e:
  1069. sys.exit(f"Error: Cannot write to output directory {args.output_dir}: {e}")
  1070. # Define expected output files
  1071. if args.use_soft_constraint:
  1072. output_files = ['Dp.nii.gz', 'Dt.nii.gz', 'Fp.nii.gz', 'delta.nii.gz', 'fitting_mask.nii.gz', 'parameter_config.json']
  1073. else:
  1074. output_files = ['Dp.nii.gz', 'Dt.nii.gz', 'Fp.nii.gz', 'fitting_mask.nii.gz', 'parameter_config.json']
  1075. # Check if output files already exist
  1076. existing_files = []
  1077. for filename in output_files:
  1078. filepath = os.path.join(args.output_dir, filename)
  1079. if os.path.exists(filepath):
  1080. existing_files.append(filename)
  1081. if existing_files and not args.force:
  1082. print(f"\n⚠ Warning: The following output files already exist:")
  1083. for f in existing_files:
  1084. print(f" - {f}")
  1085. return
  1086. # Check subdirectories if needed
  1087. if args.use_noisy_data and not args.skip_denoise:
  1088. denoised_dir = os.path.join(args.output_dir, 'denoised')
  1089. try:
  1090. os.makedirs(denoised_dir, exist_ok=True)
  1091. except Exception as e:
  1092. sys.exit(f"Error: Cannot create denoised directory {denoised_dir}: {e}")
  1093. if args.apply_noise_correction and args.use_noisy_data:
  1094. noise_corrected_dir = os.path.join(args.output_dir, 'noise_corrected')
  1095. try:
  1096. os.makedirs(noise_corrected_dir, exist_ok=True)
  1097. except Exception as e:
  1098. sys.exit(f"Error: Cannot create noise corrected directory {noise_corrected_dir}: {e}")
  1099. print("✔ Output directory validation completed\n")
  1100. # Load IVIM data
  1101. dwi_data, b_values_array, affine = load_ivim_data(args)
  1102. if dwi_data is None:
  1103. return
  1104. # Create fitting mask from b0
  1105. b0_data = dwi_data[:, :, :, 0]
  1106. # Load brain mask directly from data_1 directory (or the specific data directory being used)
  1107. # Use the input_dir or data_dir to construct mask path
  1108. base_data_dir = args.data_dir if args.data_dir else args.input_dir
  1109. # Extract the base directory (data_1, data_2, or data_3)
  1110. mask_path = os.path.join(os.path.dirname(base_data_dir) if 'SNR_' in base_data_dir or 'noisy' in base_data_dir or 'ground_truth' in base_data_dir else base_data_dir, 'brain_mask.nii.gz')
  1111. # Try multiple possible mask locations
  1112. possible_mask_paths = [
  1113. mask_path,
  1114. os.path.join(base_data_dir, 'brain_mask.nii.gz'),
  1115. os.path.join(os.path.dirname(os.path.abspath(__file__)), 'data_1', 'brain_mask.nii.gz'),
  1116. ]
  1117. fitting_mask = None
  1118. for mp in possible_mask_paths:
  1119. if os.path.exists(mp):
  1120. mask_img = nib.load(mp)
  1121. fitting_mask = mask_img.get_fdata().astype(bool)
  1122. print(f"\nLoaded brain mask from: {mp}")
  1123. break
  1124. if fitting_mask is None:
  1125. # Fallback to creating mask from b0 data
  1126. fitting_mask = b0_data > args.background_threshold
  1127. print(f"\nWarning: Brain mask not found, using b0 threshold: {args.background_threshold}")
  1128. print(f"\nFitting mask coverage: {fitting_mask.sum() / fitting_mask.size * 100:.1f}%")
  1129. # Apply noise correction first if requested and using noisy data
  1130. if args.apply_noise_correction and args.use_noisy_data:
  1131. print("\nApplying noise floor correction to noisy data...")
  1132. # Detect noise floor regions from original b0
  1133. background_mask, noise_floor_mask, tissue_mask = detect_noise_floor_mask(b0_data)
  1134. # Apply Rician noise correction to original noisy data
  1135. dwi_data_corrected, estimated_sigma = rician_noise_correction(
  1136. dwi_data, noise_floor_mask, method=args.noise_correction_method
  1137. )
  1138. # Save noise corrected DWI data to separate folder
  1139. noise_corrected_dir = os.path.join(args.output_dir, 'noise_corrected')
  1140. os.makedirs(noise_corrected_dir, exist_ok=True)
  1141. # Save the noise-corrected data
  1142. for i, b_val in enumerate(b_values_array):
  1143. nib.save(nib.Nifti1Image(dwi_data_corrected[:, :, :, i].astype(np.float32), affine),
  1144. os.path.join(noise_corrected_dir, f'dwi_ivim_b{int(b_val)}_noise_corrected.nii.gz'))
  1145. # Use noise corrected data for subsequent processing
  1146. dwi_data = dwi_data_corrected
  1147. b0_data = dwi_data[:, :, :, 0]
  1148. print(f"\nNoise correction completed. Estimated sigma: {estimated_sigma:.4f}")
  1149. # Apply denoising after noise correction if requested and using noisy data
  1150. if args.use_noisy_data and not args.skip_denoise:
  1151. print("\nChecking for existing denoised data...")
  1152. device = torch.device(f'cuda:{args.gpu}' if torch.cuda.is_available() else 'cpu')
  1153. # Create denoised output directory
  1154. denoised_output_dir = os.path.join(args.output_dir, 'denoised')
  1155. os.makedirs(denoised_output_dir, exist_ok=True)
  1156. # Check if all denoised files exist
  1157. all_denoised_exist = True
  1158. denoised_data = []
  1159. for i, b_val in enumerate(b_values_array):
  1160. denoised_path = os.path.join(denoised_output_dir, f'dwi_ivim_b{int(b_val)}_denoised.nii.gz')
  1161. if os.path.exists(denoised_path):
  1162. print(f" ✔ Found existing denoised file for b={int(b_val)}")
  1163. # Load the existing denoised data
  1164. denoised_img = nib.load(denoised_path)
  1165. denoised_data.append(denoised_img.get_fdata().astype(np.float32))
  1166. else:
  1167. all_denoised_exist = False
  1168. break
  1169. if all_denoised_exist:
  1170. # All denoised files exist, use them
  1171. print(" ✔ All denoised files found, skipping denoising step")
  1172. dwi_data = np.stack(denoised_data, axis=-1)
  1173. else:
  1174. # Some denoised files missing, run denoising
  1175. print(" Some denoised files missing, running MDINR denoising...")
  1176. dwi_data = denoise_dwi_mdinr_optimized(dwi_data, fitting_mask, b_values_array, args, device,
  1177. affine=affine, save_intermediate=True,
  1178. output_dir=denoised_output_dir)
  1179. print("\nDenoising completed.")
  1180. # Update b0_data with denoised data
  1181. b0_data = dwi_data[:, :, :, 0]
  1182. elif args.use_noisy_data and args.skip_denoise:
  1183. print("\nUsing noisy data without denoising.")
  1184. else:
  1185. if not args.use_noisy_data:
  1186. print("\nUsing ground truth data (no noise correction or denoising needed).")
  1187. # Set device
  1188. device = torch.device(f'cuda:{args.gpu}' if torch.cuda.is_available() else 'cpu')
  1189. print(f"Using device: {device}")
  1190. # Configuration is now hardcoded in the script
  1191. # No need to load external config files
  1192. # Ensure output directory exists
  1193. os.makedirs(args.output_dir, exist_ok=True)
  1194. # Fit IVIM parameters
  1195. if args.use_soft_constraint:
  1196. Dp_map, Dt_map, Fp_map, delta_map = fit_ivim_parameters_soft(
  1197. dwi_data, fitting_mask, b_values_array, args, device
  1198. )
  1199. # Save IVIM parameters
  1200. nib.save(nib.Nifti1Image(Dp_map, affine), os.path.join(args.output_dir, 'Dp.nii.gz'))
  1201. nib.save(nib.Nifti1Image(Dt_map, affine), os.path.join(args.output_dir, 'Dt.nii.gz'))
  1202. nib.save(nib.Nifti1Image(Fp_map, affine), os.path.join(args.output_dir, 'Fp.nii.gz'))
  1203. nib.save(nib.Nifti1Image(delta_map, affine), os.path.join(args.output_dir, 'delta.nii.gz'))
  1204. print(f" ✔ Saved IVIM parameters with delta to: {args.output_dir}")
  1205. else:
  1206. Dp_map, Dt_map, Fp_map = fit_ivim_parameters(
  1207. dwi_data, fitting_mask, b_values_array, args, device
  1208. )
  1209. # Save IVIM parameters
  1210. nib.save(nib.Nifti1Image(Dp_map, affine), os.path.join(args.output_dir, 'Dp.nii.gz'))
  1211. nib.save(nib.Nifti1Image(Dt_map, affine), os.path.join(args.output_dir, 'Dt.nii.gz'))
  1212. nib.save(nib.Nifti1Image(Fp_map, affine), os.path.join(args.output_dir, 'Fp.nii.gz'))
  1213. print(f" ✔ Saved IVIM parameters to: {args.output_dir}")
  1214. # Save fitting mask
  1215. nib.save(nib.Nifti1Image(fitting_mask.astype(np.float32), affine),
  1216. os.path.join(args.output_dir, 'fitting_mask.nii.gz'))
  1217. # Save parameter configuration as JSON
  1218. config_info = {
  1219. 'IVIM_parameters': IVIM_PARAMETER_RANGES,
  1220. 'b_value_threshold': 1000,
  1221. 'min_b_value_threshold': 10,
  1222. 'soft_constraint': args.use_soft_constraint,
  1223. 'delta_min': args.delta_min if args.use_soft_constraint else None,
  1224. 'delta_max': args.delta_max if args.use_soft_constraint else None,
  1225. 'delta_regularization': args.delta_regularization if args.use_soft_constraint else None
  1226. }
  1227. config_path = os.path.join(args.output_dir, 'parameter_config.json')
  1228. with open(config_path, 'w') as f:
  1229. json.dump(config_info, f, indent=4)
  1230. print(f" ✔ Saved parameter configuration to: {config_path}")
  1231. print(f"\n✔ Completed IVIM fitting")
  1232. # process_patient function removed - direct IVIM processing only
  1233. # Removed - functionality merged into process_ivim
  1234. def main():
  1235. parser = argparse.ArgumentParser(description='IVIM Parameter Fitting Pipeline')
  1236. # Input/Output arguments
  1237. _this_dir = os.path.dirname(os.path.abspath(__file__))
  1238. parser.add_argument('--input_dir', type=str,
  1239. default=os.path.join(_this_dir, 'data_1'),
  1240. help='Input directory with IVIM data')
  1241. parser.add_argument('--data_dir', type=str,
  1242. default=None,
  1243. help='Data directory (overrides input_dir)')
  1244. parser.add_argument('--output_dir', type=str,
  1245. default=os.path.join(_this_dir, 'results', 'inr'),
  1246. help='Output directory for fitted parameters')
  1247. parser.add_argument('--denoised_dir', type=str,
  1248. default=None,
  1249. help='Output directory for denoised data (defaults to output_dir if not specified)')
  1250. parser.add_argument('--noise_corrected_dir', type=str,
  1251. default=None,
  1252. help='Output directory for noise corrected data (defaults to output_dir if not specified)')
  1253. # Processing options
  1254. parser.add_argument('--use_noisy_data', action='store_true', default=True,
  1255. help='Use noisy data instead of ground truth')
  1256. parser.add_argument('--apply_noise_correction', action='store_true', default=True,
  1257. help='Apply noise floor correction')
  1258. parser.add_argument('--noise_correction_method', choices=['rayleigh_mean', 'simple_subtraction'],
  1259. default='rayleigh_mean', help='Noise correction method')
  1260. parser.add_argument('--background_threshold', type=float, default=200,
  1261. help='Background threshold for creating masks')
  1262. parser.add_argument('--skip_denoise', action='store_true', help='Skip MDINR denoising step')
  1263. parser.add_argument('--force', action='store_true', help='Force reprocessing even if output files exist')
  1264. # Config parameter removed - configuration is now hardcoded
  1265. # Denoising parameters
  1266. parser.add_argument('--denoise_hidden_size', type=int, default=512)
  1267. parser.add_argument('--denoise_num_layers', type=int, default=4)
  1268. parser.add_argument('--denoise_dropout', type=float, default=0.0)
  1269. parser.add_argument('--denoise_npe', type=int, default=64)
  1270. parser.add_argument('--denoise_maxf', type=float, default=40)
  1271. parser.add_argument('--denoise_use_coords', type=int, default=0)
  1272. parser.add_argument('--denoise_epochs', type=int, default=800)
  1273. parser.add_argument('--denoise_batch_size', type=int, default=8000)
  1274. parser.add_argument('--denoise_lr', type=float, default=1e-4)
  1275. parser.add_argument('--denoise_patch_size', type=int, default=3, help='Patch size for denoising (will create NxNxN patches)')
  1276. # IVIM parameters
  1277. parser.add_argument('--ivim_hidden_size', type=int, default=512)
  1278. parser.add_argument('--ivim_num_layers', type=int, default=4)
  1279. parser.add_argument('--ivim_first_omega_0', type=float, default=30.0)
  1280. parser.add_argument('--ivim_hidden_omega_0', type=float, default=30.0)
  1281. parser.add_argument('--ivim_epochs', type=int, default=800)
  1282. parser.add_argument('--ivim_batch_size', type=int, default=6000)
  1283. parser.add_argument('--ivim_lr', type=float, default=5e-5)
  1284. parser.add_argument('--ivim_use_dwi_signals', action='store_true', default=False,
  1285. help='Use DWI signals as additional input for IVIM fitting')
  1286. parser.add_argument('--ivim_use_b0_signal', action='store_true', default=True,
  1287. help='Use normalized B0 signal as additional input for IVIM fitting')
  1288. parser.add_argument('--ivim_b0_patch_size', type=int, default=1,
  1289. help='Patch size for B0 input in IVIM fitting (default: 3, creates 3x3x3 patches)')
  1290. # Soft constraint parameters
  1291. parser.add_argument('--use_soft_constraint', action='store_true', default=True,
  1292. help='Use soft b0 constraint for IVIM fitting (allows curve to not pass exactly through (0,1))')
  1293. parser.add_argument('--delta_min', type=float, default=-0.01,
  1294. help='Minimum delta offset value for soft constraint')
  1295. parser.add_argument('--delta_max', type=float, default=0.01,
  1296. help='Maximum delta offset value for soft constraint')
  1297. parser.add_argument('--delta_regularization', type=float, default=0.01,
  1298. help='Regularization weight for delta offset (encourages delta close to 0)')
  1299. # Removed ADC parameters - only IVIM fitting is performed
  1300. # General parameters
  1301. parser.add_argument('--gpu', type=int, default=0, help='GPU device to use')
  1302. parser.add_argument('--seed', type=int, default=42, help='Random seed')
  1303. args = parser.parse_args()
  1304. # Set default directories if not provided
  1305. if args.denoised_dir is None:
  1306. args.denoised_dir = args.output_dir
  1307. if args.noise_corrected_dir is None:
  1308. args.noise_corrected_dir = args.output_dir
  1309. # Set random seeds
  1310. np.random.seed(args.seed)
  1311. torch.manual_seed(args.seed)
  1312. if torch.cuda.is_available():
  1313. torch.cuda.manual_seed(args.seed)
  1314. # Verify input directory
  1315. if not os.path.isdir(args.input_dir):
  1316. sys.exit(f'Error: Input directory not found: {args.input_dir}')
  1317. # Create output directory
  1318. os.makedirs(args.output_dir, exist_ok=True)
  1319. # No patient processing - direct IVIM data processing
  1320. print(f'='*80)
  1321. print('IVIM Parameter Fitting Pipeline')
  1322. print(f'='*80)
  1323. print(f'Input directory: {args.input_dir}')
  1324. print(f'Output directory: {args.output_dir}')
  1325. print(f'Use noisy data: {"Yes" if args.use_noisy_data else "No (Ground Truth)"}')
  1326. print(f'Denoising: {"Disabled" if args.skip_denoise else "Enabled"}')
  1327. print(f'Noise correction: {"Enabled" if args.apply_noise_correction else "Disabled"}')
  1328. print(f'Soft constraint: {"Enabled" if args.use_soft_constraint else "Disabled"}')
  1329. if args.use_soft_constraint:
  1330. print(f' Delta range: [{args.delta_min}, {args.delta_max}]')
  1331. print(f' Delta regularization: {args.delta_regularization}')
  1332. print(f'='*80)
  1333. # Process IVIM data
  1334. try:
  1335. process_ivim(args)
  1336. except Exception as e:
  1337. print(f'✖ Processing failed: {e}')
  1338. import traceback
  1339. traceback.print_exc()
  1340. print(f'\n{"="*80}')
  1341. print('Pipeline completed!')
  1342. print(f'{"="*80}')
  1343. if __name__ == '__main__':
  1344. main()

train.py at commit 9a92905, no license · at the source

Overview

Authors: Yunxiang Li1, Yen‐Peng Liao1, Yan Dai1, Jie Deng1, You Zhang1
  1. Department of Radiation Oncology, UT Southwestern Medical Center, Dallas, Texas, USA
Journal: Medical physics, volume 53, issue 8, article e70599
Dates: received 4 December 2025; accepted 1 July 2026; published online 28 July 2026; in print August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1002/mp.70599 · PMID 42519880 · PMCID PMC13411821 · OpenAlex W7171561471
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), other condition (population), systems (subfield)
Methods: Statistics, Preprocessing, Connectivity, Machine learning
Keywords: diffusion‐weighted imaging, implicit neural representation, intravoxel incoherent motion
MeSH: Diffusion Magnetic Resonance Imaging*, Image Processing, Computer-Assisted*, Brain, Humans, Imaging, Three-Dimensional, Motion, Movement, Phantoms, Imaging, Signal-To-Noise Ratio (* major topic)
Topic: MRI in cancer diagnosis (Radiology, Nuclear Medicine and Imaging, Medicine), according to OpenAlex
Funding: NCI NIH HHS (R01 CA258987, R01 CA240808, R01 CA280135); NIBIB NIH HHS (R01 EB034691); NIH HHS (R01 CA240808, R01 CA258987, R01 CA280135, R01 EB034691); National Institutes of Health (R01 CA240808, R01 CA280135, R01 EB034691, R01 CA258987)
Citations: not cited yet (Europe PMC); 52 references in the paper

Abstract

Background: Intravoxel incoherent motion (IVIM) diffusion‐weighted imaging has important value in treatment response monitoring. However, traditional voxel‐wise independent fitting methods are highly sensitive to noise and do not utilize spatial correlation, resulting in unstable parameter estimation. Although existing deep learning methods have shown improvements, they are still limited by local receptive fields.

Purpose: To address this, we propose a two‐stage IVIM parameter estimation framework based on Implicit Neural Representation (IVIM‐INR).

Methods: Our IVIM‐INR method achieves global spatial perception through coordinate encoding and enhances spatial context modeling by leveraging local 3D patch information from multi‐b‐value images. The first stage INR performs signal denoising, and the second stage INR accurately fits IVIM parameters.

Results: Evaluation on brain digital phantoms, AAPM breast IVIM‐dMRI Challenge data, and clinical Glioblastoma (GBM) patient data demonstrates significant advantages of the proposed method over existing techniques. In brain simulation data, when SNR = 50, the normalized mean absolute errors (NMAEs) in tumor regions were 0.16±0.10 for Dp, 0.02±0.01 for Dt, and 0.07±0.05 for Fp, all lower than comparison methods. In the tumor tissues of the 100 cases from the AAPM breast IVIM‐dMRI challenge dataset, Fp error was reduced by 58% compared to ConvNet. Intraclass correlation coefficient (ICC) analysis of real clinical data indicates that our method achieves the best performance with an ICC of Dt in normal tissues reaching 0.629.

Conclusions: By combining INR's continuous function modeling capability with spatial‐aware feature design, IVIM‐INR overcomes the inherent limitations of traditional methods under noisy conditions, providing a more reliable tool for clinical IVIM quantitative analysis.

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 4 matches between paragraphs and lines of code.

Kent0n-Li/IVIM-INR

License: none: the authors keep all their rights
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 9a92905c473df52d259d29f19aeec5e9f589e56d, 22 April 2026
Languages: Python (1)
Size: 2 files, 1 script
Software Heritage: not archived
Found in: the text, “Implementation details”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NiBabel (1 file), NumPy (1 file), pydicom (1 file), PyTorch (1 file), SciPy (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
2 files

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;
  • 1 script, each with its path and the digest of its content;
  • 4 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.

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 3, 28 September 2026

  • Publisher: n/a → Wiley

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 3 keywords, 9 MeSH terms, 4 funders, 45 references.

Cite

This paper

Li, Y., Liao, Y., Dai, Y., Deng, J., & Zhang, Y. (2026). Accurate estimation of intravoxel incoherent motion parameters based on implicit neural representation. Medical physics, 53(8), e70599. https://doi.org/10.1002/mp.70599

BibTeX

@article{li2026accurate,
author = {Li, Yunxiang and Liao, Yen‐Peng and Dai, Yan and Deng, Jie and Zhang, You},
title = {{Accurate estimation of intravoxel incoherent motion parameters based on implicit neural representation}},
journal = {Medical physics},
year = {2026},
month = aug,
volume = {53},
number = {8},
pages = {e70599},
publisher = {Wiley},
issn = {0094-2405},
doi = {10.1002/mp.70599},
url = {https://doi.org/10.1002/mp.70599},
pmid = {42519880},
pmcid = {PMC13411821}
}

RIS

TY - JOUR
AU - Li, Yunxiang
AU - Liao, Yen‐Peng
AU - Dai, Yan
AU - Deng, Jie
AU - Zhang, You
TI - Accurate estimation of intravoxel incoherent motion parameters based on implicit neural representation
T2 - Medical physics
J2 - Med Phys
PY - 2026
DA - 2026/08/01
VL - 53
IS - 8
SP - e70599
SN - 0094-2405
PB - Wiley
DO - 10.1002/mp.70599
UR - https://doi.org/10.1002/mp.70599
LA - en
ER -

CSL-JSON

{
"id": "10.1002/mp.70599",
"type": "article-journal",
"title": "Accurate estimation of intravoxel incoherent motion parameters based on implicit neural representation",
"container-title": "Medical physics",
"author": [
{
"family": "Li",
"given": "Yunxiang"
},
{
"family": "Liao",
"given": "Yen‐Peng"
},
{
"family": "Dai",
"given": "Yan"
},
{
"family": "Deng",
"given": "Jie"
},
{
"family": "Zhang",
"given": "You"
}
],
"container-title-short": "Med Phys",
"volume": "53",
"issue": "8",
"page": "e70599",
"DOI": "10.1002/mp.70599",
"PMID": "42519880",
"PMCID": "PMC13411821",
"ISSN": "0094-2405",
"publisher": "Wiley",
"URL": "https://doi.org/10.1002/mp.70599",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
1
]
]
}
}

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.1002/nbm.70277 [code]
Hierarchical Bayesian Modelling Improves Microstructural Parameter Mapping in Diffusion and Exchange MRI Data.
Journal: NMR in biomedicine
In common: NiBabel, SciPy, NumPy, structural MRI / diffusion, 6 references
[2] doi:10.1038/s43856-026-01614-6 [code]
Simulation-based inference at the theoretical limit for fast, robust microstructural MRI with minimal diffusion data.
Journal: Communications medicine
In common: PyTorch, SciPy, NumPy, structural MRI / diffusion, 5 references
[3] doi:10.1002/mrm.70496 [code]
A Deep Nonlinear Subspace Modeling and Reconstruction for Diffusion-Weighted Imaging Using Denoising Auto-Encoder.
Journal: Magnetic resonance in medicine
In common: pydicom, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion, 1 reference
[4] doi:10.1002/mrm.70405 [code]
DeepRelaxo: Fast Mono-Exponential Magnitude Brain R2* Mapping With Reduced Echoes Using Self-Supervised Deep Learning.
Journal: Magnetic resonance in medicine
In common: pydicom, NiBabel, PyTorch, 2 other tools, 1 reference
[5] 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, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion, other condition
[6] doi:10.1002/alz.71530 [code]
Differential associations of plasma biomarkers with Alzheimer's disease and small vessel disease: A multimodal imaging study.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: NiBabel, PyTorch, SciPy, 1 other tool, structural MRI / diffusion, 2 references
[7] doi:10.3389/fnins.2026.1841093 [code]
Brain protein burden is related to intravoxel incoherent motion: PET-MR imaging study.
Journal: Frontiers in neuroscience
In common: NiBabel, SciPy, NumPy, structural MRI / diffusion, 2 references
[8] doi:10.1186/s13244-026-02296-3 [code]
A pre-trained foundation model framework for multiplanar MRI classification of extramural vascular invasion and mesorectal fascia invasion in rectal cancer.
Journal: Insights into imaging
In common: pydicom, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion, other condition
[9] 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, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion
[10] 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, NiBabel, PyTorch, 2 other tools, structural MRI / diffusion

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.