OSCR

SUITPy: A Python-based toolbox for the analysis of cerebellar functional and anatomical imaging data across the human lifespan.

Code ↔ Paper

6 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 6 matches
  1. [1] § Methods › Cerebellar U-Net isolation › U-Net architecture ↔ SUITPy/isolation.py, lines 554–631 · score 0.78 · corresponding skip connection, transposed convolution, feature maps, concatenated, stride, Block
  2. [2] § Methods › Cerebellar U-Net isolation › Preprocessing ↔ SUITPy/isolation.py, lines 981–1120 · score 0.74 · normalized mutual information, transformation matrix, template space, iteratively, cropped, affine
  3. [3] § Methods › Cerebellar U-Net isolation › U-Net architecture ↔ SUITPy/isolation.py, lines 554–631 · score 0.70 · corresponding skip connection, transposed convolution, convolutional layer, concatenates, strided, Blocks
  4. [4] § Methods › Cerebellar normalization ↔ SUITPy/normalization.py, lines 18–67 · score 0.61 · antsRegistrationSyN, MNI152NLin2009cSymC, space, template, masking, cerebellum
  5. [5] § Methods › Cerebellar U-Net isolation › U-Net architecture ↔ SUITPy/isolation.py, lines 415–493 · score 0.58 · Convolutional Block, convolution layer, padding, kernel, affine, training
  6. [6] § Methods › Cerebellar U-Net isolation › Postprocessing ↔ SUITPy/isolation.py, lines 1242–1259 · score 0.50 · connected component, largest, cluster, voxels, isolation, masks

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,469 lines · 55 KB · MIT · 5 matches

  1. """
  2. Cerebellar Isolation using a Unet model
  3. @authors: Yao Li, Carlos Hernandez-Castillo, Joern Diedrichsen
  4. """
  5. import sys
  6. import argparse
  7. import os
  8. import nibabel as nib
  9. import ants
  10. import numpy as np
  11. import nitools
  12. from typing import Tuple, Union
  13. import pickle
  14. import warnings
  15. class _Conv3dN:
  16. """
  17. numpy implementation of 3D convolution layer
  18. """
  19. def __init__(
  20. self,
  21. in_channels: int,
  22. out_channels: int,
  23. kernel_size: Union[int, Tuple[int, int, int]],
  24. stride: Union[int, Tuple[int, int, int]] = (1, 1, 1),
  25. padding: Union[int, Tuple[int, int, int]] = (0, 0, 0),
  26. dilation: Union[int, Tuple[int, int, int]] = (1, 1, 1),
  27. bias: bool = True,
  28. padding_mode: str = 'zeros',
  29. ):
  30. self.in_channels = in_channels
  31. self.out_channels = out_channels
  32. self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3
  33. self.stride = stride if isinstance(stride, tuple) else (stride,) * 3
  34. self.padding = padding if isinstance(padding, tuple) else (padding,) * 3
  35. self.dilation = dilation if isinstance(dilation, tuple) else (dilation,) * 3
  36. self.bias = bias
  37. self.padding_mode = padding_mode
  38. self.weight = np.random.randn(out_channels, in_channels, *self.kernel_size)
  39. if bias:
  40. self.bias_term = np.zeros(out_channels)
  41. else:
  42. self.bias_term = None
  43. def _pad_input(self, x: np.ndarray) -> np.ndarray:
  44. if all(p == 0 for p in self.padding):
  45. return x
  46. pd, ph, pw = self.padding
  47. if self.padding_mode == 'zeros':
  48. return np.pad(x,
  49. ((0, 0), (0, 0),
  50. (pd, pd), (ph, ph), (pw, pw)),
  51. mode='constant')
  52. elif self.padding_mode == 'reflect':
  53. return np.pad(x,
  54. ((0, 0), (0, 0),
  55. (pd, pd), (ph, ph), (pw, pw)),
  56. mode='reflect')
  57. else:
  58. raise NotImplementedError(f"Padding mode {self.padding_mode} not implemented")
  59. def _im2col(self, x: np.ndarray) -> Tuple[np.ndarray, Tuple[int, int, int]]:
  60. batch_size, in_channels, depth, height, width = x.shape
  61. kd, kh, kw = self.kernel_size
  62. sd, sh, sw = self.stride
  63. pd, ph, pw = self.padding
  64. dd, dh, dw = self.dilation
  65. x_padded = self._pad_input(x)
  66. out_depth = (depth + 2 * pd - dd * (kd - 1) - 1) // sd + 1
  67. out_height = (height + 2 * ph - dh * (kh - 1) - 1) // sh + 1
  68. out_width = (width + 2 * pw - dw * (kw - 1) - 1) // sw + 1
  69. cols = np.zeros((batch_size, in_channels, kd, kh, kw, out_depth, out_height, out_width))
  70. for d in range(kd):
  71. for h in range(kh):
  72. for w in range(kw):
  73. d_start = d * dd
  74. h_start = h * dh
  75. w_start = w * dw
  76. d_slice = slice(d_start, d_start + out_depth * sd, sd)
  77. h_slice = slice(h_start, h_start + out_height * sh, sh)
  78. w_slice = slice(w_start, w_start + out_width * sw, sw)
  79. cols[:, :, d, h, w, :, :, :] = x_padded[:, :, d_slice, h_slice, w_slice]
  80. # Reshape for matrix multiplication (batch, out_d, out_h, out_w, in_ch, kd, kh, kw)
  81. cols = cols.transpose(0, 5, 6, 7, 1, 2, 3, 4)
  82. cols = cols.reshape(batch_size * out_depth * out_height * out_width, -1)
  83. return cols, (out_depth, out_height, out_width)
  84. def load_state(self, weight: np.ndarray, bias_term: np.ndarray):
  85. self.weight = weight
  86. self.bias_term = bias_term
  87. def forward(self, x: np.ndarray) -> np.ndarray:
  88. if x.ndim != 5:
  89. raise ValueError(f"Input must have 5 dimensions (N, C, D, H, W), got {x.ndim}")
  90. batch_size, in_channels, depth, height, width = x.shape
  91. if in_channels != self.in_channels:
  92. raise ValueError(f"Expected {self.in_channels} input channels, got {in_channels}")
  93. cols, (out_depth, out_height, out_width) = self._im2col(x)
  94. # Reshape weights for matrix multiplication (out_ch, in_ch * kd * kh * kw)
  95. weight_flat = self.weight.reshape(self.out_channels, -1)
  96. # Perform convolution via matrix multiplication (batch * out_d * out_h * out_w, out_ch)
  97. output_flat = cols @ weight_flat.T
  98. output = output_flat.reshape(batch_size, out_depth, out_height, out_width, self.out_channels)
  99. output = output.transpose(0, 4, 1, 2, 3) # (batch, out_ch, out_d, out_h, out_w)
  100. if self.bias:
  101. output += self.bias_term.reshape(1, -1, 1, 1, 1)
  102. return output
  103. def __call__(self, x: np.ndarray) -> np.ndarray:
  104. return self.forward(x)
  105. class _ConvTranspose3dN:
  106. """
  107. numpy implementation of 3D transpose convolution layer
  108. """
  109. def __init__(
  110. self,
  111. in_channels: int,
  112. out_channels: int,
  113. kernel_size: Union[int, Tuple[int, int, int]],
  114. stride: Union[int, Tuple[int, int, int]] = (1, 1, 1),
  115. padding: Union[int, Tuple[int, int, int]] = (0, 0, 0),
  116. output_padding: Union[int, Tuple[int, int, int]] = (0, 0, 0),
  117. dilation: Union[int, Tuple[int, int, int]] = (1, 1, 1),
  118. bias: bool = True,
  119. padding_mode: str = "zeros",
  120. ):
  121. self.in_channels = in_channels
  122. self.out_channels = out_channels
  123. self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3
  124. self.stride = stride if isinstance(stride, tuple) else (stride,) * 3
  125. self.padding = padding if isinstance(padding, tuple) else (padding,) * 3
  126. self.output_padding = output_padding if isinstance(output_padding, tuple) else (output_padding,) * 3
  127. self.dilation = dilation if isinstance(dilation, tuple) else (dilation,) * 3
  128. self.bias = bias
  129. self.padding_mode = padding_mode
  130. self.weight = np.random.randn(in_channels, out_channels, *self.kernel_size)
  131. if bias:
  132. self.bias_term = np.zeros(out_channels)
  133. else:
  134. self.bias_term = None
  135. def load_state(self, weight: np.ndarray, bias_term: np.ndarray):
  136. self.weight = weight
  137. self.bias_term = bias_term
  138. def _dilate_input(self, x: np.ndarray) -> np.ndarray:
  139. batch_size, in_channels, depth, height, width = x.shape
  140. sd, sh, sw = self.stride
  141. dilated_depth = (depth - 1) * sd + 1
  142. dilated_height = (height - 1) * sh + 1
  143. dilated_width = (width - 1) * sw + 1
  144. dilated_x = np.zeros((batch_size, in_channels, dilated_depth, dilated_height, dilated_width))
  145. # Insert original values at strided positions
  146. dilated_x[:, :, ::sd, ::sh, ::sw] = x
  147. return dilated_x
  148. def _apply_padding(self, x: np.ndarray) -> np.ndarray:
  149. pd, ph, pw = self.padding
  150. kd, kh, kw = self.kernel_size
  151. pad_d_before = kd - 1 - pd
  152. pad_d_after = kd - 1 - pd + self.output_padding[0]
  153. pad_h_before = kh - 1 - ph
  154. pad_h_after = kh - 1 - ph + self.output_padding[1]
  155. pad_w_before = kw - 1 - pw
  156. pad_w_after = kw - 1 - pw + self.output_padding[2]
  157. padding = ((0, 0), (0, 0),
  158. (pad_d_before, pad_d_after),
  159. (pad_h_before, pad_h_after),
  160. (pad_w_before, pad_w_after))
  161. if self.padding_mode == "zeros":
  162. return np.pad(x, padding, mode='constant')
  163. else:
  164. raise NotImplementedError(f"Padding mode {self.padding_mode} not implemented")
  165. def _im2col(self, x: np.ndarray) -> Tuple[np.ndarray, Tuple[int, int, int]]:
  166. batch_size, in_channels, depth, height, width = x.shape
  167. kd, kh, kw = self.kernel_size
  168. sd, sh, sw = self.stride
  169. pd, ph, pw = self.padding
  170. opd, oph, opw = self.output_padding
  171. dd, dh, dw = self.dilation
  172. dilated_x = self._dilate_input(x)
  173. padded_x = self._apply_padding(dilated_x)
  174. # Output dimensions
  175. out_depth = (depth - 1) * sd - 2 * pd + dd * (kd - 1) + opd + 1
  176. out_height = (height - 1) * sh - 2 * ph + dh * (kh - 1) + oph + 1
  177. out_width = (width - 1) * sw - 2 * pw + dw * (kw - 1) + opw + 1
  178. cols = np.zeros((batch_size, in_channels, kd, kh, kw, out_depth, out_height, out_width))
  179. for d in range(kd):
  180. for h in range(kh):
  181. for w in range(kw):
  182. d_start = d * dd
  183. h_start = h * dh
  184. w_start = w * dw
  185. d_slice = slice(d_start, d_start + out_depth)
  186. h_slice = slice(h_start, h_start + out_height)
  187. w_slice = slice(w_start, w_start + out_width)
  188. cols[:, :, d, h, w, :, :, :] = padded_x[:, :, d_slice, h_slice, w_slice]
  189. # Reshape for matrix multiplication (batch, out_d, out_h, out_w, in_ch, kd, kh, kw)
  190. cols = cols.transpose(0, 5, 6, 7, 1, 2, 3, 4)
  191. cols = cols.reshape(batch_size * out_depth * out_height * out_width, -1)
  192. return cols, (out_depth, out_height, out_width)
  193. def forward(self, x: np.ndarray) -> np.ndarray:
  194. if x.ndim != 5:
  195. raise ValueError(f"Input must have 5 dimensions (N, C, D, H, W), got {x.ndim}")
  196. batch_size, in_channels, depth, height, width = x.shape
  197. if in_channels != self.in_channels:
  198. raise ValueError(f"Expected {self.in_channels} input channels, got {in_channels}")
  199. cols, (out_depth, out_height, out_width) = self._im2col(x)
  200. # Reshape weights for matrix multiplication (out_ch, in_ch * kd * kh * kw)
  201. weight = self.weight.transpose(1, 0, 2, 3, 4)
  202. weight_flip = np.flip(weight, (2, 3, 4))
  203. weight_flat = weight_flip.reshape(self.out_channels, -1)
  204. # Perform convolution via matrix multiplication (batch * out_d * out_h * out_w, out_ch)
  205. output_flat = cols @ weight_flat.T
  206. output = output_flat.reshape(batch_size, out_depth, out_height, out_width, self.out_channels)
  207. output = output.transpose(0, 4, 1, 2, 3) # (batch, out_ch, out_d, out_h, out_w)
  208. if self.bias:
  209. output += self.bias_term.reshape(1, -1, 1, 1, 1)
  210. return output
  211. def __call__(self, x: np.ndarray) -> np.ndarray:
  212. return self.forward(x)
  213. class _InstanceNorm3dN:
  214. """
  215. numpy implementation of 3D instance normalization layer
  216. """
  217. def __init__(
  218. self,
  219. num_features: int,
  220. eps: float = 1e-5,
  221. momentum: float = 0.1,
  222. affine: bool = False,
  223. track_running_stats: bool = False
  224. ):
  225. self.num_features = num_features
  226. self.eps = eps
  227. self.momentum = momentum
  228. self.affine = affine
  229. self.track_running_stats = track_running_stats
  230. if self.affine:
  231. self.weight = np.ones(num_features)
  232. self.bias = np.zeros(num_features)
  233. else:
  234. self.weight = None
  235. self.bias = None
  236. if self.track_running_stats:
  237. raise RuntimeError('track_running_stats currently not supported.')
  238. # self.running_mean = np.zeros(num_features)
  239. # self.running_var = np.ones(num_features)
  240. # self.num_batches_tracked = 0
  241. else:
  242. self.running_mean = None
  243. self.running_var = None
  244. self.num_batches_tracked = None
  245. def _check_input_dim(self, x: np.ndarray) -> None:
  246. if x.ndim != 5:
  247. raise ValueError(f"Expected 5D input (got {x.ndim}D input)")
  248. if x.shape[1] != self.num_features:
  249. raise ValueError(f"Expected {self.num_features} channels in input (got {x.shape[1]} channels)")
  250. def _compute_instance_stats(self, x: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
  251. """Compute mean and variance per instance per channel across spatial dimensions (D, H, W)"""
  252. n, c, d, h, w = x.shape
  253. x_reshaped = x.reshape(n, c, -1)
  254. mean = np.mean(x_reshaped, axis=2)
  255. var = np.var(x_reshaped, axis=2)
  256. return mean, var
  257. def _normalize(self, x: np.ndarray, mean: np.ndarray, var: np.ndarray) -> np.ndarray:
  258. """Normalize input using computed statistics"""
  259. n, c, d, h, w = x.shape
  260. mean = mean.reshape(n, c, 1, 1, 1)
  261. var = var.reshape(n, c, 1, 1, 1)
  262. x_normalized = (x - mean) / np.sqrt(var + self.eps)
  263. return x_normalized
  264. def forward(self, x: np.ndarray) -> np.ndarray:
  265. self._check_input_dim(x)
  266. if self.track_running_stats:
  267. n, c, d, h, w = x.shape
  268. mean = self.running_mean.reshape(1, c, 1, 1, 1)
  269. var = self.running_var.reshape(1, c, 1, 1, 1)
  270. else:
  271. # If not tracking running stats, compute instance stats even during inference
  272. mean, var = self._compute_instance_stats(x)
  273. if not self.track_running_stats:
  274. x_normalized = self._normalize(x, mean, var)
  275. else:
  276. x_normalized = (x - mean) / np.sqrt(var + self.eps)
  277. if self.affine:
  278. weight = self.weight.reshape(1, -1, 1, 1, 1)
  279. bias = self.bias.reshape(1, -1, 1, 1, 1)
  280. x_normalized = x_normalized * weight + bias
  281. return x_normalized
  282. def __call__(self, x: np.ndarray) -> np.ndarray:
  283. return self.forward(x)
  284. def load_state_dict(self, state_dict: dict) -> None:
  285. if self.affine:
  286. self.weight = state_dict['weight']
  287. self.bias = state_dict['bias']
  288. if self.track_running_stats:
  289. self.running_mean = state_dict['running_mean']
  290. self.running_var = state_dict['running_var']
  291. self.num_batches_tracked = state_dict['num_batches_tracked']
  292. class _LeakyReLUN:
  293. """
  294. numpy implementation of Leaky ReLUN layer
  295. """
  296. def __init__(self, negative_slope: float = 0.01):
  297. self.negative_slope = negative_slope
  298. def forward(self, x: np.ndarray) -> np.ndarray:
  299. return np.maximum(0, x) + self.negative_slope * np.minimum(0, x)
  300. def __call__(self, x: np.ndarray) -> np.ndarray:
  301. return self.forward(x)
  302. class _SoftmaxN:
  303. """
  304. numpy implementation of Softmax layer
  305. """
  306. def __init__(self, dim: int = -1,):
  307. self.dim = dim
  308. def forward(self, x: np.ndarray) -> np.ndarray:
  309. dim = self.dim
  310. x_shifted = x - np.max(x, axis=dim, keepdims=True)
  311. exp_x = np.exp(x_shifted)
  312. softmax_output = exp_x / np.sum(exp_x, axis=dim, keepdims=True)
  313. return softmax_output
  314. def __call__(self, x: np.ndarray) -> np.ndarray:
  315. return self.forward(x)
  316. class _ConvNetN:
  317. """
  318. convolution block in Unet
  319. """
  320. def __init__(self, in_channels: int, out_channels: int):
  321. """
  322. Args:
  323. in_channels: (int)
  324. the number of input channels
  325. out_channels: (int)
  326. the number of output channels
  327. """
  328. self.conv1 = _Conv3dN(in_channels=in_channels, out_channels=out_channels, kernel_size=3, padding=1)
  329. self.norm1 = _InstanceNorm3dN(num_features=out_channels, affine=True, track_running_stats=False)
  330. self.activate1= _LeakyReLUN()
  331. self.conv2 = _Conv3dN(in_channels=out_channels, out_channels=out_channels, kernel_size=3, padding=1)
  332. self.norm2 = _InstanceNorm3dN(num_features=out_channels, affine=True, track_running_stats=False)
  333. self.activate2 = _LeakyReLUN()
  334. def load_state(
  335. self,
  336. c_weight1: np.ndarray,
  337. c_bias1: np.ndarray,
  338. n_state_dict1: dict,
  339. c_weight2: np.ndarray,
  340. c_bias2: np.ndarray,
  341. n_state_dict2: dict,
  342. ):
  343. """
  344. load the previously trained parameters
  345. Args:
  346. c_weight1: (ndarray)
  347. the weight of the first convolution layer
  348. c_bias1: (ndarray)
  349. the bias of the first convolution layer
  350. n_state_dict1: (dict)
  351. the state dictionary of the first instance normalization layer
  352. c_weight2: (ndarray)
  353. the weight of the second convolution layer
  354. c_bias2: (ndarray)
  355. the bias of the second convolution layer
  356. n_state_dict2: (dict)
  357. the state dictionary of the second instance normalization layer
  358. Returns:
  359. """
  360. self.conv1.load_state(weight=c_weight1, bias_term=c_bias1)
  361. self.norm1.load_state_dict(state_dict=n_state_dict1)
  362. self.conv2.load_state(weight=c_weight2, bias_term=c_bias2)
  363. self.norm2.load_state_dict(state_dict=n_state_dict2)
  364. def forward(self, x: np.ndarray) -> np.ndarray:
  365. """
  366. calculation
  367. Args:
  368. x: (ndarray)
  369. input
  370. Returns:
  371. y: (ndarray)
  372. output
  373. """
  374. result1 = self.conv1(x)
  375. result2 = self.norm1(result1)
  376. result3 = self.activate1(result2)
  377. result4 = self.conv2(result3)
  378. result5 = self.norm2(result4)
  379. result6 = self.activate2(result5)
  380. return result6
  381. def __call__(self, x: np.ndarray) -> np.ndarray:
  382. return self.forward(x)
  383. class _DownSampleN:
  384. """
  385. downsample block in unet
  386. """
  387. def __init__(self, channels: int):
  388. """
  389. Args:
  390. channels: (int)
  391. the number of input channels
  392. """
  393. self.conv1 = _Conv3dN(in_channels=channels, out_channels=channels, kernel_size=3, stride=2, padding=1)
  394. self.norm1 = _InstanceNorm3dN(num_features=channels, affine=True, track_running_stats=False)
  395. self.activate1 = _LeakyReLUN()
  396. def load_state(
  397. self,
  398. c_weight: np.ndarray,
  399. c_bias: np.ndarray,
  400. n_state_dict: dict
  401. ):
  402. """
  403. load the previously trained parameters
  404. Args:
  405. c_weight: (ndarray)
  406. the weight of the convolution layer
  407. c_bias: (ndarray)
  408. the bias of the convolution layer
  409. n_state_dict: (dict)
  410. the state dictionary of the instance normalization layer
  411. """
  412. self.conv1.load_state(weight=c_weight, bias_term=c_bias)
  413. self.norm1.load_state_dict(state_dict=n_state_dict)
  414. def forward(self, x: np.ndarray) -> np.ndarray:
  415. """
  416. calculation
  417. Args:
  418. x: (ndarray)
  419. input
  420. Returns:
  421. y: (ndarray)
  422. output
  423. """
  424. result1 = self.conv1(x)
  425. result2 = self.norm1(result1)
  426. result3 = self.activate1(result2)
  427. return result3
  428. def __call__(self, x: np.ndarray) -> np.ndarray:
  429. return self.forward(x)
  430. class _UpSampleN:
  431. """
  432. upsample block in unet
  433. """
  434. def __init__(self, channels: int):
  435. """
  436. Args:
  437. channels: (int)
  438. the number of input channels
  439. """
  440. self.conv1 = _ConvTranspose3dN(in_channels=channels, out_channels=channels, kernel_size=2, stride=2, padding=0)
  441. self.norm1 = _InstanceNorm3dN(num_features=channels, affine=True, track_running_stats=False)
  442. self.activate1 = _LeakyReLUN()
  443. self.conv2 = _Conv3dN(in_channels=channels * 2, out_channels=channels, kernel_size=3, padding=1)
  444. self.norm2 = _InstanceNorm3dN(num_features=channels, affine=True, track_running_stats=False)
  445. self.activate2 = _LeakyReLUN()
  446. def load_state(
  447. self,
  448. c_weight1: np.ndarray,
  449. c_bias1: np.ndarray,
  450. n_state_dict1: dict,
  451. c_weight2: np.ndarray,
  452. c_bias2: np.ndarray,
  453. n_state_dict2: dict
  454. ):
  455. """
  456. load the previously trained parameters
  457. Args:
  458. c_weight1: (ndarray)
  459. the weight of the transpose convolution layer
  460. c_bias1: (ndarray)
  461. the bias of the transpose convolution layer
  462. n_state_dict1: (dict)
  463. the state dictionary of the first instance normalization layer
  464. c_weight2: (ndarray)
  465. the weight of the convolution layer
  466. c_bias2: (ndarray)
  467. the bias of the convolution layer
  468. n_state_dict2: (dict)
  469. the state dictionary of the second instance normalization layer
  470. Returns:
  471. """
  472. self.conv1.load_state(weight=c_weight1, bias_term=c_bias1)
  473. self.norm1.load_state_dict(state_dict=n_state_dict1)
  474. self.conv2.load_state(weight=c_weight2, bias_term=c_bias2)
  475. self.norm2.load_state_dict(state_dict=n_state_dict2)
  476. def forward(self, x: np.ndarray, feature_map: np.ndarray) -> np.ndarray:
  477. """
  478. calculation
  479. Args:
  480. x: (ndarray)
  481. input
  482. feature_map: (ndarray)
  483. The feature map form the corresponding skip connection
  484. Returns:
  485. y: (ndarray)
  486. output
  487. """
  488. result1 = self.conv1(x)
  489. result2 = self.norm1(result1)
  490. result3 = self.activate1(result2)
  491. result3 = np.concatenate([result3, feature_map], axis=1)
  492. result4 = self.conv2(result3)
  493. result5 = self.norm2(result4)
  494. result6 = self.activate2(result5)
  495. return result6
  496. def __call__(self, x: np.ndarray, feature_map: np.ndarray) -> np.ndarray:
  497. return self.forward(x, feature_map)
  498. class UNetN:
  499. """
  500. Unet model architecture.
  501. """
  502. def __init__(self, init_features: int = 16):
  503. """
  504. Args:
  505. init_features: int
  506. Number of filters in the first convolutional block.
  507. """
  508. self.enc1 = _ConvNetN(in_channels=2, out_channels=init_features)
  509. self.down1 = _DownSampleN(channels=init_features)
  510. self.enc2 = _ConvNetN(in_channels=init_features, out_channels=init_features * 2)
  511. self.down2 = _DownSampleN(init_features * 2)
  512. self.enc3 = _ConvNetN(in_channels=init_features * 2, out_channels=init_features * 4)
  513. self.down3 = _DownSampleN(init_features * 4)
  514. self.enc4 = _ConvNetN(in_channels=init_features * 4, out_channels=init_features * 8)
  515. self.down4 = _DownSampleN(init_features * 8)
  516. self.bottleneck = _ConvNetN(in_channels=init_features * 8, out_channels=init_features * 8)
  517. self.up4 = _UpSampleN(init_features * 8)
  518. self.dec4 = _ConvNetN(in_channels=init_features * 8, out_channels=init_features * 4)
  519. self.up3 = _UpSampleN(init_features * 4)
  520. self.dec3 = _ConvNetN(in_channels=init_features * 4, out_channels=init_features * 2)
  521. self.up2 = _UpSampleN(init_features * 2)
  522. self.dec2 = _ConvNetN(in_channels=init_features * 2, out_channels=init_features)
  523. self.up1 = _UpSampleN(init_features)
  524. self.dec1 = _ConvNetN(in_channels=init_features, out_channels=2)
  525. self.out = _SoftmaxN(dim=1)
  526. def load_state_dict(self, state_dict: dict):
  527. """
  528. Load the state dict of the Unet model.
  529. Args:
  530. state_dict: (dict)
  531. The state dict of the Unet model
  532. """
  533. convs = {'enc1': self.enc1, 'enc2': self.enc2, 'enc3': self.enc3, 'enc4': self.enc4, 'bottleneck': self.bottleneck, 'dec1': self.dec1, 'dec2': self.dec2, 'dec3': self.dec3, 'dec4': self.dec4}
  534. downs = {'down1': self.down1, 'down2': self.down2, 'down3': self.down3, 'down4': self.down4}
  535. ups = {'up1': self.up1, 'up2': self.up2, 'up3': self.up3, 'up4': self.up4}
  536. for c in convs:
  537. c_weight1 = state_dict[f'{c}.layer.0.weight']
  538. c_bias1 = state_dict[f'{c}.layer.0.bias']
  539. n1 = {
  540. 'weight': state_dict[f'{c}.layer.1.weight'],
  541. 'bias': state_dict[f'{c}.layer.1.bias'],
  542. # 'running_mean': state_dict[f'{c}.layer.1.running_mean'],
  543. # 'running_var': state_dict[f'{c}.layer.1.running_var'],
  544. # 'num_batches_tracked': state_dict[f'{c}.layer.1.num_batches_tracked']
  545. }
  546. c_weight2 = state_dict[f'{c}.layer.3.weight']
  547. c_bias2 = state_dict[f'{c}.layer.3.bias']
  548. n2 = {
  549. 'weight': state_dict[f'{c}.layer.4.weight'],
  550. 'bias': state_dict[f'{c}.layer.4.bias'],
  551. # 'running_mean': state_dict[f'{c}.layer.4.running_mean'],
  552. # 'running_var': state_dict[f'{c}.layer.4.running_var'],
  553. # 'num_batches_tracked': state_dict[f'{c}.layer.4.num_batches_tracked']
  554. }
  555. convs[c].load_state(c_weight1=c_weight1, c_bias1=c_bias1, n_state_dict1=n1, c_weight2=c_weight2, c_bias2=c_bias2, n_state_dict2=n2)
  556. for d in downs:
  557. c_weight = state_dict[f'{d}.layer.0.weight']
  558. c_bias = state_dict[f'{d}.layer.0.bias']
  559. n = {
  560. 'weight': state_dict[f'{d}.layer.1.weight'],
  561. 'bias': state_dict[f'{d}.layer.1.bias'],
  562. # 'running_mean': state_dict[f'{d}.layer.1.running_mean'],
  563. # 'running_var': state_dict[f'{d}.layer.1.running_var'],
  564. # 'num_batches_tracked': state_dict[f'{d}.layer.1.num_batches_tracked']
  565. }
  566. downs[d].load_state(c_weight=c_weight, c_bias=c_bias, n_state_dict=n)
  567. for u in ups:
  568. c_weight1 = state_dict[f'{u}.up.0.weight']
  569. c_bias1 = state_dict[f'{u}.up.0.bias']
  570. n1 = {
  571. 'weight': state_dict[f'{u}.up.1.weight'],
  572. 'bias': state_dict[f'{u}.up.1.bias'],
  573. # 'running_mean': state_dict[f'{u}.up.1.running_mean'],
  574. # 'running_var': state_dict[f'{u}.up.1.running_var'],
  575. # 'num_batches_tracked': state_dict[f'{u}.up.1.num_batches_tracked']
  576. }
  577. c_weight2 = state_dict[f'{u}.layer.0.weight']
  578. c_bias2 = state_dict[f'{u}.layer.0.bias']
  579. n2 = {
  580. 'weight': state_dict[f'{u}.layer.1.weight'],
  581. 'bias': state_dict[f'{u}.layer.1.bias'],
  582. # 'running_mean': state_dict[f'{u}.layer.1.running_mean'],
  583. # 'running_var': state_dict[f'{u}.layer.1.running_var'],
  584. # 'num_batches_tracked': state_dict[f'{u}.layer.1.num_batches_tracked']
  585. }
  586. ups[u].load_state(c_weight1=c_weight1, c_bias1=c_bias1, n_state_dict1=n1, c_weight2=c_weight2, c_bias2=c_bias2, n_state_dict2=n2)
  587. def forward(self, t1: np.ndarray, t2: np.ndarray) -> np.ndarray:
  588. """
  589. Unet calculation
  590. Args:
  591. t1: (ndarray)
  592. the cropped T1w image
  593. t2: (ndarray)
  594. the cropped T2w image
  595. Returns:
  596. pm: (ndarray)
  597. the cerebellar probability map
  598. """
  599. result_ini = np.concatenate((t1, t2), axis=1)
  600. result_1 = self.enc1(result_ini)
  601. result_2 = self.enc2(self.down1(result_1))
  602. result_3 = self.enc3(self.down2(result_2))
  603. result_4 = self.enc4(self.down3(result_3))
  604. result_5 = self.bottleneck(self.down4(result_4))
  605. out_1 = self.dec4(self.up4(result_5, result_4))
  606. out_2 = self.dec3(self.up3(out_1, result_3))
  607. out_3 = self.dec2(self.up2(out_2, result_2))
  608. out_4 = self.dec1(self.up1(out_3, result_1))
  609. return self.out(out_4)
  610. def __call__(self, t1: np.ndarray, t2: np.ndarray) -> np.ndarray:
  611. return self.forward(t1, t2)
  612. def _load_model(params_file: str):
  613. """
  614. load model with pretrained weights
  615. Args:
  616. params_file: (string)
  617. path to the pretrained weights
  618. Returns:
  619. net: (Unet)
  620. the pretrained model
  621. """
  622. net = UNetN()
  623. if os.path.exists(params_file):
  624. with open(params_file, "rb") as f:
  625. params = pickle.load(f)
  626. net.load_state_dict(params)
  627. else:
  628. raise RuntimeError('fail to load pre-trained parameters')
  629. return net
  630. def predict(params_file: str, t1: np.ndarray = None, t2: np.ndarray = None) -> np.ndarray:
  631. """
  632. Run a prediction on a single subject using a trained UNet model
  633. Args:
  634. params_file: (string)
  635. filename of the pretrained weights
  636. t1: (ndarray)
  637. Numpy array of T1w cerebellar image (after cropping)
  638. t2: (ndarray)
  639. Numpy array of T2w cerebellar image (after cropping)
  640. Returns:
  641. mask: (ndarray)
  642. the 3D numpy array of predicted mask (template space)
  643. """
  644. net = _load_model(params_file)
  645. if t1 is None:
  646. t1 = np.zeros((128, 128, 128))
  647. else:
  648. t1 = (t1 - t1.mean()) / t1.std()
  649. if t2 is None:
  650. t2 = np.zeros((128, 128, 128))
  651. else:
  652. t2 = (t2 - t2.mean()) / t2.std()
  653. t1, t2 = t1.reshape((1, 1, 128, 128, 128)), t2.reshape((1, 1, 128, 128, 128))
  654. mask = net(t1, t2)
  655. return mask[0][0]
  656. class InputError(Exception):
  657. def __init__(self, message):
  658. super().__init__(message)
  659. def img_read(file: str, use_q_form: bool = False, verbose: bool = False) -> ants.ANTsImage:
  660. """
  661. basic function to read a nifti image as an ANTs image.
  662. checks for consistency of s-form and q-form and
  663. uses s-form by default, or q-form if use_q_form is set to True.
  664. Args:
  665. file: (string)
  666. image path
  667. use_q_form: (bool)
  668. set to True to use q-form
  669. verbose: (bool)
  670. whether to print warning info
  671. Returns:
  672. img (ANTs image)
  673. An ANTs image
  674. """
  675. nib_img = nib.load(file)
  676. s_form, _ = nib_img.get_sform(coded=True)
  677. q_form, _ = nib_img.get_qform(coded=True)
  678. if s_form is None and q_form is None:
  679. raise InputError(f'Both S-form and Q-form are None in {file}')
  680. elif s_form is not None and q_form is None:
  681. if use_q_form:
  682. raise InputError(f'Q-form is None in {file}, Please set use_q_form=False')
  683. else:
  684. nib_img.set_qform(s_form)
  685. elif s_form is None and q_form is not None:
  686. if not use_q_form:
  687. if verbose:
  688. print(f'Warning: S-form is None in {file}. Using Q-form.')
  689. nib_img.set_sform(q_form)
  690. else:
  691. if (s_form != q_form).any():
  692. if not use_q_form:
  693. if verbose:
  694. print(f'Warning: S- and Q-form indicate different image orientations in {file}. Using S-form. Set use_q_form=TRUE to use the q-form.')
  695. nib_img.set_qform(s_form)
  696. else:
  697. if verbose:
  698. print(f'Warning: S- and Q-form indicate different image orientations in {file}. Using Q-form.')
  699. nib_img.set_sform(q_form)
  700. new_img = ants.from_nibabel_nifti(nib_img)
  701. return new_img
  702. def normalized_mutual_information(image1: np.ndarray, image2: np.ndarray, bins:int = 100) -> float:
  703. """
  704. Compute the normalized mutual information between two images
  705. Args:
  706. image1: (ndarray)
  707. First image
  708. image2: (ndarray)
  709. Second image
  710. bins: (int)
  711. number of bins
  712. Returns:
  713. mi: (float)
  714. normalized mutual information between image1 and image2
  715. """
  716. joint_hist, _, _ = np.histogram2d(
  717. image1.ravel(),
  718. image2.ravel(),
  719. bins=bins
  720. )
  721. joint_prob = joint_hist / np.sum(joint_hist)
  722. p_x = np.sum(joint_prob, axis=1)
  723. p_y = np.sum(joint_prob, axis=0)
  724. # H(X) = -Σ p(x) log p(x)
  725. h_x = -1 * np.sum(p_x[p_x > 0] * np.log2(p_x[p_x > 0]))
  726. # H(Y) = -Σ p(y) log p(y)
  727. h_y = -1 * np.sum(p_y[p_y > 0] * np.log2(p_y[p_y > 0]))
  728. # H(X,Y) = -ΣΣ p(x,y) log p(x,y)
  729. h_xy = -1 * np.sum(joint_prob[joint_prob > 0] * np.log2(joint_prob[joint_prob > 0]))
  730. mi = (h_x + h_y) / h_xy
  731. return mi
  732. def registration(img: ants.ANTsImage, brain_mask: ants.ANTsImage = None, template_name: str = 'MNI152NLin6Asym', type_of_transform: str = 'Similarity', max_iterations: int = 5, mi_lower: float = 1.22, mi_upper:float = 1.23) -> Tuple[ants.ANTsTransform, int]:
  733. """
  734. register the image to this template. The function runs ants registration multiple time and find the best registration results evaluated using normalized mutual information.
  735. Args:
  736. img: (ANTsImage)
  737. image to be registered
  738. brain_mask: (ANTsImage)
  739. brain mask of the image
  740. template: (string)
  741. The name of template used. It uses MNI152NLin6Asym by default. (Note that the template must be stored in SUITPy/templates directory and come with brain mask)
  742. type_of_transform: (string)
  743. transform type (Affine by default, check ANTsPY[https://antspy.readthedocs.io/en/latest/registration.html] for details)
  744. max_iterations: (int)
  745. maximum number of iterations
  746. mi_lower: (float)
  747. lower bound of normalized mutual information (range (1, 2)); setting this value too high would result in an error, while setting it too low might cause poor alignment.
  748. mi_upper: (float)
  749. upper bound of normalized mutual information (range (1, 2)); setting this value too high would result in a warning, while setting it too low would prevent the function from searching for better results.
  750. Returns:
  751. trans: (ANTsTransform)
  752. the transformation from the subject space to the template space
  753. status: (int)
  754. the status of the registration (0: for success mi>upper, 1: found one of mi>mi_lower, 2: failure)
  755. """
  756. base_dir = os.path.dirname(__file__)
  757. template = ants.image_read(os.path.join(base_dir, f'templates/tpl-{template_name}_T1w.nii.gz'))
  758. template_brain_mask = ants.image_read(os.path.join(base_dir, f'templates/tpl-{template_name}-brain_mask.nii.gz'))
  759. template_brain = ants.mask_image(template, template_brain_mask)
  760. cnt = 0
  761. status = 2
  762. trans_final = None
  763. best_mi = 0
  764. while cnt < max_iterations:
  765. if brain_mask is not None:
  766. brain = ants.mask_image(img, brain_mask)
  767. result = ants.registration(fixed=template_brain, moving=brain, type_of_transform=type_of_transform)
  768. trans = ants.read_transform(result['fwdtransforms'][0])
  769. else:
  770. result = ants.registration(fixed=template, moving=img, type_of_transform=type_of_transform)
  771. trans = ants.read_transform(result['fwdtransforms'][0])
  772. img_transformed = ants.apply_ants_transform_to_image(trans, img, template)
  773. brain_transformed = ants.mask_image(img_transformed, template_brain_mask)
  774. mi = normalized_mutual_information(brain_transformed.numpy(), template_brain.numpy())
  775. if mi > best_mi:
  776. best_mi = mi
  777. trans_final = trans
  778. if mi > mi_upper:
  779. status = 0
  780. break
  781. if mi > mi_lower:
  782. status = 1
  783. cnt += 1
  784. return trans_final, status
  785. class RegistrationError(Exception):
  786. def __init__(self, message):
  787. super().__init__(message)
  788. class TemplateCerebellarBoundingBox:
  789. """
  790. Basic cerebellar bounding box class, which defines the cropped area.
  791. All other template implementations should be registered to this template.
  792. """
  793. def __init__(self, template_name: str ='MNI152NLin6Asym', bounding_box: np.ndarray = None, cerebellar_center: np.ndarray = None, cropped_size: np.ndarray = None):
  794. """
  795. Create a bounding box
  796. Args:
  797. template_name: (string)
  798. The name of template used. It uses MNI152NLin6Asym by default.
  799. bounding_box: (ndarray)
  800. cerebellar bounding box in MNI space (in mm) (reserved for future development)
  801. cerebellar_center: (ndarray)
  802. (reserved for future development)
  803. cropped_size: (ndarray)
  804. (reserved for future development)
  805. """
  806. self.template_name = template_name
  807. self.cropped_size = (128, 128, 128)
  808. if bounding_box is not None:
  809. self.bounding_box = bounding_box
  810. else:
  811. if cerebellar_center is not None and cropped_size is not None:
  812. # reserved for future development
  813. pass
  814. else:
  815. self.bounding_box = np.array([[64, -114, -88], [-64, 14, 40]])
  816. self.lowerleft = self.bounding_box[0]
  817. self.upperright = self.bounding_box[1]
  818. base_dir = os.path.dirname(__file__)
  819. self.template = ants.image_read(os.path.join(base_dir, f'templates/tpl-{self.template_name}_T1w.nii.gz'))
  820. self.nib_template = nib.load(os.path.join(base_dir, f'templates/tpl-{self.template_name}_T1w.nii.gz'))
  821. self.brainmask = ants.image_read(os.path.join(base_dir, f'templates/tpl-{self.template_name}-brain_mask.nii.gz'))
  822. self.brain = ants.mask_image(self.template, self.brainmask)
  823. self.affine = nib.load(os.path.join(base_dir, f'templates/tpl-{self.template_name}_T1w.nii.gz')).affine
  824. def get_crop_indices(self) -> np.ndarray:
  825. """
  826. calculate the lower left and upper right indices of the cropped area (in voxels).
  827. Returns:
  828. indices: (ndarray)
  829. a 2 * 3 ndarray consists of two vertices defining the bounding box
  830. """
  831. return nitools.coords_to_voxelidxs(self.bounding_box.T, self.nib_template).T
  832. def get_cropped_affine(self) -> np.ndarray:
  833. """
  834. get the cropped area affine
  835. Returns:
  836. affine: (ndarray)
  837. The affine form for the cropped image
  838. """
  839. # This function needs to fix. It will fail if the template affine is not diagonal
  840. affine = np.diag([self.affine[0, 0], self.affine[1, 1], self.affine[2, 2], 1])
  841. affine[0, 3] = abs(affine[0, 0]) * self.lowerleft[0]
  842. affine[1, 3] = abs(affine[1, 1]) * self.lowerleft[1]
  843. affine[2, 3] = abs(affine[2, 2]) * self.lowerleft[2]
  844. return affine
  845. def registration(self, img: ants.ANTsImage, brain_mask: ants.ANTsImage = None, type_of_transform: str = 'Similarity', max_iterations: int = 5, mi_lower: float = 1.22, mi_upper:float = 1.23) -> Tuple[ants.ANTsTransform, int]:
  846. """
  847. register the image to this template
  848. Args:
  849. img: (ANTsImage)
  850. image to be registered
  851. brain_mask: (ANTsImage)
  852. brain mask of the image
  853. type_of_transform: (string)
  854. transform type (Affine by default, check ANTsPY[https://antspy.readthedocs.io/en/latest/registration.html] for details)
  855. max_iterations: (int)
  856. maximum number of iterations
  857. mi_lower: (float)
  858. lower boundary of normalized mutual information
  859. mi_upper: (float)
  860. upper boundary of normalized mutual information
  861. Returns:
  862. trans: (ANTsTransform)
  863. the transformation from the subject space to the template space
  864. status: (int)
  865. the status of the registration (0 for success, 1 for uncertainty, 2 for failure)
  866. """
  867. trans, status = registration(img=img, brain_mask=brain_mask,template_name=self.template_name, type_of_transform=type_of_transform, max_iterations=max_iterations, mi_lower=mi_lower, mi_upper=mi_upper)
  868. return trans, status
  869. def crop(self, img: ants.ANTsImage, trans: ants.ANTsTransform = None) -> Tuple[ants.ANTsImage, ants.ANTsImage]:
  870. """
  871. Crop the cerebellar area using the bounding box.
  872. Args:
  873. img: (ANTsImage)
  874. image to be cropped
  875. trans: (ANTsTransform)
  876. transformation matrix from the image space to the template space (only use it if img is not in the MNI template)
  877. Returns:
  878. cropped_img: (ANTsImage)
  879. cropped image
  880. img : (ANTsImage)
  881. the whole transformed image
  882. """
  883. start_indices, end_indices = self.get_crop_indices()
  884. if trans is not None:
  885. img = ants.apply_ants_transform_to_image(trans, img, self.template)
  886. return ants.crop_indices(img, tuple(start_indices.astype(int)), tuple(end_indices.astype(int))), img
  887. def template2subject(self, img: ants.ANTsImage, trans: ants.ANTsTransform, ref: ants.ANTsImage) -> ants.ANTsImage:
  888. """
  889. transform the image from template space to the subject space
  890. Args:
  891. img: (ANTsImage)
  892. the image to be transformed
  893. trans: (ANTs transformation)
  894. transformation matrix (from subject space to template space)
  895. ref: (ANTsImage)
  896. reference image
  897. Returns:
  898. img: (ANTsImage)
  899. the transformed image in subject space
  900. """
  901. trans_inv = ants.invert_ants_transform(trans)
  902. result = ants.apply_ants_transform_to_image(trans_inv, img, ref)
  903. return result
  904. def subject_preprocess(t1_file: str = None, t2_file: str = None, brain_mask_file: str = None, label_file: str = None,
  905. BoundingBox: TemplateCerebellarBoundingBox = TemplateCerebellarBoundingBox(),
  906. type_of_transform: str = 'Similarity', max_iterations: int = 5, use_q_form: bool = False) -> Tuple[ants.ANTsTransform, ants.ANTsImage, ants.ANTsImage, ants.ANTsImage, ants.ANTsImage, ants.ANTsImage]:
  907. """
  908. function to preprocess a single subject.
  909. 1. Transform the image from subject space to the template space
  910. 2. Using a pre-defined bounding box to crop the image in the template space
  911. Args:
  912. t1_file: (string)
  913. file to T1w image
  914. t2_file: (string)
  915. file to T2w image
  916. brain_mask_file: (string)
  917. file to the brain mask image (can be used to improve affine registration)
  918. label_file: (string)
  919. file to label image (Optional, this image will be transformed into the template space using the same transformation.)
  920. BoundingBox: (TemplateCerebellarBoundingBox)
  921. the bounding box
  922. type_of_transform: (string)
  923. reserved for future use (see ANTspy)
  924. max_iterations: (int)
  925. maximum number of registration iterations
  926. use_q_form: (bool)
  927. set to True to use q-form
  928. Returns:
  929. trans: (ANTsTransform)
  930. transformation from subject space to template space
  931. t1_crop: (ANTsImage)
  932. cropped cerebellar area from transformed T1w image
  933. t2_crop: (ANTsImage)
  934. cropped cerebellar area from transformed T2w image
  935. label_crop: (ANTsImage)
  936. cropped cerebellar area from transformed label image
  937. t1_whole: (ANTsImage)
  938. whole transformed T1w image
  939. t2_whole: (ANTsImage)
  940. whole transformed T2w image
  941. """
  942. if t1_file is not None:
  943. t1 = img_read(file=t1_file, use_q_form=use_q_form, verbose=True)
  944. else:
  945. t1 = None
  946. if t2_file is not None:
  947. t2 = img_read(file=t2_file, use_q_form=use_q_form, verbose=True)
  948. else:
  949. t2 = None
  950. if brain_mask_file is not None:
  951. brain_mask = img_read(file=brain_mask_file, use_q_form=use_q_form)
  952. else:
  953. brain_mask = None
  954. # Read additional images
  955. if label_file is not None:
  956. label = img_read(file=label_file, use_q_form=use_q_form)
  957. else:
  958. label = None
  959. # If T1 and T2 are both given, but not in the same spacing, reslice T2 to T1 first
  960. if t2 is not None and t1 is not None:
  961. if ants.get_spacing(t1) != ants.get_spacing(t2):
  962. t2 = ants.registration(fixed=t1, moving=t2, type_of_transform='Rigid')['warpedmovout']
  963. if t1 is not None:
  964. trans, status = BoundingBox.registration(img=t1, brain_mask=brain_mask, type_of_transform=type_of_transform, max_iterations=max_iterations)
  965. else:
  966. trans, status = BoundingBox.registration(img=t2, brain_mask=brain_mask, type_of_transform=type_of_transform, max_iterations=max_iterations)
  967. if status == 2:
  968. raise RegistrationError('invalid registration detected')
  969. if status == 1:
  970. warnings.warn(f'low-quality registration detected, please double check the results for {t1_file if t1 is not None else t2_file}')
  971. if t1 is not None:
  972. t1_crop, t1_whole = BoundingBox.crop(t1, trans)
  973. else:
  974. t1_crop = None
  975. t1_whole = None
  976. if t2 is not None:
  977. t2_crop, t2_whole = BoundingBox.crop(t2, trans)
  978. else:
  979. t2_crop = None
  980. t2_whole = None
  981. if label is not None:
  982. label_crop, _ = BoundingBox.crop(label, trans)
  983. else:
  984. label_crop = None
  985. return trans, t1_crop, t2_crop, label_crop, t1_whole, t2_whole
  986. def threshold(img: ants.ANTsImage, lower: float = 0.5, upper: float = 1.0) -> ants.ANTsImage:
  987. """
  988. remove all other values from the image
  989. Args:
  990. img: (ANTsImage)
  991. the input image
  992. lower: (float)
  993. lower threshold
  994. upper: (float)
  995. upper threshold
  996. Returns:
  997. image : (ANTsImage)
  998. the thresholded image
  999. """
  1000. img[img < lower] = 0
  1001. img[img > upper] = 0
  1002. return img
  1003. def remove_islands(img: ants.ANTsImage) -> ants.ANTsImage:
  1004. """ Removes parts of the mask that is not connected to the largest cluster
  1005. Args:
  1006. img (ANTsImage): the input image
  1007. Returns:
  1008. mask (ANTsImage): Image containing the largest connected component
  1009. """
  1010. clusters = ants.image_to_cluster_images(img)
  1011. mask = None
  1012. voxels = 0
  1013. for temp in clusters:
  1014. if temp.numpy().sum() > voxels:
  1015. mask = temp
  1016. voxels = temp.numpy().sum()
  1017. return mask
  1018. def subject_postprocess(mask: ants.ANTsImage, trans: ants.ANTsTransform, BoundingBox: TemplateCerebellarBoundingBox, ref: ants.ANTsImage) -> ants.ANTsImage:
  1019. """
  1020. transform the predicted cerebellum mask to the original space
  1021. Args:
  1022. mask: (ANTsImage)
  1023. the predicted cerebellum mask from the template space
  1024. trans: (ANTsTransform)
  1025. the transformation from subject space to template space
  1026. BoundingBox: (TemplateCerebellarBoundingBox)
  1027. the bounding box
  1028. ref: (ANTsImage)
  1029. the reference image
  1030. Returns:
  1031. result: (ANTsImage)
  1032. the final cerebellum mask from the subject space
  1033. """
  1034. result = BoundingBox.template2subject(mask, trans, ref)
  1035. # threshold and binarize the image
  1036. result = threshold(result)
  1037. result[result != 0] = 1
  1038. result = remove_islands(result)
  1039. return result
  1040. def isolate(t1_file: str = None, t2_file: str = None,
  1041. brain_mask_file: str = None,
  1042. label_file: str = None,
  1043. result_folder: str = None,
  1044. template: str = 'MNI152NLin6Asym',
  1045. type_of_transform: str = 'Similarity',
  1046. max_iterations: int = 5,
  1047. params: str = 'pre_trained_numpy.pkl',
  1048. save_cropped_files: bool = False,
  1049. use_q_form: bool = False,
  1050. verbose: bool = True) -> ants.ANTsImage:
  1051. """
  1052. main function for cerebellum isolation
  1053. Args:
  1054. t1_file: (string)
  1055. filename and path to T1w image, optional
  1056. t2_file: (string)
  1057. filename and path to T2w image, optional
  1058. brain_mask_file: (string)
  1059. filename and path to brain mask, optional
  1060. label_file: (string)
  1061. filename and path to label image, optional (reserved, currently has no effect)
  1062. result_folder: (string)
  1063. path to output folder (optional, otherwise it is saved to input image folder)
  1064. template: (string)
  1065. template to use (reserved)
  1066. type_of_transform: (string)
  1067. reserved for future use (see ANTspy)
  1068. max_iterations: (int)
  1069. maximum number of registration iterations (optional, default 5)
  1070. params: (string)
  1071. path to params file (reserved)
  1072. save_cropped_files: (bool)
  1073. set to True to save files cropped to window
  1074. use_q_form: (bool)
  1075. set to True to use q-form
  1076. verbose: (bool)
  1077. whether to print out status information during processing
  1078. """
  1079. if t1_file is not None:
  1080. result_folder = os.path.dirname(os.path.abspath(t1_file)) if result_folder is None else result_folder
  1081. basename = os.path.splitext(os.path.basename(t1_file))
  1082. elif t2_file is not None:
  1083. result_folder = os.path.dirname(os.path.abspath(t2_file)) if result_folder is None else result_folder
  1084. basename = os.path.splitext(os.path.basename(t2_file))
  1085. else:
  1086. raise RuntimeError('Must specify either t1_file or t2_file')
  1087. # Strip .nii or .nii.gz extension
  1088. if basename[1] == '.gz':
  1089. basename = os.path.splitext(basename[0])
  1090. basename = basename[0]
  1091. # find paramter file and template bounding box
  1092. base_dir = os.path.dirname(os.path.abspath(__file__))
  1093. params_file = os.path.join(base_dir, 'parameters', params)
  1094. BoundingBox = TemplateCerebellarBoundingBox(template_name=template)
  1095. try:
  1096. # Crop the images to the Unet input window
  1097. if verbose:
  1098. print(f"preprocessing {t1_file if t1_file is not None else t2_file}")
  1099. trans, t1_crop, t2_crop, label_crop, _, _ = subject_preprocess(t1_file=t1_file,
  1100. t2_file=t2_file,
  1101. brain_mask_file=brain_mask_file,
  1102. label_file=label_file,
  1103. BoundingBox=BoundingBox,
  1104. type_of_transform=type_of_transform,
  1105. max_iterations=max_iterations,
  1106. use_q_form=use_q_form)
  1107. if isinstance(t1_crop, ants.core.ants_image.ANTsImage):
  1108. t1_crop_data = t1_crop.numpy()
  1109. else:
  1110. t1_crop_data = t1_crop
  1111. if isinstance(t2_crop, ants.core.ants_image.ANTsImage):
  1112. t2_crop_data = t2_crop.numpy()
  1113. else:
  1114. t2_crop_data = t2_crop
  1115. if isinstance(label_crop, ants.core.ants_image.ANTsImage):
  1116. label_crop_data = label_crop.numpy()
  1117. else:
  1118. label_crop_data = label_crop
  1119. # Do a forward pass through the Unet model
  1120. if verbose:
  1121. print('isolating cerebellum using UNet model')
  1122. mask_template = predict(params_file=params_file, t1=t1_crop_data, t2=t2_crop_data)
  1123. mask_template = nib.Nifti1Image(mask_template, BoundingBox.get_cropped_affine())
  1124. mask_template = ants.from_nibabel_nifti(mask_template)
  1125. # Postprocess and transform the mask back to subject space
  1126. if verbose:
  1127. print('postprocessing')
  1128. if t1_file is not None:
  1129. mask_subject = subject_postprocess(mask=mask_template, trans=trans, BoundingBox=BoundingBox, ref=img_read(file=t1_file, use_q_form=use_q_form))
  1130. else:
  1131. mask_subject = subject_postprocess(mask=mask_template, trans=trans, BoundingBox=BoundingBox, ref=img_read(file=t2_file, use_q_form=use_q_form))
  1132. # use the original header info for the dseg mask
  1133. if t1_file is not None:
  1134. ref = nib.load(t1_file)
  1135. else:
  1136. ref = nib.load(t2_file)
  1137. mask_subject = nib.Nifti1Image(mask_subject.numpy().astype(int), affine=ref.affine, header=ref.header)
  1138. os.makedirs(result_folder, exist_ok=True)
  1139. ofname = f'{basename}_cerebellum_dseg.nii.gz'
  1140. if verbose:
  1141. print(f"saving results to {ofname}")
  1142. nib.save(mask_subject, os.path.join(result_folder, ofname))
  1143. if save_cropped_files:
  1144. if verbose:
  1145. print(f"saving intermediate results to {result_folder}")
  1146. if t1_crop is not None:
  1147. ants.image_write(t1_crop, os.path.join(result_folder, f'{basename}_crop.nii.gz'))
  1148. else:
  1149. ants.image_write(t2_crop, os.path.join(result_folder, f'{basename}_crop.nii.gz'))
  1150. ants.image_write(mask_template, os.path.join(result_folder, f'{basename}_cerebellum_crop_dseg.nii.gz'))
  1151. ants.write_transform(trans, os.path.join(result_folder, f'{basename}_trans.mat'))
  1152. except RegistrationError as e:
  1153. print(f'Caught registration error for {t1_file if t1_file is not None else t2_file} : {e}')
  1154. print(f'Isolation fails on {t1_file if t1_file is not None else t2_file}. No results were saved')
  1155. mask_subject = None
  1156. except InputError as e:
  1157. print(f'Caught input file error for {t1_file if t1_file is not None else t2_file} : {e}')
  1158. print(f'Isolation fails on {t1_file if t1_file is not None else t2_file}. No results were saved')
  1159. mask_subject = None
  1160. return img_read(file=os.path.join(result_folder, ofname), use_q_form=use_q_form) if mask_subject is not None else None
  1161. if __name__ == '__main__':
  1162. parser = argparse.ArgumentParser()
  1163. parser.add_argument('--T1', type=str, help='path to T1w image')
  1164. parser.add_argument('--T2', type=str, help='path to T2w image')
  1165. parser.add_argument('--brain_mask', type=str, help='path to brain mask image')
  1166. parser.add_argument('--label', type=str, help='path to label image (reserved, currently has no effect)')
  1167. parser.add_argument('--result_folder', type=str, help='path to save the isolation image (results will be saved to '
  1168. 'T1w image folder (or T2w image folder if no T1w image is '
  1169. 'specified))')
  1170. parser.add_argument('--template', type=str, default='MNI152NLin6Asym',
  1171. help='template for registration (MNI152NLin6Asym by '
  1172. 'default)')
  1173. parser.add_argument('--type_of_transform', type=str, default='Similarity', help='reserved for future use (see ANTspy)')
  1174. parser.add_argument('--max_iterations', type=int, default=5, help='maximum number of registration iterations (optional, default 5)')
  1175. parser.add_argument('--params', type=str, default='pre_trained.pkl', help='pretrained parameter file')
  1176. parser.add_argument('--save_cropped_files', action='store_true', help='whether to save files cropped to UNet input window')
  1177. parser.add_argument('--use_q_form', action='store_true', help='whether to use q-form')
  1178. parser.add_argument('--verbose', action='store_true', help='whether to print out status information during processing')
  1179. args = parser.parse_args()
  1180. print(args)
  1181. if args.T1 is None and args.T2 is None:
  1182. raise RuntimeError('Must specify either t1_file or t2_file')
  1183. if args.result_folder is None:
  1184. if args.T1 is None:
  1185. args.result_folder = os.path.dirname(os.path.abspath(args.T2))
  1186. else:
  1187. args.result_folder = os.path.dirname(os.path.abspath(args.T1))
  1188. isolate(t1_file=args.T1,
  1189. t2_file=args.T2,
  1190. brain_mask_file=args.brain_mask,
  1191. label_file=args.label,
  1192. result_folder=args.result_folder,
  1193. template=args.template,
  1194. type_of_transform=args.type_of_transform,
  1195. max_iterations=args.max_iterations,
  1196. params=args.params,
  1197. save_cropped_files=args.save_cropped_files,
  1198. use_q_form=args.use_q_form,
  1199. verbose=args.verbose)

isolation.py at commit 3d20055, under MIT · at the source

Overview

Authors: Yaping Wang1,2,3, Yao Li4, Bassel Arafat1,2, Vahid Ashkanichenarlogh1,2, Caroline Nettekoven1,2,5, Ana Luísa Pinho1,2,6, Carlos R. Hernandez-Castillo4, Andre F. Marquand3, Jörn Diedrichsen1,2
  1. Western Institute for Neuroscience, Western University, London, ON, Canada
  2. Department of Computer Science, Western University, London, ON, Canada
  3. Donders Institute for Brain, Cognition and Behaviour, Radboud University Medical Centre, Nijmegen, The Netherlands
  4. Faculty of Computer Science, Dalhousie University, Halifax, Canada
  5. Department of Experimental Psychology, University of Oxford, Oxford, United Kingdom
  6. Department of Psychology, Western University, London, ON, Canada
Journal: Imaging neuroscience (Cambridge, Mass.), volume 4, article IMAG.a.1323
Dates: received 21 May 2026; accepted 24 June 2026; published online 6 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1162/imag.a.1323 · PMID 42569404 · PMCID PMC13449926 · OpenAlex W7167902735
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), methods / tools (subfield)
Methods: Connectivity, Statistics, Machine learning, fMRI & imaging
Keywords: human cerebellum, MRI, deep learning, cerebellar isolation, spatial normalization, voxel-wise inference, neuroimaging toolbox, lifespan
Journal subjects: Software Toolbox
Topic: Vestibular and auditory disorders (Neurology, Neuroscience), according to OpenAlex
Funding: Raynor Cerebellum Project; Gouvernement du Canada | Canadian Institutes of Health Research (Instituts de Recherche en Santé du Canada) (PJT-191815); Canada First Research Excellence Fund (BrainsCAN); Wellcome Trust Early Career Award (306553/Z/23/Z); Junior Research Fellowship Grant from Linacre College, University of Oxford
Citations: not cited yet (Europe PMC); 57 references in the paper

Abstract

The human cerebellum plays a central role in motor, emotional, and cognitive functions, and is implicated in many brain disorders. To improve the analysis of functional and anatomical imaging from the cerebellum, we introduce SUITPy, an improved and fully revised Python implementation of the widely used SUIT toolbox. For this new version, we developed a U-Net based model to automatically isolate the cerebellum from adjacent cortical tissue, which achieves higher fidelity than existing algorithms. The isolation works robustly without manual corrections for imaging data across the lifespan. We show that isolation and subsequent normalization to a cerebellum-only template lead to a more precise alignment of cerebellar structures across participants compared to normalization using a whole-brain template. We also show the utility of the cerebellar mask to prevent contamination of cerebellar functional data from surrounding cortical structures. The toolbox also provides functionality for visualizing cerebellar data on a flatmap, along with a range of anatomical and functional cerebellar atlases, thereby offering an essential tool that enables accurate cerebellar analysis across the lifespan.

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

DiedrichsenLab/SUITPy

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 3d20055022e5339687f412da92979520c7568da2, 12 September 2026
Languages: Python (16), Jupyter (8)
Size: 95 files, 24 scripts
Software Heritage: not archived
Found in: “Data and Code Availability”
Holds: README, license file, CITATION.cff, environment (environment.yml, requirements-build-docs.txt, requirements-dev.txt, requirements-min.txt, setup.cfg, setup.py), tests, continuous integration, documentation, 8 notebooks
Tools: NiBabel (17 files), NumPy (11 files), Matplotlib (8 files), Nilearn (8 files), ANTs (4 files), pandas (3 files), Plotly (2 files), SciPy (2 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
26 files

The paper's code and data availability statement is in the Data section.

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;
  • 24 scripts, each with its path and the digest of its content;
  • 6 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 and Code Availability

The manual isolation dataset (except for the EssCer dataset) is available upon request. The SUITPy toolbox (version 2.1) can be installed via pip with source code available at https://github.com/DiedrichsenLab/SUITPy. Tutorials and documentation can be found at https://suitpy.readthedocs.io.

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, 27 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 9 authors, 8 keywords, 5 funders, 55 references.

Cite

This paper

Wang, Y., Li, Y., Arafat, B., Ashkanichenarlogh, V., Nettekoven, C., Pinho, A. L., Hernandez-Castillo, C. R., Marquand, A. F., & Diedrichsen, J. (2026). SUITPy: A Python-based toolbox for the analysis of cerebellar functional and anatomical imaging data across the human lifespan. Imaging neuroscience (Cambridge, Mass.), 4, IMAG.a.1323. https://doi.org/10.1162/imag.a.1323

BibTeX

@article{wang2026suitpy,
author = {Wang, Yaping and Li, Yao and Arafat, Bassel and Ashkanichenarlogh, Vahid and Nettekoven, Caroline and Pinho, Ana Luísa and Hernandez-Castillo, Carlos R. and Marquand, Andre F. and Diedrichsen, Jörn},
title = {{SUITPy: A Python-based toolbox for the analysis of cerebellar functional and anatomical imaging data across the human lifespan}},
journal = {Imaging neuroscience (Cambridge, Mass.)},
year = {2026},
month = aug,
volume = {4},
pages = {IMAG.a.1323},
publisher = {MIT Press},
issn = {2837-6056},
doi = {10.1162/imag.a.1323},
url = {https://doi.org/10.1162/imag.a.1323},
pmid = {42569404},
pmcid = {PMC13449926}
}

RIS

TY - JOUR
AU - Wang, Yaping
AU - Li, Yao
AU - Arafat, Bassel
AU - Ashkanichenarlogh, Vahid
AU - Nettekoven, Caroline
AU - Pinho, Ana Luísa
AU - Hernandez-Castillo, Carlos R.
AU - Marquand, Andre F.
AU - Diedrichsen, Jörn
TI - SUITPy: A Python-based toolbox for the analysis of cerebellar functional and anatomical imaging data across the human lifespan
T2 - Imaging neuroscience (Cambridge, Mass.)
J2 - Imaging Neurosci (Camb)
PY - 2026
DA - 2026/08/06
VL - 4
SP - IMAG.a.1323
SN - 2837-6056
PB - MIT Press
DO - 10.1162/imag.a.1323
UR - https://doi.org/10.1162/imag.a.1323
LA - en
ER -

CSL-JSON

{
"id": "10.1162/imag.a.1323",
"type": "article-journal",
"title": "SUITPy: A Python-based toolbox for the analysis of cerebellar functional and anatomical imaging data across the human lifespan",
"container-title": "Imaging neuroscience (Cambridge, Mass.)",
"author": [
{
"family": "Wang",
"given": "Yaping"
},
{
"family": "Li",
"given": "Yao"
},
{
"family": "Arafat",
"given": "Bassel"
},
{
"family": "Ashkanichenarlogh",
"given": "Vahid"
},
{
"family": "Nettekoven",
"given": "Caroline"
},
{
"family": "Pinho",
"given": "Ana Luísa"
},
{
"family": "Hernandez-Castillo",
"given": "Carlos R."
},
{
"family": "Marquand",
"given": "Andre F."
},
{
"family": "Diedrichsen",
"given": "Jörn"
}
],
"container-title-short": "Imaging Neurosci (Camb)",
"volume": "4",
"page": "IMAG.a.1323",
"DOI": "10.1162/imag.a.1323",
"PMID": "42569404",
"PMCID": "PMC13449926",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://doi.org/10.1162/imag.a.1323",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
6
]
]
}
}

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.21203/rs.3.rs-9326213/v1 [code]
Multi-task fMRI outperforms resting-state fMRI for revealing task-invariant organization of the human brain
Journal: Research Square (preprint)
In common: ANTs, Nilearn, Plotly, 5 other tools, 3 references, 3 authors
[2] doi:10.1038/s41467-026-72940-5 [code]
Cerebellar growth is associated with domain-specific cerebral maturation and socio-linguistic behavior.
Journal: Nature communications
In common: NiBabel, pandas, SciPy, 2 other tools, 9 references, author Jörn Diedrichsen
[3] doi:10.1016/j.neuroimage.2026.121930
Quantifying cerebellar signal detectability in MEG and EEG in epilepsy using anatomically informed source modeling.
Journal: NeuroImage
In common: 9 references
[4] doi:10.1186/s12888-026-08071-4 [code]
Cerebellar gray matter volume difference in first-episode bipolar and unipolar depression.
Journal: BMC psychiatry
In common: 8 references
[5] doi:10.1038/s41467-026-72931-6 [code]
Three parsimonious spatiotemporal patterns in cerebellum reveal individual traits in function and behavior.
Journal: Nature communications
In common: Nilearn, NiBabel, pandas, 3 other tools, 4 references
[6] doi:10.1162/imag.a.1276 [code]
High-resolution whole-brain magnetic resonance spectroscopic imaging in youth at risk for psychosis.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: ANTs, Nilearn, Plotly, 5 other tools, 1 reference
[7] doi:10.1162/imag.a.1362 [code]
Human fMRI at 11.7T: Assessing feasibility, stability, and reliability on the Iseult scanner.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: ANTs, Nilearn, NiBabel, 4 other tools, methods / tools, 2 references
[8] doi:10.1038/s41592-026-03159-x [code]
Siibra: a software tool suite for realizing a Multilevel Human Brain Atlas from complex data resources.
Journal: Nature methods
In common: Nilearn, Plotly, NiBabel, 4 other tools, methods / tools, 2 references
[9] doi:10.1038/s41467-026-71151-2 [code]
Common and distinct neural correlates of social interaction processing and theory of mind in narratives.
Journal: Nature communications
In common: ANTs, Nilearn, NiBabel, 4 other tools, 2 references
[10] doi:10.1162/imag.a.1269 [code]
From early to contemporary normative modeling: Mapping individual differences in neurophysiological signals.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Nilearn, Plotly, NiBabel, 4 other tools, 2 references

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.