SUITPy: A Python-based toolbox for the analysis of cerebellar functional and anatomical imaging data across the human lifespan.
The 6 matches
- [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] § 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] § 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] § Methods › Cerebellar normalization ↔ SUITPy/normalization.py, lines 18–67 · score 0.61 · antsRegistrationSyN, MNI152NLin2009cSymC, space, template, masking, cerebellum
- [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] § 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
- """
- Cerebellar Isolation using a Unet model
- @authors: Yao Li, Carlos Hernandez-Castillo, Joern Diedrichsen
- """
- import sys
- import argparse
- import os
- import nibabel as nib
- import ants
- import numpy as np
- import nitools
- from typing import Tuple, Union
- import pickle
- import warnings
- class _Conv3dN:
- """
- numpy implementation of 3D convolution layer
- """
- def __init__(
- self,
- in_channels: int,
- out_channels: int,
- kernel_size: Union[int, Tuple[int, int, int]],
- stride: Union[int, Tuple[int, int, int]] = (1, 1, 1),
- padding: Union[int, Tuple[int, int, int]] = (0, 0, 0),
- dilation: Union[int, Tuple[int, int, int]] = (1, 1, 1),
- bias: bool = True,
- padding_mode: str = 'zeros',
- ):
- self.in_channels = in_channels
- self.out_channels = out_channels
- self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3
- self.stride = stride if isinstance(stride, tuple) else (stride,) * 3
- self.padding = padding if isinstance(padding, tuple) else (padding,) * 3
- self.dilation = dilation if isinstance(dilation, tuple) else (dilation,) * 3
- self.bias = bias
- self.padding_mode = padding_mode
- self.weight = np.random.randn(out_channels, in_channels, *self.kernel_size)
- if bias:
- self.bias_term = np.zeros(out_channels)
- else:
- self.bias_term = None
- def _pad_input(self, x: np.ndarray) -> np.ndarray:
- if all(p == 0 for p in self.padding):
- return x
- pd, ph, pw = self.padding
- if self.padding_mode == 'zeros':
- return np.pad(x,
- ((0, 0), (0, 0),
- (pd, pd), (ph, ph), (pw, pw)),
- mode='constant')
- elif self.padding_mode == 'reflect':
- return np.pad(x,
- ((0, 0), (0, 0),
- (pd, pd), (ph, ph), (pw, pw)),
- mode='reflect')
- else:
- raise NotImplementedError(f"Padding mode {self.padding_mode} not implemented")
- def _im2col(self, x: np.ndarray) -> Tuple[np.ndarray, Tuple[int, int, int]]:
- batch_size, in_channels, depth, height, width = x.shape
- kd, kh, kw = self.kernel_size
- sd, sh, sw = self.stride
- pd, ph, pw = self.padding
- dd, dh, dw = self.dilation
- x_padded = self._pad_input(x)
- out_depth = (depth + 2 * pd - dd * (kd - 1) - 1) // sd + 1
- out_height = (height + 2 * ph - dh * (kh - 1) - 1) // sh + 1
- out_width = (width + 2 * pw - dw * (kw - 1) - 1) // sw + 1
- cols = np.zeros((batch_size, in_channels, kd, kh, kw, out_depth, out_height, out_width))
- for d in range(kd):
- for h in range(kh):
- for w in range(kw):
- d_start = d * dd
- h_start = h * dh
- w_start = w * dw
- d_slice = slice(d_start, d_start + out_depth * sd, sd)
- h_slice = slice(h_start, h_start + out_height * sh, sh)
- w_slice = slice(w_start, w_start + out_width * sw, sw)
- cols[:, :, d, h, w, :, :, :] = x_padded[:, :, d_slice, h_slice, w_slice]
- # Reshape for matrix multiplication (batch, out_d, out_h, out_w, in_ch, kd, kh, kw)
- cols = cols.transpose(0, 5, 6, 7, 1, 2, 3, 4)
- cols = cols.reshape(batch_size * out_depth * out_height * out_width, -1)
- return cols, (out_depth, out_height, out_width)
- def load_state(self, weight: np.ndarray, bias_term: np.ndarray):
- self.weight = weight
- self.bias_term = bias_term
- def forward(self, x: np.ndarray) -> np.ndarray:
- if x.ndim != 5:
- raise ValueError(f"Input must have 5 dimensions (N, C, D, H, W), got {x.ndim}")
- batch_size, in_channels, depth, height, width = x.shape
- if in_channels != self.in_channels:
- raise ValueError(f"Expected {self.in_channels} input channels, got {in_channels}")
- cols, (out_depth, out_height, out_width) = self._im2col(x)
- # Reshape weights for matrix multiplication (out_ch, in_ch * kd * kh * kw)
- weight_flat = self.weight.reshape(self.out_channels, -1)
- # Perform convolution via matrix multiplication (batch * out_d * out_h * out_w, out_ch)
- output_flat = cols @ weight_flat.T
- output = output_flat.reshape(batch_size, out_depth, out_height, out_width, self.out_channels)
- output = output.transpose(0, 4, 1, 2, 3) # (batch, out_ch, out_d, out_h, out_w)
- if self.bias:
- output += self.bias_term.reshape(1, -1, 1, 1, 1)
- return output
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- class _ConvTranspose3dN:
- """
- numpy implementation of 3D transpose convolution layer
- """
- def __init__(
- self,
- in_channels: int,
- out_channels: int,
- kernel_size: Union[int, Tuple[int, int, int]],
- stride: Union[int, Tuple[int, int, int]] = (1, 1, 1),
- padding: Union[int, Tuple[int, int, int]] = (0, 0, 0),
- output_padding: Union[int, Tuple[int, int, int]] = (0, 0, 0),
- dilation: Union[int, Tuple[int, int, int]] = (1, 1, 1),
- bias: bool = True,
- padding_mode: str = "zeros",
- ):
- self.in_channels = in_channels
- self.out_channels = out_channels
- self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3
- self.stride = stride if isinstance(stride, tuple) else (stride,) * 3
- self.padding = padding if isinstance(padding, tuple) else (padding,) * 3
- self.output_padding = output_padding if isinstance(output_padding, tuple) else (output_padding,) * 3
- self.dilation = dilation if isinstance(dilation, tuple) else (dilation,) * 3
- self.bias = bias
- self.padding_mode = padding_mode
- self.weight = np.random.randn(in_channels, out_channels, *self.kernel_size)
- if bias:
- self.bias_term = np.zeros(out_channels)
- else:
- self.bias_term = None
- def load_state(self, weight: np.ndarray, bias_term: np.ndarray):
- self.weight = weight
- self.bias_term = bias_term
- def _dilate_input(self, x: np.ndarray) -> np.ndarray:
- batch_size, in_channels, depth, height, width = x.shape
- sd, sh, sw = self.stride
- dilated_depth = (depth - 1) * sd + 1
- dilated_height = (height - 1) * sh + 1
- dilated_width = (width - 1) * sw + 1
- dilated_x = np.zeros((batch_size, in_channels, dilated_depth, dilated_height, dilated_width))
- # Insert original values at strided positions
- dilated_x[:, :, ::sd, ::sh, ::sw] = x
- return dilated_x
- def _apply_padding(self, x: np.ndarray) -> np.ndarray:
- pd, ph, pw = self.padding
- kd, kh, kw = self.kernel_size
- pad_d_before = kd - 1 - pd
- pad_d_after = kd - 1 - pd + self.output_padding[0]
- pad_h_before = kh - 1 - ph
- pad_h_after = kh - 1 - ph + self.output_padding[1]
- pad_w_before = kw - 1 - pw
- pad_w_after = kw - 1 - pw + self.output_padding[2]
- padding = ((0, 0), (0, 0),
- (pad_d_before, pad_d_after),
- (pad_h_before, pad_h_after),
- (pad_w_before, pad_w_after))
- if self.padding_mode == "zeros":
- return np.pad(x, padding, mode='constant')
- else:
- raise NotImplementedError(f"Padding mode {self.padding_mode} not implemented")
- def _im2col(self, x: np.ndarray) -> Tuple[np.ndarray, Tuple[int, int, int]]:
- batch_size, in_channels, depth, height, width = x.shape
- kd, kh, kw = self.kernel_size
- sd, sh, sw = self.stride
- pd, ph, pw = self.padding
- opd, oph, opw = self.output_padding
- dd, dh, dw = self.dilation
- dilated_x = self._dilate_input(x)
- padded_x = self._apply_padding(dilated_x)
- # Output dimensions
- out_depth = (depth - 1) * sd - 2 * pd + dd * (kd - 1) + opd + 1
- out_height = (height - 1) * sh - 2 * ph + dh * (kh - 1) + oph + 1
- out_width = (width - 1) * sw - 2 * pw + dw * (kw - 1) + opw + 1
- cols = np.zeros((batch_size, in_channels, kd, kh, kw, out_depth, out_height, out_width))
- for d in range(kd):
- for h in range(kh):
- for w in range(kw):
- d_start = d * dd
- h_start = h * dh
- w_start = w * dw
- d_slice = slice(d_start, d_start + out_depth)
- h_slice = slice(h_start, h_start + out_height)
- w_slice = slice(w_start, w_start + out_width)
- cols[:, :, d, h, w, :, :, :] = padded_x[:, :, d_slice, h_slice, w_slice]
- # Reshape for matrix multiplication (batch, out_d, out_h, out_w, in_ch, kd, kh, kw)
- cols = cols.transpose(0, 5, 6, 7, 1, 2, 3, 4)
- cols = cols.reshape(batch_size * out_depth * out_height * out_width, -1)
- return cols, (out_depth, out_height, out_width)
- def forward(self, x: np.ndarray) -> np.ndarray:
- if x.ndim != 5:
- raise ValueError(f"Input must have 5 dimensions (N, C, D, H, W), got {x.ndim}")
- batch_size, in_channels, depth, height, width = x.shape
- if in_channels != self.in_channels:
- raise ValueError(f"Expected {self.in_channels} input channels, got {in_channels}")
- cols, (out_depth, out_height, out_width) = self._im2col(x)
- # Reshape weights for matrix multiplication (out_ch, in_ch * kd * kh * kw)
- weight = self.weight.transpose(1, 0, 2, 3, 4)
- weight_flip = np.flip(weight, (2, 3, 4))
- weight_flat = weight_flip.reshape(self.out_channels, -1)
- # Perform convolution via matrix multiplication (batch * out_d * out_h * out_w, out_ch)
- output_flat = cols @ weight_flat.T
- output = output_flat.reshape(batch_size, out_depth, out_height, out_width, self.out_channels)
- output = output.transpose(0, 4, 1, 2, 3) # (batch, out_ch, out_d, out_h, out_w)
- if self.bias:
- output += self.bias_term.reshape(1, -1, 1, 1, 1)
- return output
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- class _InstanceNorm3dN:
- """
- numpy implementation of 3D instance normalization layer
- """
- def __init__(
- self,
- num_features: int,
- eps: float = 1e-5,
- momentum: float = 0.1,
- affine: bool = False,
- track_running_stats: bool = False
- ):
- self.num_features = num_features
- self.eps = eps
- self.momentum = momentum
- self.affine = affine
- self.track_running_stats = track_running_stats
- if self.affine:
- self.weight = np.ones(num_features)
- self.bias = np.zeros(num_features)
- else:
- self.weight = None
- self.bias = None
- if self.track_running_stats:
- raise RuntimeError('track_running_stats currently not supported.')
- # self.running_mean = np.zeros(num_features)
- # self.running_var = np.ones(num_features)
- # self.num_batches_tracked = 0
- else:
- self.running_mean = None
- self.running_var = None
- self.num_batches_tracked = None
- def _check_input_dim(self, x: np.ndarray) -> None:
- if x.ndim != 5:
- raise ValueError(f"Expected 5D input (got {x.ndim}D input)")
- if x.shape[1] != self.num_features:
- raise ValueError(f"Expected {self.num_features} channels in input (got {x.shape[1]} channels)")
- def _compute_instance_stats(self, x: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
- """Compute mean and variance per instance per channel across spatial dimensions (D, H, W)"""
- n, c, d, h, w = x.shape
- x_reshaped = x.reshape(n, c, -1)
- mean = np.mean(x_reshaped, axis=2)
- var = np.var(x_reshaped, axis=2)
- return mean, var
- def _normalize(self, x: np.ndarray, mean: np.ndarray, var: np.ndarray) -> np.ndarray:
- """Normalize input using computed statistics"""
- n, c, d, h, w = x.shape
- mean = mean.reshape(n, c, 1, 1, 1)
- var = var.reshape(n, c, 1, 1, 1)
- x_normalized = (x - mean) / np.sqrt(var + self.eps)
- return x_normalized
- def forward(self, x: np.ndarray) -> np.ndarray:
- self._check_input_dim(x)
- if self.track_running_stats:
- n, c, d, h, w = x.shape
- mean = self.running_mean.reshape(1, c, 1, 1, 1)
- var = self.running_var.reshape(1, c, 1, 1, 1)
- else:
- # If not tracking running stats, compute instance stats even during inference
- mean, var = self._compute_instance_stats(x)
- if not self.track_running_stats:
- x_normalized = self._normalize(x, mean, var)
- else:
- x_normalized = (x - mean) / np.sqrt(var + self.eps)
- if self.affine:
- weight = self.weight.reshape(1, -1, 1, 1, 1)
- bias = self.bias.reshape(1, -1, 1, 1, 1)
- x_normalized = x_normalized * weight + bias
- return x_normalized
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- def load_state_dict(self, state_dict: dict) -> None:
- if self.affine:
- self.weight = state_dict['weight']
- self.bias = state_dict['bias']
- if self.track_running_stats:
- self.running_mean = state_dict['running_mean']
- self.running_var = state_dict['running_var']
- self.num_batches_tracked = state_dict['num_batches_tracked']
- class _LeakyReLUN:
- """
- numpy implementation of Leaky ReLUN layer
- """
- def __init__(self, negative_slope: float = 0.01):
- self.negative_slope = negative_slope
- def forward(self, x: np.ndarray) -> np.ndarray:
- return np.maximum(0, x) + self.negative_slope * np.minimum(0, x)
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- class _SoftmaxN:
- """
- numpy implementation of Softmax layer
- """
- def __init__(self, dim: int = -1,):
- self.dim = dim
- def forward(self, x: np.ndarray) -> np.ndarray:
- dim = self.dim
- x_shifted = x - np.max(x, axis=dim, keepdims=True)
- exp_x = np.exp(x_shifted)
- softmax_output = exp_x / np.sum(exp_x, axis=dim, keepdims=True)
- return softmax_output
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- class _ConvNetN:
- """
- convolution block in Unet
- """
- def __init__(self, in_channels: int, out_channels: int):
- """
- Args:
- in_channels: (int)
- the number of input channels
- out_channels: (int)
- the number of output channels
- """
- self.conv1 = _Conv3dN(in_channels=in_channels, out_channels=out_channels, kernel_size=3, padding=1)
- self.norm1 = _InstanceNorm3dN(num_features=out_channels, affine=True, track_running_stats=False)
- self.activate1= _LeakyReLUN()
- self.conv2 = _Conv3dN(in_channels=out_channels, out_channels=out_channels, kernel_size=3, padding=1)
- self.norm2 = _InstanceNorm3dN(num_features=out_channels, affine=True, track_running_stats=False)
- self.activate2 = _LeakyReLUN()
- def load_state(
- self,
- c_weight1: np.ndarray,
- c_bias1: np.ndarray,
- n_state_dict1: dict,
- c_weight2: np.ndarray,
- c_bias2: np.ndarray,
- n_state_dict2: dict,
- ):
- """
- load the previously trained parameters
- Args:
- c_weight1: (ndarray)
- the weight of the first convolution layer
- c_bias1: (ndarray)
- the bias of the first convolution layer
- n_state_dict1: (dict)
- the state dictionary of the first instance normalization layer
- c_weight2: (ndarray)
- the weight of the second convolution layer
- c_bias2: (ndarray)
- the bias of the second convolution layer
- n_state_dict2: (dict)
- the state dictionary of the second instance normalization layer
- Returns:
- """
- self.conv1.load_state(weight=c_weight1, bias_term=c_bias1)
- self.norm1.load_state_dict(state_dict=n_state_dict1)
- self.conv2.load_state(weight=c_weight2, bias_term=c_bias2)
- self.norm2.load_state_dict(state_dict=n_state_dict2)
- def forward(self, x: np.ndarray) -> np.ndarray:
- """
- calculation
- Args:
- x: (ndarray)
- input
- Returns:
- y: (ndarray)
- output
- """
- result1 = self.conv1(x)
- result2 = self.norm1(result1)
- result3 = self.activate1(result2)
- result4 = self.conv2(result3)
- result5 = self.norm2(result4)
- result6 = self.activate2(result5)
- return result6
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- class _DownSampleN:
- """
- downsample block in unet
- """
- def __init__(self, channels: int):
- """
- Args:
- channels: (int)
- the number of input channels
- """
- self.conv1 = _Conv3dN(in_channels=channels, out_channels=channels, kernel_size=3, stride=2, padding=1)
- self.norm1 = _InstanceNorm3dN(num_features=channels, affine=True, track_running_stats=False)
- self.activate1 = _LeakyReLUN()
- def load_state(
- self,
- c_weight: np.ndarray,
- c_bias: np.ndarray,
- n_state_dict: dict
- ):
- """
- load the previously trained parameters
- Args:
- c_weight: (ndarray)
- the weight of the convolution layer
- c_bias: (ndarray)
- the bias of the convolution layer
- n_state_dict: (dict)
- the state dictionary of the instance normalization layer
- """
- self.conv1.load_state(weight=c_weight, bias_term=c_bias)
- self.norm1.load_state_dict(state_dict=n_state_dict)
- def forward(self, x: np.ndarray) -> np.ndarray:
- """
- calculation
- Args:
- x: (ndarray)
- input
- Returns:
- y: (ndarray)
- output
- """
- result1 = self.conv1(x)
- result2 = self.norm1(result1)
- result3 = self.activate1(result2)
- return result3
- def __call__(self, x: np.ndarray) -> np.ndarray:
- return self.forward(x)
- class _UpSampleN:
- """
- upsample block in unet
- """
- def __init__(self, channels: int):
- """
- Args:
- channels: (int)
- the number of input channels
- """
- self.conv1 = _ConvTranspose3dN(in_channels=channels, out_channels=channels, kernel_size=2, stride=2, padding=0)
- self.norm1 = _InstanceNorm3dN(num_features=channels, affine=True, track_running_stats=False)
- self.activate1 = _LeakyReLUN()
- self.conv2 = _Conv3dN(in_channels=channels * 2, out_channels=channels, kernel_size=3, padding=1)
- self.norm2 = _InstanceNorm3dN(num_features=channels, affine=True, track_running_stats=False)
- self.activate2 = _LeakyReLUN()
- def load_state(
- self,
- c_weight1: np.ndarray,
- c_bias1: np.ndarray,
- n_state_dict1: dict,
- c_weight2: np.ndarray,
- c_bias2: np.ndarray,
- n_state_dict2: dict
- ):
- """
- load the previously trained parameters
- Args:
- c_weight1: (ndarray)
- the weight of the transpose convolution layer
- c_bias1: (ndarray)
- the bias of the transpose convolution layer
- n_state_dict1: (dict)
- the state dictionary of the first instance normalization layer
- c_weight2: (ndarray)
- the weight of the convolution layer
- c_bias2: (ndarray)
- the bias of the convolution layer
- n_state_dict2: (dict)
- the state dictionary of the second instance normalization layer
- Returns:
- """
- self.conv1.load_state(weight=c_weight1, bias_term=c_bias1)
- self.norm1.load_state_dict(state_dict=n_state_dict1)
- self.conv2.load_state(weight=c_weight2, bias_term=c_bias2)
- self.norm2.load_state_dict(state_dict=n_state_dict2)
- def forward(self, x: np.ndarray, feature_map: np.ndarray) -> np.ndarray:
- """
- calculation
- Args:
- x: (ndarray)
- input
- feature_map: (ndarray)
- The feature map form the corresponding skip connection
- Returns:
- y: (ndarray)
- output
- """
- result1 = self.conv1(x)
- result2 = self.norm1(result1)
- result3 = self.activate1(result2)
- result3 = np.concatenate([result3, feature_map], axis=1)
- result4 = self.conv2(result3)
- result5 = self.norm2(result4)
- result6 = self.activate2(result5)
- return result6
- def __call__(self, x: np.ndarray, feature_map: np.ndarray) -> np.ndarray:
- return self.forward(x, feature_map)
- class UNetN:
- """
- Unet model architecture.
- """
- def __init__(self, init_features: int = 16):
- """
- Args:
- init_features: int
- Number of filters in the first convolutional block.
- """
- self.enc1 = _ConvNetN(in_channels=2, out_channels=init_features)
- self.down1 = _DownSampleN(channels=init_features)
- self.enc2 = _ConvNetN(in_channels=init_features, out_channels=init_features * 2)
- self.down2 = _DownSampleN(init_features * 2)
- self.enc3 = _ConvNetN(in_channels=init_features * 2, out_channels=init_features * 4)
- self.down3 = _DownSampleN(init_features * 4)
- self.enc4 = _ConvNetN(in_channels=init_features * 4, out_channels=init_features * 8)
- self.down4 = _DownSampleN(init_features * 8)
- self.bottleneck = _ConvNetN(in_channels=init_features * 8, out_channels=init_features * 8)
- self.up4 = _UpSampleN(init_features * 8)
- self.dec4 = _ConvNetN(in_channels=init_features * 8, out_channels=init_features * 4)
- self.up3 = _UpSampleN(init_features * 4)
- self.dec3 = _ConvNetN(in_channels=init_features * 4, out_channels=init_features * 2)
- self.up2 = _UpSampleN(init_features * 2)
- self.dec2 = _ConvNetN(in_channels=init_features * 2, out_channels=init_features)
- self.up1 = _UpSampleN(init_features)
- self.dec1 = _ConvNetN(in_channels=init_features, out_channels=2)
- self.out = _SoftmaxN(dim=1)
- def load_state_dict(self, state_dict: dict):
- """
- Load the state dict of the Unet model.
- Args:
- state_dict: (dict)
- The state dict of the Unet model
- """
- 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}
- downs = {'down1': self.down1, 'down2': self.down2, 'down3': self.down3, 'down4': self.down4}
- ups = {'up1': self.up1, 'up2': self.up2, 'up3': self.up3, 'up4': self.up4}
- for c in convs:
- c_weight1 = state_dict[f'{c}.layer.0.weight']
- c_bias1 = state_dict[f'{c}.layer.0.bias']
- n1 = {
- 'weight': state_dict[f'{c}.layer.1.weight'],
- 'bias': state_dict[f'{c}.layer.1.bias'],
- # 'running_mean': state_dict[f'{c}.layer.1.running_mean'],
- # 'running_var': state_dict[f'{c}.layer.1.running_var'],
- # 'num_batches_tracked': state_dict[f'{c}.layer.1.num_batches_tracked']
- }
- c_weight2 = state_dict[f'{c}.layer.3.weight']
- c_bias2 = state_dict[f'{c}.layer.3.bias']
- n2 = {
- 'weight': state_dict[f'{c}.layer.4.weight'],
- 'bias': state_dict[f'{c}.layer.4.bias'],
- # 'running_mean': state_dict[f'{c}.layer.4.running_mean'],
- # 'running_var': state_dict[f'{c}.layer.4.running_var'],
- # 'num_batches_tracked': state_dict[f'{c}.layer.4.num_batches_tracked']
- }
- 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)
- for d in downs:
- c_weight = state_dict[f'{d}.layer.0.weight']
- c_bias = state_dict[f'{d}.layer.0.bias']
- n = {
- 'weight': state_dict[f'{d}.layer.1.weight'],
- 'bias': state_dict[f'{d}.layer.1.bias'],
- # 'running_mean': state_dict[f'{d}.layer.1.running_mean'],
- # 'running_var': state_dict[f'{d}.layer.1.running_var'],
- # 'num_batches_tracked': state_dict[f'{d}.layer.1.num_batches_tracked']
- }
- downs[d].load_state(c_weight=c_weight, c_bias=c_bias, n_state_dict=n)
- for u in ups:
- c_weight1 = state_dict[f'{u}.up.0.weight']
- c_bias1 = state_dict[f'{u}.up.0.bias']
- n1 = {
- 'weight': state_dict[f'{u}.up.1.weight'],
- 'bias': state_dict[f'{u}.up.1.bias'],
- # 'running_mean': state_dict[f'{u}.up.1.running_mean'],
- # 'running_var': state_dict[f'{u}.up.1.running_var'],
- # 'num_batches_tracked': state_dict[f'{u}.up.1.num_batches_tracked']
- }
- c_weight2 = state_dict[f'{u}.layer.0.weight']
- c_bias2 = state_dict[f'{u}.layer.0.bias']
- n2 = {
- 'weight': state_dict[f'{u}.layer.1.weight'],
- 'bias': state_dict[f'{u}.layer.1.bias'],
- # 'running_mean': state_dict[f'{u}.layer.1.running_mean'],
- # 'running_var': state_dict[f'{u}.layer.1.running_var'],
- # 'num_batches_tracked': state_dict[f'{u}.layer.1.num_batches_tracked']
- }
- 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)
- def forward(self, t1: np.ndarray, t2: np.ndarray) -> np.ndarray:
- """
- Unet calculation
- Args:
- t1: (ndarray)
- the cropped T1w image
- t2: (ndarray)
- the cropped T2w image
- Returns:
- pm: (ndarray)
- the cerebellar probability map
- """
- result_ini = np.concatenate((t1, t2), axis=1)
- result_1 = self.enc1(result_ini)
- result_2 = self.enc2(self.down1(result_1))
- result_3 = self.enc3(self.down2(result_2))
- result_4 = self.enc4(self.down3(result_3))
- result_5 = self.bottleneck(self.down4(result_4))
- out_1 = self.dec4(self.up4(result_5, result_4))
- out_2 = self.dec3(self.up3(out_1, result_3))
- out_3 = self.dec2(self.up2(out_2, result_2))
- out_4 = self.dec1(self.up1(out_3, result_1))
- return self.out(out_4)
- def __call__(self, t1: np.ndarray, t2: np.ndarray) -> np.ndarray:
- return self.forward(t1, t2)
- def _load_model(params_file: str):
- """
- load model with pretrained weights
- Args:
- params_file: (string)
- path to the pretrained weights
- Returns:
- net: (Unet)
- the pretrained model
- """
- net = UNetN()
- if os.path.exists(params_file):
- with open(params_file, "rb") as f:
- params = pickle.load(f)
- net.load_state_dict(params)
- else:
- raise RuntimeError('fail to load pre-trained parameters')
- return net
- def predict(params_file: str, t1: np.ndarray = None, t2: np.ndarray = None) -> np.ndarray:
- """
- Run a prediction on a single subject using a trained UNet model
- Args:
- params_file: (string)
- filename of the pretrained weights
- t1: (ndarray)
- Numpy array of T1w cerebellar image (after cropping)
- t2: (ndarray)
- Numpy array of T2w cerebellar image (after cropping)
- Returns:
- mask: (ndarray)
- the 3D numpy array of predicted mask (template space)
- """
- net = _load_model(params_file)
- if t1 is None:
- t1 = np.zeros((128, 128, 128))
- else:
- t1 = (t1 - t1.mean()) / t1.std()
- if t2 is None:
- t2 = np.zeros((128, 128, 128))
- else:
- t2 = (t2 - t2.mean()) / t2.std()
- t1, t2 = t1.reshape((1, 1, 128, 128, 128)), t2.reshape((1, 1, 128, 128, 128))
- mask = net(t1, t2)
- return mask[0][0]
- class InputError(Exception):
- def __init__(self, message):
- super().__init__(message)
- def img_read(file: str, use_q_form: bool = False, verbose: bool = False) -> ants.ANTsImage:
- """
- basic function to read a nifti image as an ANTs image.
- checks for consistency of s-form and q-form and
- uses s-form by default, or q-form if use_q_form is set to True.
- Args:
- file: (string)
- image path
- use_q_form: (bool)
- set to True to use q-form
- verbose: (bool)
- whether to print warning info
- Returns:
- img (ANTs image)
- An ANTs image
- """
- nib_img = nib.load(file)
- s_form, _ = nib_img.get_sform(coded=True)
- q_form, _ = nib_img.get_qform(coded=True)
- if s_form is None and q_form is None:
- raise InputError(f'Both S-form and Q-form are None in {file}')
- elif s_form is not None and q_form is None:
- if use_q_form:
- raise InputError(f'Q-form is None in {file}, Please set use_q_form=False')
- else:
- nib_img.set_qform(s_form)
- elif s_form is None and q_form is not None:
- if not use_q_form:
- if verbose:
- print(f'Warning: S-form is None in {file}. Using Q-form.')
- nib_img.set_sform(q_form)
- else:
- if (s_form != q_form).any():
- if not use_q_form:
- if verbose:
- 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.')
- nib_img.set_qform(s_form)
- else:
- if verbose:
- print(f'Warning: S- and Q-form indicate different image orientations in {file}. Using Q-form.')
- nib_img.set_sform(q_form)
- new_img = ants.from_nibabel_nifti(nib_img)
- return new_img
- def normalized_mutual_information(image1: np.ndarray, image2: np.ndarray, bins:int = 100) -> float:
- """
- Compute the normalized mutual information between two images
- Args:
- image1: (ndarray)
- First image
- image2: (ndarray)
- Second image
- bins: (int)
- number of bins
- Returns:
- mi: (float)
- normalized mutual information between image1 and image2
- """
- joint_hist, _, _ = np.histogram2d(
- image1.ravel(),
- image2.ravel(),
- bins=bins
- )
- joint_prob = joint_hist / np.sum(joint_hist)
- p_x = np.sum(joint_prob, axis=1)
- p_y = np.sum(joint_prob, axis=0)
- # H(X) = -Σ p(x) log p(x)
- h_x = -1 * np.sum(p_x[p_x > 0] * np.log2(p_x[p_x > 0]))
- # H(Y) = -Σ p(y) log p(y)
- h_y = -1 * np.sum(p_y[p_y > 0] * np.log2(p_y[p_y > 0]))
- # H(X,Y) = -ΣΣ p(x,y) log p(x,y)
- h_xy = -1 * np.sum(joint_prob[joint_prob > 0] * np.log2(joint_prob[joint_prob > 0]))
- mi = (h_x + h_y) / h_xy
- return mi
- 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]:
- """
- register the image to this template. The function runs ants registration multiple time and find the best registration results evaluated using normalized mutual information.
- Args:
- img: (ANTsImage)
- image to be registered
- brain_mask: (ANTsImage)
- brain mask of the image
- template: (string)
- 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)
- type_of_transform: (string)
- transform type (Affine by default, check ANTsPY[https://antspy.readthedocs.io/en/latest/registration.html] for details)
- max_iterations: (int)
- maximum number of iterations
- mi_lower: (float)
- 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.
- mi_upper: (float)
- 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.
- Returns:
- trans: (ANTsTransform)
- the transformation from the subject space to the template space
- status: (int)
- the status of the registration (0: for success mi>upper, 1: found one of mi>mi_lower, 2: failure)
- """
- base_dir = os.path.dirname(__file__)
- template = ants.image_read(os.path.join(base_dir, f'templates/tpl-{template_name}_T1w.nii.gz'))
- template_brain_mask = ants.image_read(os.path.join(base_dir, f'templates/tpl-{template_name}-brain_mask.nii.gz'))
- template_brain = ants.mask_image(template, template_brain_mask)
- cnt = 0
- status = 2
- trans_final = None
- best_mi = 0
- while cnt < max_iterations:
- if brain_mask is not None:
- brain = ants.mask_image(img, brain_mask)
- result = ants.registration(fixed=template_brain, moving=brain, type_of_transform=type_of_transform)
- trans = ants.read_transform(result['fwdtransforms'][0])
- else:
- result = ants.registration(fixed=template, moving=img, type_of_transform=type_of_transform)
- trans = ants.read_transform(result['fwdtransforms'][0])
- img_transformed = ants.apply_ants_transform_to_image(trans, img, template)
- brain_transformed = ants.mask_image(img_transformed, template_brain_mask)
- mi = normalized_mutual_information(brain_transformed.numpy(), template_brain.numpy())
- if mi > best_mi:
- best_mi = mi
- trans_final = trans
- if mi > mi_upper:
- status = 0
- break
- if mi > mi_lower:
- status = 1
- cnt += 1
- return trans_final, status
- class RegistrationError(Exception):
- def __init__(self, message):
- super().__init__(message)
- class TemplateCerebellarBoundingBox:
- """
- Basic cerebellar bounding box class, which defines the cropped area.
- All other template implementations should be registered to this template.
- """
- def __init__(self, template_name: str ='MNI152NLin6Asym', bounding_box: np.ndarray = None, cerebellar_center: np.ndarray = None, cropped_size: np.ndarray = None):
- """
- Create a bounding box
- Args:
- template_name: (string)
- The name of template used. It uses MNI152NLin6Asym by default.
- bounding_box: (ndarray)
- cerebellar bounding box in MNI space (in mm) (reserved for future development)
- cerebellar_center: (ndarray)
- (reserved for future development)
- cropped_size: (ndarray)
- (reserved for future development)
- """
- self.template_name = template_name
- self.cropped_size = (128, 128, 128)
- if bounding_box is not None:
- self.bounding_box = bounding_box
- else:
- if cerebellar_center is not None and cropped_size is not None:
- # reserved for future development
- pass
- else:
- self.bounding_box = np.array([[64, -114, -88], [-64, 14, 40]])
- self.lowerleft = self.bounding_box[0]
- self.upperright = self.bounding_box[1]
- base_dir = os.path.dirname(__file__)
- self.template = ants.image_read(os.path.join(base_dir, f'templates/tpl-{self.template_name}_T1w.nii.gz'))
- self.nib_template = nib.load(os.path.join(base_dir, f'templates/tpl-{self.template_name}_T1w.nii.gz'))
- self.brainmask = ants.image_read(os.path.join(base_dir, f'templates/tpl-{self.template_name}-brain_mask.nii.gz'))
- self.brain = ants.mask_image(self.template, self.brainmask)
- self.affine = nib.load(os.path.join(base_dir, f'templates/tpl-{self.template_name}_T1w.nii.gz')).affine
- def get_crop_indices(self) -> np.ndarray:
- """
- calculate the lower left and upper right indices of the cropped area (in voxels).
- Returns:
- indices: (ndarray)
- a 2 * 3 ndarray consists of two vertices defining the bounding box
- """
- return nitools.coords_to_voxelidxs(self.bounding_box.T, self.nib_template).T
- def get_cropped_affine(self) -> np.ndarray:
- """
- get the cropped area affine
- Returns:
- affine: (ndarray)
- The affine form for the cropped image
- """
- # This function needs to fix. It will fail if the template affine is not diagonal
- affine = np.diag([self.affine[0, 0], self.affine[1, 1], self.affine[2, 2], 1])
- affine[0, 3] = abs(affine[0, 0]) * self.lowerleft[0]
- affine[1, 3] = abs(affine[1, 1]) * self.lowerleft[1]
- affine[2, 3] = abs(affine[2, 2]) * self.lowerleft[2]
- return affine
- 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]:
- """
- register the image to this template
- Args:
- img: (ANTsImage)
- image to be registered
- brain_mask: (ANTsImage)
- brain mask of the image
- type_of_transform: (string)
- transform type (Affine by default, check ANTsPY[https://antspy.readthedocs.io/en/latest/registration.html] for details)
- max_iterations: (int)
- maximum number of iterations
- mi_lower: (float)
- lower boundary of normalized mutual information
- mi_upper: (float)
- upper boundary of normalized mutual information
- Returns:
- trans: (ANTsTransform)
- the transformation from the subject space to the template space
- status: (int)
- the status of the registration (0 for success, 1 for uncertainty, 2 for failure)
- """
- 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)
- return trans, status
- def crop(self, img: ants.ANTsImage, trans: ants.ANTsTransform = None) -> Tuple[ants.ANTsImage, ants.ANTsImage]:
- """
- Crop the cerebellar area using the bounding box.
- Args:
- img: (ANTsImage)
- image to be cropped
- trans: (ANTsTransform)
- transformation matrix from the image space to the template space (only use it if img is not in the MNI template)
- Returns:
- cropped_img: (ANTsImage)
- cropped image
- img : (ANTsImage)
- the whole transformed image
- """
- start_indices, end_indices = self.get_crop_indices()
- if trans is not None:
- img = ants.apply_ants_transform_to_image(trans, img, self.template)
- return ants.crop_indices(img, tuple(start_indices.astype(int)), tuple(end_indices.astype(int))), img
- def template2subject(self, img: ants.ANTsImage, trans: ants.ANTsTransform, ref: ants.ANTsImage) -> ants.ANTsImage:
- """
- transform the image from template space to the subject space
- Args:
- img: (ANTsImage)
- the image to be transformed
- trans: (ANTs transformation)
- transformation matrix (from subject space to template space)
- ref: (ANTsImage)
- reference image
- Returns:
- img: (ANTsImage)
- the transformed image in subject space
- """
- trans_inv = ants.invert_ants_transform(trans)
- result = ants.apply_ants_transform_to_image(trans_inv, img, ref)
- return result
- def subject_preprocess(t1_file: str = None, t2_file: str = None, brain_mask_file: str = None, label_file: str = None,
- BoundingBox: TemplateCerebellarBoundingBox = TemplateCerebellarBoundingBox(),
- 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]:
- """
- function to preprocess a single subject.
- 1. Transform the image from subject space to the template space
- 2. Using a pre-defined bounding box to crop the image in the template space
- Args:
- t1_file: (string)
- file to T1w image
- t2_file: (string)
- file to T2w image
- brain_mask_file: (string)
- file to the brain mask image (can be used to improve affine registration)
- label_file: (string)
- file to label image (Optional, this image will be transformed into the template space using the same transformation.)
- BoundingBox: (TemplateCerebellarBoundingBox)
- the bounding box
- type_of_transform: (string)
- reserved for future use (see ANTspy)
- max_iterations: (int)
- maximum number of registration iterations
- use_q_form: (bool)
- set to True to use q-form
- Returns:
- trans: (ANTsTransform)
- transformation from subject space to template space
- t1_crop: (ANTsImage)
- cropped cerebellar area from transformed T1w image
- t2_crop: (ANTsImage)
- cropped cerebellar area from transformed T2w image
- label_crop: (ANTsImage)
- cropped cerebellar area from transformed label image
- t1_whole: (ANTsImage)
- whole transformed T1w image
- t2_whole: (ANTsImage)
- whole transformed T2w image
- """
- if t1_file is not None:
- t1 = img_read(file=t1_file, use_q_form=use_q_form, verbose=True)
- else:
- t1 = None
- if t2_file is not None:
- t2 = img_read(file=t2_file, use_q_form=use_q_form, verbose=True)
- else:
- t2 = None
- if brain_mask_file is not None:
- brain_mask = img_read(file=brain_mask_file, use_q_form=use_q_form)
- else:
- brain_mask = None
- # Read additional images
- if label_file is not None:
- label = img_read(file=label_file, use_q_form=use_q_form)
- else:
- label = None
- # If T1 and T2 are both given, but not in the same spacing, reslice T2 to T1 first
- if t2 is not None and t1 is not None:
- if ants.get_spacing(t1) != ants.get_spacing(t2):
- t2 = ants.registration(fixed=t1, moving=t2, type_of_transform='Rigid')['warpedmovout']
- if t1 is not None:
- trans, status = BoundingBox.registration(img=t1, brain_mask=brain_mask, type_of_transform=type_of_transform, max_iterations=max_iterations)
- else:
- trans, status = BoundingBox.registration(img=t2, brain_mask=brain_mask, type_of_transform=type_of_transform, max_iterations=max_iterations)
- if status == 2:
- raise RegistrationError('invalid registration detected')
- if status == 1:
- warnings.warn(f'low-quality registration detected, please double check the results for {t1_file if t1 is not None else t2_file}')
- if t1 is not None:
- t1_crop, t1_whole = BoundingBox.crop(t1, trans)
- else:
- t1_crop = None
- t1_whole = None
- if t2 is not None:
- t2_crop, t2_whole = BoundingBox.crop(t2, trans)
- else:
- t2_crop = None
- t2_whole = None
- if label is not None:
- label_crop, _ = BoundingBox.crop(label, trans)
- else:
- label_crop = None
- return trans, t1_crop, t2_crop, label_crop, t1_whole, t2_whole
- def threshold(img: ants.ANTsImage, lower: float = 0.5, upper: float = 1.0) -> ants.ANTsImage:
- """
- remove all other values from the image
- Args:
- img: (ANTsImage)
- the input image
- lower: (float)
- lower threshold
- upper: (float)
- upper threshold
- Returns:
- image : (ANTsImage)
- the thresholded image
- """
- img[img < lower] = 0
- img[img > upper] = 0
- return img
- def remove_islands(img: ants.ANTsImage) -> ants.ANTsImage:
- """ Removes parts of the mask that is not connected to the largest cluster
- Args:
- img (ANTsImage): the input image
- Returns:
- mask (ANTsImage): Image containing the largest connected component
- """
- clusters = ants.image_to_cluster_images(img)
- mask = None
- voxels = 0
- for temp in clusters:
- if temp.numpy().sum() > voxels:
- mask = temp
- voxels = temp.numpy().sum()
- return mask
- def subject_postprocess(mask: ants.ANTsImage, trans: ants.ANTsTransform, BoundingBox: TemplateCerebellarBoundingBox, ref: ants.ANTsImage) -> ants.ANTsImage:
- """
- transform the predicted cerebellum mask to the original space
- Args:
- mask: (ANTsImage)
- the predicted cerebellum mask from the template space
- trans: (ANTsTransform)
- the transformation from subject space to template space
- BoundingBox: (TemplateCerebellarBoundingBox)
- the bounding box
- ref: (ANTsImage)
- the reference image
- Returns:
- result: (ANTsImage)
- the final cerebellum mask from the subject space
- """
- result = BoundingBox.template2subject(mask, trans, ref)
- # threshold and binarize the image
- result = threshold(result)
- result[result != 0] = 1
- result = remove_islands(result)
- return result
- def isolate(t1_file: str = None, t2_file: str = None,
- brain_mask_file: str = None,
- label_file: str = None,
- result_folder: str = None,
- template: str = 'MNI152NLin6Asym',
- type_of_transform: str = 'Similarity',
- max_iterations: int = 5,
- params: str = 'pre_trained_numpy.pkl',
- save_cropped_files: bool = False,
- use_q_form: bool = False,
- verbose: bool = True) -> ants.ANTsImage:
- """
- main function for cerebellum isolation
- Args:
- t1_file: (string)
- filename and path to T1w image, optional
- t2_file: (string)
- filename and path to T2w image, optional
- brain_mask_file: (string)
- filename and path to brain mask, optional
- label_file: (string)
- filename and path to label image, optional (reserved, currently has no effect)
- result_folder: (string)
- path to output folder (optional, otherwise it is saved to input image folder)
- template: (string)
- template to use (reserved)
- type_of_transform: (string)
- reserved for future use (see ANTspy)
- max_iterations: (int)
- maximum number of registration iterations (optional, default 5)
- params: (string)
- path to params file (reserved)
- save_cropped_files: (bool)
- set to True to save files cropped to window
- use_q_form: (bool)
- set to True to use q-form
- verbose: (bool)
- whether to print out status information during processing
- """
- if t1_file is not None:
- result_folder = os.path.dirname(os.path.abspath(t1_file)) if result_folder is None else result_folder
- basename = os.path.splitext(os.path.basename(t1_file))
- elif t2_file is not None:
- result_folder = os.path.dirname(os.path.abspath(t2_file)) if result_folder is None else result_folder
- basename = os.path.splitext(os.path.basename(t2_file))
- else:
- raise RuntimeError('Must specify either t1_file or t2_file')
- # Strip .nii or .nii.gz extension
- if basename[1] == '.gz':
- basename = os.path.splitext(basename[0])
- basename = basename[0]
- # find paramter file and template bounding box
- base_dir = os.path.dirname(os.path.abspath(__file__))
- params_file = os.path.join(base_dir, 'parameters', params)
- BoundingBox = TemplateCerebellarBoundingBox(template_name=template)
- try:
- # Crop the images to the Unet input window
- if verbose:
- print(f"preprocessing {t1_file if t1_file is not None else t2_file}")
- trans, t1_crop, t2_crop, label_crop, _, _ = subject_preprocess(t1_file=t1_file,
- t2_file=t2_file,
- brain_mask_file=brain_mask_file,
- label_file=label_file,
- BoundingBox=BoundingBox,
- type_of_transform=type_of_transform,
- max_iterations=max_iterations,
- use_q_form=use_q_form)
- if isinstance(t1_crop, ants.core.ants_image.ANTsImage):
- t1_crop_data = t1_crop.numpy()
- else:
- t1_crop_data = t1_crop
- if isinstance(t2_crop, ants.core.ants_image.ANTsImage):
- t2_crop_data = t2_crop.numpy()
- else:
- t2_crop_data = t2_crop
- if isinstance(label_crop, ants.core.ants_image.ANTsImage):
- label_crop_data = label_crop.numpy()
- else:
- label_crop_data = label_crop
- # Do a forward pass through the Unet model
- if verbose:
- print('isolating cerebellum using UNet model')
- mask_template = predict(params_file=params_file, t1=t1_crop_data, t2=t2_crop_data)
- mask_template = nib.Nifti1Image(mask_template, BoundingBox.get_cropped_affine())
- mask_template = ants.from_nibabel_nifti(mask_template)
- # Postprocess and transform the mask back to subject space
- if verbose:
- print('postprocessing')
- if t1_file is not None:
- mask_subject = subject_postprocess(mask=mask_template, trans=trans, BoundingBox=BoundingBox, ref=img_read(file=t1_file, use_q_form=use_q_form))
- else:
- mask_subject = subject_postprocess(mask=mask_template, trans=trans, BoundingBox=BoundingBox, ref=img_read(file=t2_file, use_q_form=use_q_form))
- # use the original header info for the dseg mask
- if t1_file is not None:
- ref = nib.load(t1_file)
- else:
- ref = nib.load(t2_file)
- mask_subject = nib.Nifti1Image(mask_subject.numpy().astype(int), affine=ref.affine, header=ref.header)
- os.makedirs(result_folder, exist_ok=True)
- ofname = f'{basename}_cerebellum_dseg.nii.gz'
- if verbose:
- print(f"saving results to {ofname}")
- nib.save(mask_subject, os.path.join(result_folder, ofname))
- if save_cropped_files:
- if verbose:
- print(f"saving intermediate results to {result_folder}")
- if t1_crop is not None:
- ants.image_write(t1_crop, os.path.join(result_folder, f'{basename}_crop.nii.gz'))
- else:
- ants.image_write(t2_crop, os.path.join(result_folder, f'{basename}_crop.nii.gz'))
- ants.image_write(mask_template, os.path.join(result_folder, f'{basename}_cerebellum_crop_dseg.nii.gz'))
- ants.write_transform(trans, os.path.join(result_folder, f'{basename}_trans.mat'))
- except RegistrationError as e:
- print(f'Caught registration error for {t1_file if t1_file is not None else t2_file} : {e}')
- print(f'Isolation fails on {t1_file if t1_file is not None else t2_file}. No results were saved')
- mask_subject = None
- except InputError as e:
- print(f'Caught input file error for {t1_file if t1_file is not None else t2_file} : {e}')
- print(f'Isolation fails on {t1_file if t1_file is not None else t2_file}. No results were saved')
- mask_subject = None
- return img_read(file=os.path.join(result_folder, ofname), use_q_form=use_q_form) if mask_subject is not None else None
- if __name__ == '__main__':
- parser = argparse.ArgumentParser()
- parser.add_argument('--T1', type=str, help='path to T1w image')
- parser.add_argument('--T2', type=str, help='path to T2w image')
- parser.add_argument('--brain_mask', type=str, help='path to brain mask image')
- parser.add_argument('--label', type=str, help='path to label image (reserved, currently has no effect)')
- parser.add_argument('--result_folder', type=str, help='path to save the isolation image (results will be saved to '
- 'T1w image folder (or T2w image folder if no T1w image is '
- 'specified))')
- parser.add_argument('--template', type=str, default='MNI152NLin6Asym',
- help='template for registration (MNI152NLin6Asym by '
- 'default)')
- parser.add_argument('--type_of_transform', type=str, default='Similarity', help='reserved for future use (see ANTspy)')
- parser.add_argument('--max_iterations', type=int, default=5, help='maximum number of registration iterations (optional, default 5)')
- parser.add_argument('--params', type=str, default='pre_trained.pkl', help='pretrained parameter file')
- parser.add_argument('--save_cropped_files', action='store_true', help='whether to save files cropped to UNet input window')
- parser.add_argument('--use_q_form', action='store_true', help='whether to use q-form')
- parser.add_argument('--verbose', action='store_true', help='whether to print out status information during processing')
- args = parser.parse_args()
- print(args)
- if args.T1 is None and args.T2 is None:
- raise RuntimeError('Must specify either t1_file or t2_file')
- if args.result_folder is None:
- if args.T1 is None:
- args.result_folder = os.path.dirname(os.path.abspath(args.T2))
- else:
- args.result_folder = os.path.dirname(os.path.abspath(args.T1))
- isolate(t1_file=args.T1,
- t2_file=args.T2,
- brain_mask_file=args.brain_mask,
- label_file=args.label,
- result_folder=args.result_folder,
- template=args.template,
- type_of_transform=args.type_of_transform,
- max_iterations=args.max_iterations,
- params=args.params,
- save_cropped_files=args.save_cropped_files,
- use_q_form=args.use_q_form,
- verbose=args.verbose)
isolation.py at commit 3d20055, under MIT · at the source
Overview
- Western Institute for Neuroscience, Western University, London, ON, Canada
- Department of Computer Science, Western University, London, ON, Canada
- Donders Institute for Brain, Cognition and Behaviour, Radboud University Medical Centre, Nijmegen, The Netherlands
- Faculty of Computer Science, Dalhousie University, Halifax, Canada
- Department of Experimental Psychology, University of Oxford, Oxford, United Kingdom
- Department of Psychology, Western University, London, ON, Canada
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
3d20055022e5339687f412da92979520c7568da2, 12 September 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
26 files
- SUITPy/
__init__.py , Python, 26 lines - SUITPy/
atlas.py , Python, 405 lines - SUITPy/
flatmap.py , Python, 896 lines - SUITPy/
isolation.py , Python, 1,469 lines, 5 matches - SUITPy/
normalization.py , Python, 295 lines, 1 match - SUITPy/
reslice.py , Python, 142 lines - SUITPy/
utils.py , Python, 820 lines - docs/
source/ , Python, 57 linesconf.py - docs/
source/ , Jupyter, 164 linestutorials/ 1.quickstart_fMRI.ipynb - docs/
source/ , Jupyter, 188 linestutorials/ 2.quickstart_anatomical. ipynb - docs/
source/ , Jupyter, 68 linestutorials/ 3.isolate_example.ipynb - docs/
source/ , Jupyter, 155 linestutorials/ 4.normalize_example.ipyn b - docs/
source/ , Jupyter, 114 linestutorials/ 5.reslice_example.ipynb - docs/
source/ , Jupyter, 189 linestutorials/ 6.flatmap_example.ipynb - setup.py, Python, 79 lines
- tests/
resample_example_dataset , Python, 17 lines.py - tests/
test_atlas.py , Python, 26 lines - tests/
test_estimate_volume_ind , Jupyter, 193 linesivid.ipynb - tests/
test_flatmap.py , Python, 64 lines - tests/
test_isolate.py , Python, 19 lines - tests/
test_normalize.py , Python, 29 lines - tests/
test_reslice.ipynb , Jupyter, 73 lines - tests/
test_reslice.py , Python, 52 lines - tests/
test_roi_summarization.p , Python, 30 linesy - LICENSE, License, 9 lines
- README.rst, Text, 57 lines
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://
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://
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/
url = {https://
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/
VL - 4
SP - IMAG.a.1323
SN - 2837-6056
PB - MIT Press
DO - 10.1162/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1162/
"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":
"volume": "4",
"page": "IMAG.a.1323",
"DOI": "10.1162/
"PMID": "42569404",
"PMCID": "PMC13449926",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://
"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 brainJournal: 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 communicationsIn 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: NeuroImageIn 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 psychiatryIn 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 communicationsIn 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 methodsIn 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 communicationsIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 24 scripts, and 6 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:f44ee012fadb66ef…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
