Image Distortion Correction for Diffusion MR Imaging Using a Transformer-based U-Net.
The 6 matches
- [1] § Materials and Methods › Network architecture › TransDisCo (proposed method) ↔ Baseline_Transformers/models/nnFormer/Swin_Unet_s_ACDC_2laterdown.py, lines 420–533 · score 0.77 · SW MSA, Swin Transformer block, Patch merging, layer normalization, partition, upsampling
- [2] § Materials and Methods › Network architecture › TransDisCo (proposed method) ↔ Baseline_Transformers/models/nnFormer/Swin_Unet_l_gelunorm.py, lines 445–544 · score 0.76 · SW MSA, Swin Transformer block, Patch merging, layer normalization, partition, upsampling
- [3] § Materials and Methods › Network architecture ↔ Baseline_Transformers/models/CoTr/training/network_training/nnUNetTrainer.py, lines 233–264 · score 0.56 · LeakyReLU, Conv3d, kernel, architectural, Network, Transformer
- [4] § Materials and Methods › Network architecture › CNN-DisCo (ablation study model) ↔ Baseline_Transformers/models/configs_PVT.py, lines 1–26 · score 0.55 · Swin Transformer blocks, skip connections, convolutional layers, depth, model
- [5] § Materials and Methods › Network architecture › CNN-DisCo (ablation study model) ↔ IXI/TransMorph/models/configs_TransMorph.py, lines 1–28 · score 0.55 · Swin Transformer blocks, skip connections, convolutional layers, depth, model
- [6] § Materials and Methods › Network architecture ↔ Baseline_Transformers/models/nnFormer/generic_UNet.py, lines 184–246 · score 0.54 · LeakyReLU, Conv3d, kernel, architectural, convolutional, blocks
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 · 976 lines · 38 KB · MIT · 1 match
- from einops import rearrange
- from copy import deepcopy
- from nnformer.utilities.nd_softmax import softmax_helper
- from torch import nn
- import torch
- import numpy as np
- from nnformer.network_architecture.initialization import InitWeights_He
- from nnformer.network_architecture.neural_network import SegmentationNetwork
- import torch.nn.functional
- import torch.nn.functional as F
- import torch.utils.checkpoint as checkpoint
- from timm.models.layers import DropPath, to_3tuple, trunc_normal_
- class Mlp(nn.Module):
- """ Multilayer perceptron."""
- def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
- super().__init__()
- out_features = out_features or in_features
- hidden_features = hidden_features or in_features
- self.fc1 = nn.Linear(in_features, hidden_features)
- self.act = act_layer()
- self.fc2 = nn.Linear(hidden_features, out_features)
- self.drop = nn.Dropout(drop)
- def forward(self, x):
- x = self.fc1(x)
- x = self.act(x)
- x = self.drop(x)
- x = self.fc2(x)
- x = self.drop(x)
- return x
- def window_partition(x, window_size):
- B, S, H, W, C = x.shape
- x = x.view(B, S // window_size[0], window_size[0], H // window_size[1], window_size[1], W // window_size[2], window_size[2], C)
- windows = x.permute(0, 1, 3, 5, 2, 4, 6, 7).contiguous().view(-1, window_size[0], window_size[1], window_size[2], C)
- return windows
- def window_reverse(windows, window_size, S, H, W):
- B = int(windows.shape[0] / (S * H * W / window_size[0] / window_size[1] / window_size[2]))
- x = windows.view(B, S // window_size[0], H // window_size[1], W // window_size[2], window_size[0], window_size[1], window_size[2], -1)
- x = x.permute(0, 1, 4, 2, 5, 3, 6, 7).contiguous().view(B, S, H, W, -1)
- return x
- class WindowAttention(nn.Module):
- """ Window based multi-head self attention (W-MSA) module with relative position bias.
- It supports both of shifted and non-shifted window.
- Args:
- dim (int): Number of input channels.
- window_size (tuple[int]): The height and width of the window.
- num_heads (int): Number of attention heads.
- qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
- qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
- attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
- proj_drop (float, optional): Dropout ratio of output. Default: 0.0
- """
- def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):
- super().__init__()
- self.dim = dim
- self.window_size = window_size # Wh, Ww
- self.num_heads = num_heads
- head_dim = dim // num_heads
- self.scale = qk_scale or head_dim ** -0.5
- # define a parameter table of relative position bias
- self.relative_position_bias_table = nn.Parameter(
- torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1) * (2 * window_size[2] - 1),
- num_heads))
- # get pair-wise relative position index for each token inside the window
- coords_s = torch.arange(self.window_size[0])
- coords_h = torch.arange(self.window_size[1])
- coords_w = torch.arange(self.window_size[2])
- coords = torch.stack(torch.meshgrid([coords_s, coords_h, coords_w]))
- coords_flatten = torch.flatten(coords, 1)
- relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]
- relative_coords = relative_coords.permute(1, 2, 0).contiguous()
- relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
- relative_coords[:, :, 1] += self.window_size[1] - 1
- relative_coords[:, :, 2] += self.window_size[2] - 1
- relative_coords[:, :, 0] *= 3 * self.window_size[1] - 1
- relative_coords[:, :, 1] *= 2 * self.window_size[1] - 1
- relative_position_index = relative_coords.sum(-1)
- self.register_buffer("relative_position_index", relative_position_index)
- self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
- self.attn_drop = nn.Dropout(attn_drop)
- self.proj = nn.Linear(dim, dim)
- self.proj_drop = nn.Dropout(proj_drop)
- trunc_normal_(self.relative_position_bias_table, std=.02)
- self.softmax = nn.Softmax(dim=-1)
- def forward(self, x, mask=None):
- B_, N, C = x.shape
- qkv = self.qkv(x)
- qkv=qkv.reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
- q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
- q = q * self.scale
- attn = (q @ k.transpose(-2, -1))
- relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
- self.window_size[0] * self.window_size[1] * self.window_size[2],
- self.window_size[0] * self.window_size[1] * self.window_size[2], -1)
- relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()
- attn = attn + relative_position_bias.unsqueeze(0)
- if mask is not None:
- nW = mask.shape[0]
- attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
- attn = attn.view(-1, self.num_heads, N, N)
- attn = self.softmax(attn)
- else:
- attn = self.softmax(attn)
- attn = self.attn_drop(attn)
- x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
- x = self.proj(x)
- x = self.proj_drop(x)
- return x
- class SwinTransformerBlock(nn.Module):
- """ Swin Transformer Block.
- Args:
- dim (int): Number of input channels.
- num_heads (int): Number of attention heads.
- window_size (int): Window size.
- shift_size (int): Shift size for SW-MSA.
- mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
- qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
- qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
- drop (float, optional): Dropout rate. Default: 0.0
- attn_drop (float, optional): Attention dropout rate. Default: 0.0
- drop_path (float, optional): Stochastic depth rate. Default: 0.0
- act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
- norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
- """
- def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,
- mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
- act_layer=nn.GELU, norm_layer=nn.LayerNorm):
- super().__init__()
- self.dim = dim
- self.input_resolution = input_resolution
- self.num_heads = num_heads
- self.window_size = window_size
- self.shift_size = shift_size
- self.mlp_ratio = mlp_ratio
- if tuple(self.input_resolution) == tuple(self.window_size):
- # if window size is larger than input resolution, we don't partition windows
- self.shift_size = [0,0,0]
- #self.window_size = min(self.input_resolution)
- #assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
- self.norm1 = norm_layer(dim)
- self.attn = WindowAttention(
- dim, window_size=self.window_size, num_heads=num_heads,
- qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
- self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
- self.norm2 = norm_layer(dim)
- mlp_hidden_dim = int(dim * mlp_ratio)
- self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
- def forward(self, x, mask_matrix):
- B, L, C = x.shape
- S, H, W = self.input_resolution
- assert L == S * H * W, "input feature has wrong size"
- shortcut = x
- x = self.norm1(x)
- x = x.view(B, S, H, W, C)
- # pad feature maps to multiples of window size
- pad_r = (self.window_size[2] - W % self.window_size[2]) % self.window_size[2]
- pad_b = (self.window_size[1] - H % self.window_size[1]) % self.window_size[1]
- pad_g = (self.window_size[0] - S % self.window_size[0]) % self.window_size[0]
- x = F.pad(x, (0, 0, 0, pad_r, 0, pad_b, 0, pad_g))
- _, Sp, Hp, Wp, _ = x.shape
- # cyclic shift
- if min(self.shift_size) > 0:
- shifted_x = torch.roll(x, shifts=(-self.shift_size[0], -self.shift_size[1],-self.shift_size[2]), dims=(1, 2,3))
- attn_mask = mask_matrix
- else:
- shifted_x = x
- attn_mask = None
- # partition windows
- x_windows = window_partition(shifted_x, self.window_size)
- x_windows = x_windows.view(-1, self.window_size[0] * self.window_size[1] * self.window_size[2],
- C)
- # W-MSA/SW-MSA
- attn_windows = self.attn(x_windows, mask=attn_mask)
- # merge windows
- attn_windows = attn_windows.view(-1, self.window_size[0], self.window_size[1], self.window_size[2], C)
- shifted_x = window_reverse(attn_windows, self.window_size, Sp, Hp, Wp)
- # reverse cyclic shift
- if min(self.shift_size) > 0:
- x = torch.roll(shifted_x, shifts=(self.shift_size[0], self.shift_size[1], self.shift_size[2]), dims=(1, 2, 3))
- else:
- x = shifted_x
- if pad_r > 0 or pad_b > 0 or pad_g > 0:
- x = x[:, :S, :H, :W, :].contiguous()
- x = x.view(B, S * H * W, C)
- # FFN
- x = shortcut + self.drop_path(x)
- x = x + self.drop_path(self.mlp(self.norm2(x)))
- return x
- class PatchMerging(nn.Module):
- """ Patch Merging Layer
- Args:
- dim (int): Number of input channels.
- norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
- """
- def __init__(self, dim, norm_layer=nn.LayerNorm,tag=None):
- super().__init__()
- self.dim = dim
- if tag==0:
- self.reduction = nn.Conv3d(dim,dim*2,kernel_size=[1,2,2],stride=[1,2,2])
- else:
- self.reduction = nn.Conv3d(dim,dim*2,kernel_size=[2,2,2],stride=[2,2,2])
- self.norm = norm_layer(dim)
- def forward(self, x, S, H, W):
- B, L, C = x.shape
- assert L == H * W * S, "input feature has wrong size"
- x = x.view(B, S, H, W, C)
- x = F.gelu(x)
- x = self.norm(x)
- x=x.permute(0,4,1,2,3)
- x=self.reduction(x)
- x=x.permute(0,2,3,4,1).view(B,-1,2*C)
- return x
- class Patch_Expanding(nn.Module):
- def __init__(self, dim, norm_layer=nn.LayerNorm,tag=None):
- super().__init__()
- self.dim = dim
- self.norm = norm_layer(dim)
- if tag==0:
- self.up=nn.ConvTranspose3d(dim,dim//2,[1,2,2],[1,2,2])
- elif tag==1:
- self.up=nn.ConvTranspose3d(dim,dim//2,[2,2,2],[2,2,2])
- elif tag==2:
- self.up=nn.ConvTranspose3d(dim,dim//2,[2,2,2],[2,2,2],output_padding=[1,0,0])
- def forward(self, x, S, H, W):
- """ Forward function.
- Args:
- x: Input feature, tensor size (B, H*W, C).
- H, W: Spatial resolution of the input feature.
- """
- B, L, C = x.shape
- assert L == H * W * S, "input feature has wrong size"
- x = x.view(B, S, H, W, C)
- x = self.norm(x)
- x=x.permute(0,4,1,2,3)
- x = self.up(x)
- x=x.permute(0,2,3,4,1).view(B,-1,C//2)
- return x
- class BasicLayer(nn.Module):
- """ A basic Swin Transformer layer for one stage.
- Args:
- dim (int): Number of feature channels
- depth (int): Depths of this stage.
- num_heads (int): Number of attention head.
- window_size (int): Local window size. Default: 7.
- mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
- qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
- qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
- drop (float, optional): Dropout rate. Default: 0.0
- attn_drop (float, optional): Attention dropout rate. Default: 0.0
- drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
- norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
- downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
- use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
- """
- def __init__(self,
- dim,
- input_resolution,
- depth,
- num_heads,
- window_size=7,
- mlp_ratio=4.,
- qkv_bias=True,
- qk_scale=None,
- drop=0.,
- attn_drop=0.,
- drop_path=0.,
- norm_layer=nn.LayerNorm,
- downsample=True,
- use_checkpoint=False,
- i_layer=None):
- super().__init__()
- self.window_size = window_size
- self.shift_size = [window_size[0] // 2,window_size[1] // 2,window_size[2] // 2]
- self.depth = depth
- self.use_checkpoint = use_checkpoint
- self.i_layer=i_layer
- # build blocks
- self.blocks = nn.ModuleList([
- SwinTransformerBlock(
- dim=dim,
- input_resolution=input_resolution,
- num_heads=num_heads,
- window_size=window_size,
- shift_size=[0,0,0] if (i % 2 == 0) else self.shift_size,
- mlp_ratio=mlp_ratio,
- qkv_bias=qkv_bias,
- qk_scale=qk_scale,
- drop=drop,
- attn_drop=attn_drop,
- drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path, norm_layer=norm_layer)
- for i in range(depth)])
- # patch merging layer
- if downsample is not None:
- if i_layer==1 or i_layer==2:
- self.downsample = downsample(dim=dim, norm_layer=norm_layer,tag=1)
- else:
- self.downsample = downsample(dim=dim, norm_layer=norm_layer,tag=0)
- else:
- self.downsample = None
- def forward(self, x, S, H, W):
- # calculate attention mask for SW-MSA
- Sp = int(np.ceil(S / self.window_size[0])) * self.window_size[0]
- Hp = int(np.ceil(H / self.window_size[1])) * self.window_size[1]
- Wp = int(np.ceil(W / self.window_size[2])) * self.window_size[2]
- img_mask = torch.zeros((1, Sp, Hp, Wp, 1), device=x.device)
- s_slices = (slice(0, -self.window_size[0]),
- slice(-self.window_size[0], -self.shift_size[0]),
- slice(-self.shift_size[0], None))
- h_slices = (slice(0, -self.window_size[1]),
- slice(-self.window_size[1], -self.shift_size[1]),
- slice(-self.shift_size[1], None))
- w_slices = (slice(0, -self.window_size[2]),
- slice(-self.window_size[2], -self.shift_size[2]),
- slice(-self.shift_size[2], None))
- cnt = 0
- for s in s_slices:
- for h in h_slices:
- for w in w_slices:
- img_mask[:, s, h, w, :] = cnt
- cnt += 1
- mask_windows = window_partition(img_mask, self.window_size)
- mask_windows = mask_windows.view(-1,
- self.window_size[0] * self.window_size[1] * self.window_size[2])
- attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
- attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
- for blk in self.blocks:
- blk.H, blk.W = H, W
- if self.use_checkpoint:
- x = checkpoint.checkpoint(blk, x, attn_mask)
- else:
- x = blk(x, attn_mask)
- if self.downsample is not None:
- x_down = self.downsample(x, S, H, W)
- if self.i_layer!=1 and self.i_layer!=2:
- Ws, Wh, Ww = S , (H + 1) // 2, (W + 1) // 2
- else:
- Ws, Wh, Ww = S//2 , (H + 1) // 2, (W + 1) // 2
- return x, S, H, W, x_down, Ws, Wh, Ww
- else:
- return x, S, H, W, x, S, H, W
- class BasicLayer_up(nn.Module):
- """ A basic Swin Transformer layer for one stage.
- Args:
- dim (int): Number of feature channels
- depth (int): Depths of this stage.
- num_heads (int): Number of attention head.
- window_size (int): Local window size. Default: 7.
- mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
- qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
- qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
- drop (float, optional): Dropout rate. Default: 0.0
- attn_drop (float, optional): Attention dropout rate. Default: 0.0
- drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
- norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
- downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
- use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
- """
- def __init__(self,
- dim,
- input_resolution,
- depth,
- num_heads,
- window_size=7,
- mlp_ratio=4.,
- qkv_bias=True,
- qk_scale=None,
- drop=0.,
- attn_drop=0.,
- drop_path=0.,
- norm_layer=nn.LayerNorm,
- upsample=True,
- i_layer=None
- ):
- super().__init__()
- self.window_size = window_size
- self.shift_size = [window_size[0] // 2,window_size[1] // 2,window_size[2] // 2]
- self.depth = depth
- # build blocks
- self.blocks = nn.ModuleList([
- SwinTransformerBlock(
- dim=dim,
- input_resolution=input_resolution,
- num_heads=num_heads,
- window_size=window_size,
- shift_size=[0,0,0] if (i % 2 == 0) else self.shift_size,
- mlp_ratio=mlp_ratio,
- qkv_bias=qkv_bias,
- qk_scale=qk_scale,
- drop=drop,
- attn_drop=attn_drop,
- drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path, norm_layer=norm_layer)
- for i in range(depth)])
- # patch merging layer
- self.i_layer=i_layer
- if i_layer==1:
- self.Upsample = upsample(dim=2*dim, norm_layer=norm_layer,tag=1)
- elif i_layer==0:
- self.Upsample = upsample(dim=2*dim, norm_layer=norm_layer,tag=2)
- else:
- self.Upsample = upsample(dim=2*dim, norm_layer=norm_layer,tag=0)
- def forward(self, x,skip, S, H, W):
- """ Forward function.
- Args:
- x: Input feature, tensor size (B, H*W, C).
- H, W: Spatial resolution of the input feature.
- """
- skip = skip.flatten(2).transpose(1, 2)
- x_up = self.Upsample(x, S, H, W)
- x_up+=skip
- if self.i_layer==1:
- S, H, W = S * 2 , H * 2, W * 2
- elif self.i_layer==0:
- S, H, W = (S * 2)+1 , H * 2, W * 2
- else:
- S, H, W = S , H * 2, W * 2
- # calculate attention mask for SW-MSA
- Sp = int(np.ceil(S / self.window_size[0])) * self.window_size[0]
- Hp = int(np.ceil(H / self.window_size[1])) * self.window_size[1]
- Wp = int(np.ceil(W / self.window_size[2])) * self.window_size[2]
- img_mask = torch.zeros((1, Sp, Hp, Wp, 1), device=x.device)
- s_slices = (slice(0, -self.window_size[0]),
- slice(-self.window_size[0], -self.shift_size[0]),
- slice(-self.shift_size[0], None))
- h_slices = (slice(0, -self.window_size[1]),
- slice(-self.window_size[1], -self.shift_size[1]),
- slice(-self.shift_size[1], None))
- w_slices = (slice(0, -self.window_size[2]),
- slice(-self.window_size[2], -self.shift_size[2]),
- slice(-self.shift_size[2], None))
- cnt = 0
- for s in s_slices:
- for h in h_slices:
- for w in w_slices:
- img_mask[:, s, h, w, :] = cnt
- cnt += 1
- mask_windows = window_partition(img_mask, self.window_size)
- mask_windows = mask_windows.view(-1,
- self.window_size[0] * self.window_size[1] * self.window_size[2])
- attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
- attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
- for blk in self.blocks:
- x_up = blk(x_up, attn_mask)
- return x_up, S, H, W
- # done
- class project(nn.Module):
- def __init__(self,in_dim,out_dim,stride,padding,activate,norm,last=False):
- super().__init__()
- self.out_dim=out_dim
- self.conv1=nn.Conv3d(in_dim,out_dim,kernel_size=3,stride=stride,padding=padding)
- self.conv2=nn.Conv3d(out_dim,out_dim,kernel_size=3,stride=1,padding=1)
- self.activate=activate()
- self.norm1=norm(out_dim)
- self.last=last
- if not last:
- self.norm2=norm(out_dim)
- def forward(self,x):
- x=self.conv1(x)
- x=self.activate(x)
- #norm1
- Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
- x = x.flatten(2).transpose(1, 2)
- x = self.norm1(x)
- x = x.transpose(1, 2).view(-1, self.out_dim, Ws, Wh, Ww)
- x=self.conv2(x)
- if not self.last:
- x=self.activate(x)
- #norm2
- Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
- x = x.flatten(2).transpose(1, 2)
- x = self.norm2(x)
- x = x.transpose(1, 2).view(-1, self.out_dim, Ws, Wh, Ww)
- return x
- class PatchEmbed(nn.Module):
- """ Image to Patch Embedding
- Args:
- patch_size (int): Patch token size. Default: 4.
- in_chans (int): Number of input image channels. Default: 3.
- embed_dim (int): Number of linear projection output channels. Default: 96.
- norm_layer (nn.Module, optional): Normalization layer. Default: None
- """
- def __init__(self, patch_size=4, in_chans=4, embed_dim=96, norm_layer=None):
- super().__init__()
- patch_size = to_3tuple(patch_size)
- self.patch_size = patch_size
- self.in_chans = in_chans
- self.embed_dim = embed_dim
- self.proj1 = project(in_chans,embed_dim//2,[1,2,2],1,nn.GELU,nn.LayerNorm,False)
- self.proj2 = project(embed_dim//2,embed_dim,[1,2,2],1,nn.GELU,nn.LayerNorm,True)
- if norm_layer is not None:
- self.norm = norm_layer(embed_dim)
- else:
- self.norm = None
- def forward(self, x):
- """Forward function."""
- # padding
- _, _, S, H, W = x.size()
- if W % self.patch_size[2] != 0:
- x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
- if H % self.patch_size[1] != 0:
- x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
- if S % self.patch_size[0] != 0:
- x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - S % self.patch_size[0]))
- x = self.proj1(x)
- x = self.proj2(x)
- if self.norm is not None:
- Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
- x = x.flatten(2).transpose(1, 2)
- x = self.norm(x)
- x = x.transpose(1, 2).view(-1, self.embed_dim, Ws, Wh, Ww)
- return x
- class SwinTransformer(nn.Module):
- """ Swin Transformer backbone.
- A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
- https://arxiv.org/pdf/2103.14030
- Args:
- pretrain_img_size (int): Input image size for training the pretrained model,
- used in absolute postion embedding. Default 224.
- patch_size (int | tuple(int)): Patch size. Default: 4.
- in_chans (int): Number of input image channels. Default: 3.
- embed_dim (int): Number of linear projection output channels. Default: 96.
- depths (tuple[int]): Depths of each Swin Transformer stage.
- num_heads (tuple[int]): Number of attention head of each stage.
- window_size (int): Window size. Default: 7.
- mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
- qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
- qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
- drop_rate (float): Dropout rate.
- attn_drop_rate (float): Attention dropout rate. Default: 0.
- drop_path_rate (float): Stochastic depth rate. Default: 0.2.
- norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
- ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
- patch_norm (bool): If True, add normalization after patch embedding. Default: True.
- out_indices (Sequence[int]): Output from which stages.
- frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
- -1 means not freezing any parameters.
- use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
- """
- def __init__(self,
- pretrain_img_size=224,
- patch_size=4,
- in_chans=1 ,
- embed_dim=96,
- depths=[2, 2, 2, 2],
- num_heads=[4, 8, 16, 32],
- window_size=7,
- mlp_ratio=4.,
- qkv_bias=True,
- qk_scale=None,
- drop_rate=0.,
- attn_drop_rate=0.,
- drop_path_rate=0.2,
- norm_layer=nn.LayerNorm,
- ape=False,
- patch_norm=True,
- out_indices=(0, 1, 2, 3),
- frozen_stages=-1,
- use_checkpoint=False):
- super().__init__()
- self.pretrain_img_size = pretrain_img_size
- self.num_layers = len(depths)
- self.embed_dim = embed_dim
- self.ape = ape
- self.patch_norm = patch_norm
- self.out_indices = out_indices
- self.frozen_stages = frozen_stages
- # split image into non-overlapping patches
- self.patch_embed = PatchEmbed(
- patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim,
- norm_layer=norm_layer if self.patch_norm else None)
- # absolute position embedding
- if self.ape:
- pretrain_img_size = to_3tuple(pretrain_img_size)
- patch_size = to_3tuple(patch_size)
- patches_resolution = [pretrain_img_size[0] // patch_size[0], pretrain_img_size[1] // patch_size[1],
- pretrain_img_size[2] // patch_size[2]]
- self.absolute_pos_embed = nn.Parameter(
- torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1], patches_resolution[2]))
- trunc_normal_(self.absolute_pos_embed, std=.02)
- self.pos_drop = nn.Dropout(p=drop_rate)
- # stochastic depth
- dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
- down_size=[[1,4,4],[1,8,8],[2,16,16],[4,32,32]]
- # build layers
- self.layers = nn.ModuleList()
- for i_layer in range(self.num_layers):
- layer = BasicLayer(
- dim=int(embed_dim * 2 ** i_layer),
- input_resolution=(
- pretrain_img_size[0] // down_size[i_layer][0], pretrain_img_size[1] // down_size[i_layer][1],
- pretrain_img_size[2] // down_size[i_layer][2]),
- depth=depths[i_layer],
- num_heads=num_heads[i_layer],
- window_size=window_size,
- mlp_ratio=mlp_ratio,
- qkv_bias=qkv_bias,
- qk_scale=qk_scale,
- drop=drop_rate,
- attn_drop=attn_drop_rate,
- drop_path=dpr[sum(
- depths[:i_layer]):sum(depths[:i_layer + 1])],
- norm_layer=norm_layer,
- downsample=PatchMerging,
- use_checkpoint=use_checkpoint,
- i_layer=i_layer)
- self.layers.append(layer)
- num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
- self.num_features = num_features
- # add a norm layer for each output
- for i_layer in out_indices:
- layer = norm_layer(num_features[i_layer])
- layer_name = f'norm{i_layer}'
- self.add_module(layer_name, layer)
- self._freeze_stages()
- def _freeze_stages(self):
- if self.frozen_stages >= 0:
- self.patch_embed.eval()
- for param in self.patch_embed.parameters():
- param.requires_grad = False
- if self.frozen_stages >= 1 and self.ape:
- self.absolute_pos_embed.requires_grad = False
- if self.frozen_stages >= 2:
- self.pos_drop.eval()
- for i in range(0, self.frozen_stages - 1):
- m = self.layers[i]
- m.eval()
- for param in m.parameters():
- param.requires_grad = False
- def forward(self, x):
- """Forward function."""
- x = self.patch_embed(x)
- down=[]
- Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
- if self.ape:
- # interpolate the position embedding to the corresponding size
- absolute_pos_embed = F.interpolate(self.absolute_pos_embed, size=(Ws, Wh, Ww), align_corners=True,
- mode='trilinear')
- x = (x + absolute_pos_embed).flatten(2).transpose(1, 2)
- else:
- x = x.flatten(2).transpose(1, 2)
- x = self.pos_drop(x)
- for i in range(self.num_layers):
- layer = self.layers[i]
- x_out, S, H, W, x, Ws, Wh, Ww = layer(x, Ws, Wh, Ww)
- if i in self.out_indices:
- norm_layer = getattr(self, f'norm{i}')
- x_out = norm_layer(x_out)
- out = x_out.view(-1, S, H, W, self.num_features[i]).permute(0, 4, 1, 2, 3).contiguous()
- down.append(out)
- return down
- def train(self, mode=True):
- """Convert the model into training mode while keep layers freezed."""
- super(SwinTransformer, self).train(mode)
- self._freeze_stages()
- class encoder(nn.Module):
- def __init__(self,
- pretrain_img_size,
- embed_dim,
- patch_size=4,
- depths=[2,2,2],
- num_heads=[24,12,6],
- window_size=4,
- mlp_ratio=4.,
- qkv_bias=True,
- qk_scale=None,
- drop_rate=0.,
- attn_drop_rate=0.,
- drop_path_rate=0.2,
- norm_layer=nn.LayerNorm
- ):
- super().__init__()
- self.num_layers = len(depths)
- self.pos_drop = nn.Dropout(p=drop_rate)
- # stochastic depth
- dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
- up_size=[[2,16,16],[1,8,8],[1,4,4]]
- # build layers
- self.layers = nn.ModuleList()
- for i_layer in range(self.num_layers)[::-1]:
- layer = BasicLayer_up(
- dim=int(embed_dim * 2 ** (len(depths)-i_layer-1)),
- input_resolution=(
- pretrain_img_size[0] // up_size[i_layer][0], pretrain_img_size[1] // up_size[i_layer][1],
- pretrain_img_size[2] // up_size[i_layer][2]),
- depth=depths[i_layer],
- num_heads=num_heads[i_layer],
- window_size=window_size,
- mlp_ratio=mlp_ratio,
- qkv_bias=qkv_bias,
- qk_scale=qk_scale,
- drop=drop_rate,
- attn_drop=attn_drop_rate,
- drop_path=dpr[sum(
- depths[:i_layer]):sum(depths[:i_layer + 1])],
- norm_layer=norm_layer,
- upsample=Patch_Expanding,
- i_layer=i_layer
- )
- self.layers.append(layer)
- self.num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
- def forward(self,x,skips):
- outs=[]
- S, H, W = x.size(2), x.size(3), x.size(4)
- x = x.flatten(2).transpose(1, 2)
- x = self.pos_drop(x)
- for i in range(self.num_layers)[::-1]:
- layer = self.layers[i]
- x, S, H, W, = layer(x,skips[i], S, H, W)
- out = x.view(-1, S, H, W, self.num_features[i])
- outs.append(out)
- return outs
- class final_patch_expanding(nn.Module):
- def __init__(self,dim,num_class,patch_size):
- super().__init__()
- self.up=nn.ConvTranspose3d(dim,num_class,patch_size,patch_size)
- def forward(self,x):
- x=x.permute(0,4,1,2,3)
- x=self.up(x)
- return x
- class swintransformer(SegmentationNetwork):
- def __init__(self, input_channels, base_num_features, num_classes, num_pool, num_conv_per_stage=2,
- feat_map_mul_on_downscale=2, conv_op=nn.Conv2d,
- norm_op=nn.BatchNorm2d, norm_op_kwargs=None,
- dropout_op=nn.Dropout2d, dropout_op_kwargs=None,
- nonlin=nn.LeakyReLU, nonlin_kwargs=None, deep_supervision=True, dropout_in_localization=False,
- final_nonlin=softmax_helper, weightInitializer=InitWeights_He(1e-2), pool_op_kernel_sizes=None,
- conv_kernel_sizes=None,
- upscale_logits=False, convolutional_pooling=False, convolutional_upsampling=False,
- max_num_features=None, basic_block=None,
- seg_output_use_bias=False):
- super(swintransformer, self).__init__()
- self._deep_supervision = deep_supervision
- self.do_ds = deep_supervision
- self.num_classes=num_classes
- self.conv_op=conv_op
- self.upscale_logits_ops = []
- self.upscale_logits_ops.append(lambda x: x)
- embed_dim=96
- depths=[2, 2, 2, 2]
- num_heads=[3, 6, 12, 24]
- patch_size=[1,4,4]
- self.model_down=SwinTransformer(pretrain_img_size=[14,160,160],window_size=[3,5,5],embed_dim=embed_dim,patch_size=patch_size,depths=depths,num_heads=num_heads,in_chans=1)
- self.encoder=encoder(pretrain_img_size=[14,160,160],embed_dim=embed_dim,window_size=[3,5,5],patch_size=patch_size,num_heads=[12,6,3],depths=[2,2,2])
- self.final=[]
- for i in range(len(depths)-1):
- self.final.append(final_patch_expanding(embed_dim*2**i,self.num_classes,patch_size=patch_size))
- self.final=nn.ModuleList(self.final)
- def forward(self, x):
- seg_outputs=[]
- skips = self.model_down(x)
- neck=skips[-1]
- out=self.encoder(neck,skips)
- for i in range(len(out)):
- seg_outputs.append(self.final[-(i+1)](out[i]))
- if self._deep_supervision and self.do_ds:
- return tuple([seg_outputs[-1]] + [i(j) for i, j in
- zip(list(self.upscale_logits_ops)[::-1], seg_outputs[:-1][::-1])])
- else:
- return seg_outputs[-1]
- @staticmethod
- def compute_approx_vram_consumption(patch_size, num_pool_per_axis, base_num_features, max_num_features,
- num_modalities, num_classes, pool_op_kernel_sizes, deep_supervision=False,
- conv_per_stage=2):
- """
- This only applies for num_conv_per_stage and convolutional_upsampling=True
- not real vram consumption. just a constant term to which the vram consumption will be approx proportional
- (+ offset for parameter storage)
- :param deep_supervision:
- :param patch_size:
- :param num_pool_per_axis:
- :param base_num_features:
- :param max_num_features:
- :param num_modalities:
- :param num_classes:
- :param pool_op_kernel_sizes:
- :return:
- """
- if not isinstance(num_pool_per_axis, np.ndarray):
- num_pool_per_axis = np.array(num_pool_per_axis)
- npool = len(pool_op_kernel_sizes)
- map_size = np.array(patch_size)
- tmp = np.int64((conv_per_stage * 2 + 1) * np.prod(map_size, dtype=np.int64) * base_num_features +
- num_modalities * np.prod(map_size, dtype=np.int64) +
- num_classes * np.prod(map_size, dtype=np.int64))
- num_feat = base_num_features
- for p in range(npool):
- for pi in range(len(num_pool_per_axis)):
- map_size[pi] /= pool_op_kernel_sizes[p][pi]
- num_feat = min(num_feat * 2, max_num_features)
- num_blocks = (conv_per_stage * 2 + 1) if p < (npool - 1) else conv_per_stage # conv_per_stage + conv_per_stage for the convs of encode/decode and 1 for transposed conv
- tmp += num_blocks * np.prod(map_size, dtype=np.int64) * num_feat
- if deep_supervision and p < (npool - 2):
- tmp += np.prod(map_size, dtype=np.int64) * num_classes
- # print(p, map_size, num_feat, tmp)
- return tmp
Swin_Unet_s_ACDC_2laterdown.py at commit 6357a1d, under MIT · at the source
Overview
- Institute of Pure and Applied Sciences, University of Tsukuba, Tsukuba, Ibaraki, Japan
- Department of Radiology, The University of Tokyo Hospital, Tokyo, Japan
- Department of Radiology, The University of Tokyo, Tokyo, Japan
Abstract
The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.
Repositories
Its files are read in the Code ↔ Paper reader above, with 6 matches between paragraphs and lines of code.
MASILab/Synb0-DISCO
d813f9435cffb5b663f15b52d55ed0a27181f0b2, 31 July 2025Availability: 1 check, the latest on 28 September 2026: the link answers
- 28 September 2026: the link answers
18 files
- data_processing/
PlotLossCurves.m , MATLAB, 42 lines - data_processing/
normalize_T1.sh , Shell, 55 lines - data_processing/
prepare_input.sh , Shell, 129 lines - data_processing/
prepare_truth.sh , Shell, 78 lines - src/
inference.py , Python, 95 lines - src/
model.py , Python, 101 lines - src/
pipeline.sh , Shell, 69 lines - src/
train_lin.py , Python, 344 lines - src/
util.py , Python, 108 lines - v1_0/
build/ , Shell, 33 linesrun_synb0_run.sh - v1_0/
src/ , MATLAB, 130 linesdatRGBtriDWMRI.m - v1_0/
src/ , MATLAB, 119 linesdatRGBtriDWMRIrecon.m - v1_0/
src/ , MATLAB, 27 linessynb0.m - v1_0/
src/ , Python, 20 linessynb0.py - v1_0/
src/ , MATLAB, 17 linessynb0_run.m - v1_0/
src/ , MATLAB, 7 linestest.m - v1_0/
src/ , Python, 10 linestest.py - README.md, Text, 158 lines
junyuchen245/transmorph_transformer_for_medical_image_registration
6357a1d7fc44c36db9b1d1ccaa372409253142cf, 22 May 2025Availability: 1 check, the latest on 28 September 2026: the link answers
- 28 September 2026: the link answers
369 files
- Baseline_Transformers/
data/ , Python, 1 line__init__.py - Baseline_Transformers/
data/ , Python, 66 linesdata_utils.py - Baseline_Transformers/
data/ , Python, 95 linesdatasets.py - Baseline_Transformers/
data/ , Python, 27 linesrand.py - Baseline_Transformers/
data/ , Python, 536 linestrans.py - Baseline_Transformers/
infer_CoTr.py , Python, 149 lines - Baseline_Transformers/
infer_PVT.py , Python, 151 lines - Baseline_Transformers/
infer_ViTVNet.py , Python, 139 lines - Baseline_Transformers/
infer_nnFormer.py , Python, 116 lines - Baseline_Transformers/
losses.py , Python, 554 lines - Baseline_Transformers/
models/ , Python, 5 linesCoTr/ __init__.py - Baseline_Transformers/
models/ , Python, 5 linesCoTr/ configuration.py - Baseline_Transformers/
models/ , Python, 163 linesCoTr/ network_architecture/ CNNBackbone.py - Baseline_Transformers/
models/ , Python, 180 linesCoTr/ network_architecture/ DeTrans/ DeformableTrans.py - Baseline_Transformers/
models/ , Python, 32 linesCoTr/ network_architecture/ DeTrans/ ops/ functions/ ms_deform_attn_func.py - Baseline_Transformers/
models/ , Python, 1 lineCoTr/ network_architecture/ DeTrans/ ops/ modules/ __init__.py - Baseline_Transformers/
models/ , Python, 96 linesCoTr/ network_architecture/ DeTrans/ ops/ modules/ ms_deform_attn.py - Baseline_Transformers/
models/ , Python, 73 linesCoTr/ network_architecture/ DeTrans/ position_encoding.py - Baseline_Transformers/
models/ , Python, 275 linesCoTr/ network_architecture/ ResTranUnet.py - Baseline_Transformers/
models/ , Python, 2 linesCoTr/ network_architecture/ __init__.py - Baseline_Transformers/
models/ , Python, 828 linesCoTr/ network_architecture/ neural_network.py - Baseline_Transformers/
models/ , Python, 2 linesCoTr/ run/ __init__.py - Baseline_Transformers/
models/ , Python, 64 linesCoTr/ run/ default_configuration.py - Baseline_Transformers/
models/ , Python, 137 linesCoTr/ run/ run_training.py - Baseline_Transformers/
models/ , Python, 2 linesCoTr/ training/ __init__.py - Baseline_Transformers/
models/ , Python, 112 linesCoTr/ training/ model_restore.py - Baseline_Transformers/
models/ , Python, 2 linesCoTr/ training/ network_training/ __init__.py - Baseline_Transformers/
models/ , Python, 727 linesCoTr/ training/ network_training/ network_trainer.py - Baseline_Transformers/
models/ , Python, 731 lines, 1 matchCoTr/ training/ network_training/ nnUNetTrainer.py - Baseline_Transformers/
models/ , Python, 388 linesCoTr/ training/ network_training/ nnUNetTrainerV2_ResTrans .py - Baseline_Transformers/
models/ , Python, 493 linesPVT.py - Baseline_Transformers/
models/ , Python, 471 linesViTVNet.py - Baseline_Transformers/
models/ , Python, 43 lines, 1 matchconfigs_PVT.py - Baseline_Transformers/
models/ , Python, 28 linesconfigs_ViTVNet.py - Baseline_Transformers/
models/ , Python, 980 lines, 1 matchnnFormer/ Swin_Unet_l_gelunorm.py - Baseline_Transformers/
models/ , Python, 976 lines, 1 matchnnFormer/ Swin_Unet_s_ACDC_2laterd own.py - Baseline_Transformers/
models/ , Python, 456 lines, 1 matchnnFormer/ generic_UNet.py - Baseline_Transformers/
models/ , Python, 38 linesnnFormer/ initialization.py - Baseline_Transformers/
models/ , Python, 846 linesnnFormer/ neural_network.py - Baseline_Transformers/
train_CoTr.py , Python, 228 lines - Baseline_Transformers/
train_PVT.py , Python, 228 lines - Baseline_Transformers/
train_ViTVNet.py , Python, 228 lines - Baseline_Transformers/
train_nnFormer.py , Python, 229 lines - Baseline_Transformers/
utils.py , Python, 335 lines - Baseline_registration_mo
dels/ , Python, 1 lineCycleMorph/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesCycleMorph/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesCycleMorph/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesCycleMorph/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesCycleMorph/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 117 linesCycleMorph/ infer.py - Baseline_registration_mo
dels/ , Python, 1 lineCycleMorph/ models/ __init__.py - Baseline_registration_mo
dels/ , Python, 62 linesCycleMorph/ models/ base_model.py - Baseline_registration_mo
dels/ , Python, 26 linesCycleMorph/ models/ configs.py - Baseline_registration_mo
dels/ , Python, 208 linesCycleMorph/ models/ cycleMorph_model.py - Baseline_registration_mo
dels/ , Python, 52 linesCycleMorph/ models/ loss.py - Baseline_registration_mo
dels/ , Python, 12 linesCycleMorph/ models/ models.py - Baseline_registration_mo
dels/ , Python, 330 linesCycleMorph/ models/ networks.py - Baseline_registration_mo
dels/ , Python, 181 linesCycleMorph/ train_CycleMorph.py - Baseline_registration_mo
dels/ , Python, 1 lineCycleMorph/ util/ __init__.py - Baseline_registration_mo
dels/ , Python, 115 linesCycleMorph/ util/ get_data.py - Baseline_registration_mo
dels/ , Python, 64 linesCycleMorph/ util/ html.py - Baseline_registration_mo
dels/ , Python, 33 linesCycleMorph/ util/ png.py - Baseline_registration_mo
dels/ , Python, 84 linesCycleMorph/ util/ util.py - Baseline_registration_mo
dels/ , Python, 205 linesCycleMorph/ util/ visualizer.py - Baseline_registration_mo
dels/ , Python, 335 linesCycleMorph/ utils.py - Baseline_registration_mo
dels/ , Python, 1 lineLDDMM/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesLDDMM/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesLDDMM/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesLDDMM/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesLDDMM/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 109 linesLDDMM/ infer_LDDMM.py - Baseline_registration_mo
dels/ , Python, 2,219 linesLDDMM/ torch_lddmm.py - Baseline_registration_mo
dels/ , Python, 335 linesLDDMM/ utils.py - Baseline_registration_mo
dels/ , Python, 1 lineMIDIR/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesMIDIR/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesMIDIR/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesMIDIR/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesMIDIR/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 144 linesMIDIR/ infer.py - Baseline_registration_mo
dels/ , Python, 554 linesMIDIR/ losses.py - Baseline_registration_mo
dels/ , Python, 242 linesMIDIR/ models.py - Baseline_registration_mo
dels/ , Python, 221 linesMIDIR/ train_MIDIR.py - Baseline_registration_mo
dels/ , Python, 257 linesMIDIR/ transformation.py - Baseline_registration_mo
dels/ , Python, 335 linesMIDIR/ utils.py - Baseline_registration_mo
dels/ , Python, 1 lineNiftyReg/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesNiftyReg/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesNiftyReg/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesNiftyReg/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesNiftyReg/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 109 linesNiftyReg/ infer_NiftyReg.py - Baseline_registration_mo
dels/ , Python, 335 linesNiftyReg/ utils.py - Baseline_registration_mo
dels/ , Python, 1 lineSyN/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesSyN/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesSyN/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesSyN/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesSyN/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 96 linesSyN/ infer_SyN.py - Baseline_registration_mo
dels/ , Python, 335 linesSyN/ utils.py - Baseline_registration_mo
dels/ , Python, 1 lineVoxelMorph-diff/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesVoxelMorph-diff/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesVoxelMorph-diff/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesVoxelMorph-diff/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesVoxelMorph-diff/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 525 linesVoxelMorph-diff/ finite_differences.py - Baseline_registration_mo
dels/ , Python, 114 linesVoxelMorph-diff/ infer.py - Baseline_registration_mo
dels/ , Python, 554 linesVoxelMorph-diff/ losses.py - Baseline_registration_mo
dels/ , Python, 633 linesVoxelMorph-diff/ models.py - Baseline_registration_mo
dels/ , Python, 222 linesVoxelMorph-diff/ train_vxm_diff.py - Baseline_registration_mo
dels/ , Python, 335 linesVoxelMorph-diff/ utils.py - Baseline_registration_mo
dels/ , Python, 1 lineVoxelMorph/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesVoxelMorph/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesVoxelMorph/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesVoxelMorph/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesVoxelMorph/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 114 linesVoxelMorph/ infer.py - Baseline_registration_mo
dels/ , Python, 554 linesVoxelMorph/ losses.py - Baseline_registration_mo
dels/ , Python, 485 linesVoxelMorph/ models.py - Baseline_registration_mo
dels/ , Python, 227 linesVoxelMorph/ train_vxm.py - Baseline_registration_mo
dels/ , Python, 335 linesVoxelMorph/ utils.py - Baseline_registration_mo
dels/ , Python, 1 linedeedsBCV/ data/ __init__.py - Baseline_registration_mo
dels/ , Python, 66 linesdeedsBCV/ data/ data_utils.py - Baseline_registration_mo
dels/ , Python, 95 linesdeedsBCV/ data/ datasets.py - Baseline_registration_mo
dels/ , Python, 27 linesdeedsBCV/ data/ rand.py - Baseline_registration_mo
dels/ , Python, 536 linesdeedsBCV/ data/ trans.py - Baseline_registration_mo
dels/ , Python, 103 linesdeedsBCV/ infer_deedsBCV.py - Baseline_registration_mo
dels/ , Python, 335 linesdeedsBCV/ utils.py - Docker/
TransMorph_build_Docker/ , Python, 934 linesTransMorph.py - Docker/
TransMorph_build_Docker/ , Shell, 3 linesbuild.sh - Docker/
TransMorph_build_Docker/ , Shell, 3 linesbuild_GPU.sh - Docker/
TransMorph_build_Docker/ , Python, 372 linesconfigs_TransMorph.py - Docker/
TransMorph_build_Docker/ , Shell, 5 linesexport.sh - Docker/
TransMorph_build_Docker/ , Python, 459 linesinfer_TransMorph.py - Docker/
TransMorph_build_Docker/ , Python, 470 linesinfer_TransMorph_GPU.py - Docker/
TransMorph_build_Docker/ , Python, 173 lineslosses.py - Docker/
TransMorph_build_Docker/ , Shell, 7 linespush.sh - Docker/
TransMorph_build_Docker/ , Shell, 13 linestest.sh - Docker/
TransMorph_build_Docker/ , Shell, 13 linestest_GPU.sh - Docker/
test.sh , Shell, 12 lines - Docker/
test_GPU.sh , Shell, 12 lines - IXI/
Baseline_Transformers/ , Python, 1 linedata/ __init__.py - IXI/
Baseline_Transformers/ , Python, 66 linesdata/ data_utils.py - IXI/
Baseline_Transformers/ , Python, 82 linesdata/ datasets.py - IXI/
Baseline_Transformers/ , Python, 27 linesdata/ rand.py - IXI/
Baseline_Transformers/ , Python, 536 linesdata/ trans.py - IXI/
Baseline_Transformers/ , Python, 135 linesinfer_CoTr.py - IXI/
Baseline_Transformers/ , Python, 138 linesinfer_PVT.py - IXI/
Baseline_Transformers/ , Python, 139 linesinfer_ViTVNet.py - IXI/
Baseline_Transformers/ , Python, 110 linesinfer_nnFormer.py - IXI/
Baseline_Transformers/ , Python, 556 lineslosses.py - IXI/
Baseline_Transformers/ , Python, 5 linesmodels/ CoTr/ __init__.py - IXI/
Baseline_Transformers/ , Python, 5 linesmodels/ CoTr/ configuration.py - IXI/
Baseline_Transformers/ , Python, 163 linesmodels/ CoTr/ network_architecture/ CNNBackbone.py - IXI/
Baseline_Transformers/ , Python, 180 linesmodels/ CoTr/ network_architecture/ DeTrans/ DeformableTrans.py - IXI/
Baseline_Transformers/ , Python, 32 linesmodels/ CoTr/ network_architecture/ DeTrans/ ops/ functions/ ms_deform_attn_func.py - IXI/
Baseline_Transformers/ , Python, 1 linemodels/ CoTr/ network_architecture/ DeTrans/ ops/ modules/ __init__.py - IXI/
Baseline_Transformers/ , Python, 96 linesmodels/ CoTr/ network_architecture/ DeTrans/ ops/ modules/ ms_deform_attn.py - IXI/
Baseline_Transformers/ , Python, 73 linesmodels/ CoTr/ network_architecture/ DeTrans/ position_encoding.py - IXI/
Baseline_Transformers/ , Python, 275 linesmodels/ CoTr/ network_architecture/ ResTranUnet.py - IXI/
Baseline_Transformers/ , Python, 2 linesmodels/ CoTr/ network_architecture/ __init__.py - IXI/
Baseline_Transformers/ , Python, 828 linesmodels/ CoTr/ network_architecture/ neural_network.py - IXI/
Baseline_Transformers/ , Python, 2 linesmodels/ CoTr/ run/ __init__.py - IXI/
Baseline_Transformers/ , Python, 64 linesmodels/ CoTr/ run/ default_configuration.py - IXI/
Baseline_Transformers/ , Python, 137 linesmodels/ CoTr/ run/ run_training.py - IXI/
Baseline_Transformers/ , Python, 2 linesmodels/ CoTr/ training/ __init__.py - IXI/
Baseline_Transformers/ , Python, 112 linesmodels/ CoTr/ training/ model_restore.py - IXI/
Baseline_Transformers/ , Python, 2 linesmodels/ CoTr/ training/ network_training/ __init__.py - IXI/
Baseline_Transformers/ , Python, 727 linesmodels/ CoTr/ training/ network_training/ network_trainer.py - IXI/
Baseline_Transformers/ , Python, 731 linesmodels/ CoTr/ training/ network_training/ nnUNetTrainer.py - IXI/
Baseline_Transformers/ , Python, 388 linesmodels/ CoTr/ training/ network_training/ nnUNetTrainerV2_ResTrans .py - IXI/
Baseline_Transformers/ , Python, 493 linesmodels/ PVT.py - IXI/
Baseline_Transformers/ , Python, 471 linesmodels/ ViTVNet.py - IXI/
Baseline_Transformers/ , Python, 43 linesmodels/ configs_PVT.py - IXI/
Baseline_Transformers/ , Python, 28 linesmodels/ configs_ViTVNet.py - IXI/
Baseline_Transformers/ , Python, 980 linesmodels/ nnFormer/ Swin_Unet_l_gelunorm.py - IXI/
Baseline_Transformers/ , Python, 976 linesmodels/ nnFormer/ Swin_Unet_s_ACDC_2laterd own.py - IXI/
Baseline_Transformers/ , Python, 456 linesmodels/ nnFormer/ generic_UNet.py - IXI/
Baseline_Transformers/ , Python, 38 linesmodels/ nnFormer/ initialization.py - IXI/
Baseline_Transformers/ , Python, 846 linesmodels/ nnFormer/ neural_network.py - IXI/
Baseline_Transformers/ , Python, 213 linestrain_CoTr.py - IXI/
Baseline_Transformers/ , Python, 214 linestrain_PVT.py - IXI/
Baseline_Transformers/ , Python, 214 linestrain_ViTVNet.py - IXI/
Baseline_Transformers/ , Python, 214 linestrain_nnFormer.py - IXI/
Baseline_Transformers/ , Python, 352 linesutils.py - IXI/
Baseline_registration_me , Python, 1 linethods/ CycleMorph/ data/ __init__.py - IXI/
Baseline_registration_me , Python, 66 linesthods/ CycleMorph/ data/ data_utils.py - IXI/
Baseline_registration_me , Python, 82 linesthods/ CycleMorph/ data/ datasets.py - IXI/
Baseline_registration_me , Python, 27 linesthods/ CycleMorph/ data/ rand.py - IXI/
Baseline_registration_me , Python, 536 linesthods/ CycleMorph/ data/ trans.py - IXI/
Baseline_registration_me , Python, 137 linesthods/ CycleMorph/ infer.py - IXI/
Baseline_registration_me , Python, 719 linesthods/ CycleMorph/ losses.py - IXI/
Baseline_registration_me , Python, 1 linethods/ CycleMorph/ models/ __init__.py - IXI/
Baseline_registration_me , Python, 62 linesthods/ CycleMorph/ models/ base_model.py - IXI/
Baseline_registration_me , Python, 28 linesthods/ CycleMorph/ models/ configs.py - IXI/
Baseline_registration_me , Python, 221 linesthods/ CycleMorph/ models/ cycleMorph_model.py - IXI/
Baseline_registration_me , Python, 144 linesthods/ CycleMorph/ models/ loss.py - IXI/
Baseline_registration_me , Python, 12 linesthods/ CycleMorph/ models/ models.py - IXI/
Baseline_registration_me , Python, 489 linesthods/ CycleMorph/ models/ networks.py - IXI/
Baseline_registration_me , Python, 206 linesthods/ CycleMorph/ train.py - IXI/
Baseline_registration_me , Python, 1 linethods/ CycleMorph/ util/ __init__.py - IXI/
Baseline_registration_me , Python, 115 linesthods/ CycleMorph/ util/ get_data.py - IXI/
Baseline_registration_me , Python, 64 linesthods/ CycleMorph/ util/ html.py - IXI/
Baseline_registration_me , Python, 33 linesthods/ CycleMorph/ util/ png.py - IXI/
Baseline_registration_me , Python, 84 linesthods/ CycleMorph/ util/ util.py - IXI/
Baseline_registration_me , Python, 205 linesthods/ CycleMorph/ util/ visualizer.py - IXI/
Baseline_registration_me , Python, 521 linesthods/ CycleMorph/ utils.py - IXI/
Baseline_registration_me , Python, 1 linethods/ MIDIR/ data/ __init__.py - IXI/
Baseline_registration_me , Python, 66 linesthods/ MIDIR/ data/ data_utils.py - IXI/
Baseline_registration_me , Python, 82 linesthods/ MIDIR/ data/ datasets.py - IXI/
Baseline_registration_me , Python, 27 linesthods/ MIDIR/ data/ rand.py - IXI/
Baseline_registration_me , Python, 536 linesthods/ MIDIR/ data/ trans.py - IXI/
Baseline_registration_me , Python, 133 linesthods/ MIDIR/ infer.py - IXI/
Baseline_registration_me , Python, 556 linesthods/ MIDIR/ losses.py - IXI/
Baseline_registration_me , Python, 242 linesthods/ MIDIR/ models.py - IXI/
Baseline_registration_me , Python, 206 linesthods/ MIDIR/ train_MIDIR.py - IXI/
Baseline_registration_me , Python, 257 linesthods/ MIDIR/ transformation.py - IXI/
Baseline_registration_me , Python, 352 linesthods/ MIDIR/ utils.py - IXI/
Baseline_registration_me , Python, 1 linethods/ VoxelMorph-diff/ data/ __init__.py - IXI/
Baseline_registration_me , Python, 66 linesthods/ VoxelMorph-diff/ data/ data_utils.py - IXI/
Baseline_registration_me , Python, 82 linesthods/ VoxelMorph-diff/ data/ datasets.py - IXI/
Baseline_registration_me , Python, 27 linesthods/ VoxelMorph-diff/ data/ rand.py - IXI/
Baseline_registration_me , Python, 536 linesthods/ VoxelMorph-diff/ data/ trans.py - IXI/
Baseline_registration_me , Python, 525 linesthods/ VoxelMorph-diff/ finite_differences.py - IXI/
Baseline_registration_me , Python, 108 linesthods/ VoxelMorph-diff/ infer.py - IXI/
Baseline_registration_me , Python, 554 linesthods/ VoxelMorph-diff/ losses.py - IXI/
Baseline_registration_me , Python, 633 linesthods/ VoxelMorph-diff/ models.py - IXI/
Baseline_registration_me , Python, 210 linesthods/ VoxelMorph-diff/ train_vxm_diff.py - IXI/
Baseline_registration_me , Python, 352 linesthods/ VoxelMorph-diff/ utils.py - IXI/
Baseline_registration_me , Python, 1 linethods/ VoxelMorph/ data/ __init__.py - IXI/
Baseline_registration_me , Python, 66 linesthods/ VoxelMorph/ data/ data_utils.py - IXI/
Baseline_registration_me , Python, 82 linesthods/ VoxelMorph/ data/ datasets.py - IXI/
Baseline_registration_me , Python, 27 linesthods/ VoxelMorph/ data/ rand.py - IXI/
Baseline_registration_me , Python, 536 linesthods/ VoxelMorph/ data/ trans.py - IXI/
Baseline_registration_me , Python, 136 linesthods/ VoxelMorph/ infer.py - IXI/
Baseline_registration_me , Python, 556 linesthods/ VoxelMorph/ losses.py - IXI/
Baseline_registration_me , Python, 485 linesthods/ VoxelMorph/ models.py - IXI/
Baseline_registration_me , Python, 212 linesthods/ VoxelMorph/ train_vxm.py - IXI/
Baseline_registration_me , Python, 352 linesthods/ VoxelMorph/ utils.py - IXI/
Baseline_traditional_met , Python, 1 linehods/ LDDMM/ data_IXI/ __init__.py - IXI/
Baseline_traditional_met , Python, 66 lineshods/ LDDMM/ data_IXI/ data_utils.py - IXI/
Baseline_traditional_met , Python, 82 lineshods/ LDDMM/ data_IXI/ datasets.py - IXI/
Baseline_traditional_met , Python, 27 lineshods/ LDDMM/ data_IXI/ rand.py - IXI/
Baseline_traditional_met , Python, 536 lineshods/ LDDMM/ data_IXI/ trans.py - IXI/
Baseline_traditional_met , Python, 119 lineshods/ LDDMM/ infer_IXI.py - IXI/
Baseline_traditional_met , Python, 2,219 lineshods/ LDDMM/ torch_lddmm.py - IXI/
Baseline_traditional_met , Python, 277 lineshods/ LDDMM/ utils.py - IXI/
Baseline_traditional_met , Python, 1 linehods/ NiftyReg/ data_IXI/ __init__.py - IXI/
Baseline_traditional_met , Python, 66 lineshods/ NiftyReg/ data_IXI/ data_utils.py - IXI/
Baseline_traditional_met , Python, 82 lineshods/ NiftyReg/ data_IXI/ datasets.py - IXI/
Baseline_traditional_met , Python, 27 lineshods/ NiftyReg/ data_IXI/ rand.py - IXI/
Baseline_traditional_met , Python, 536 lineshods/ NiftyReg/ data_IXI/ trans.py - IXI/
Baseline_traditional_met , Python, 117 lineshods/ NiftyReg/ infer_IXI.py - IXI/
Baseline_traditional_met , Python, 277 lineshods/ NiftyReg/ utils.py - IXI/
Baseline_traditional_met , Python, 1 linehods/ SyN/ data_IXI/ __init__.py - IXI/
Baseline_traditional_met , Python, 66 lineshods/ SyN/ data_IXI/ data_utils.py - IXI/
Baseline_traditional_met , Python, 82 lineshods/ SyN/ data_IXI/ datasets.py - IXI/
Baseline_traditional_met , Python, 27 lineshods/ SyN/ data_IXI/ rand.py - IXI/
Baseline_traditional_met , Python, 536 lineshods/ SyN/ data_IXI/ trans.py - IXI/
Baseline_traditional_met , Python, 110 lineshods/ SyN/ infer_IXI.py - IXI/
Baseline_traditional_met , Python, 277 lineshods/ SyN/ utils.py - IXI/
Baseline_traditional_met , Python, 1 linehods/ deedsBCV/ data_IXI/ __init__.py - IXI/
Baseline_traditional_met , Python, 66 lineshods/ deedsBCV/ data_IXI/ data_utils.py - IXI/
Baseline_traditional_met , Python, 82 lineshods/ deedsBCV/ data_IXI/ datasets.py - IXI/
Baseline_traditional_met , Python, 27 lineshods/ deedsBCV/ data_IXI/ rand.py - IXI/
Baseline_traditional_met , Python, 536 lineshods/ deedsBCV/ data_IXI/ trans.py - IXI/
Baseline_traditional_met , Python, 109 lineshods/ deedsBCV/ infer_IXI.py - IXI/
Baseline_traditional_met , Python, 277 lineshods/ deedsBCV/ utils.py - IXI/
TransMorph/ , Python, 1 linedata/ __init__.py - IXI/
TransMorph/ , Python, 66 linesdata/ data_utils.py - IXI/
TransMorph/ , Python, 82 linesdata/ datasets.py - IXI/
TransMorph/ , Python, 27 linesdata/ rand.py - IXI/
TransMorph/ , Python, 536 linesdata/ trans.py - IXI/
TransMorph/ , Python, 115 linesinfer_TransMorph.py - IXI/
TransMorph/ , Python, 109 linesinfer_TransMorph_Bayes.p y - IXI/
TransMorph/ , Python, 107 linesinfer_TransMorph_bspl.py - IXI/
TransMorph/ , Python, 109 linesinfer_TransMorph_diff.py - IXI/
TransMorph/ , Python, 556 lineslosses.py - IXI/
TransMorph/ , Python, 894 linesmodels/ TransMorph.py - IXI/
TransMorph/ , Python, 890 linesmodels/ TransMorph_Bayes.py - IXI/
TransMorph/ , Python, 173 linesmodels/ TransMorph_bspl.py - IXI/
TransMorph/ , Python, 624 linesmodels/ TransMorph_diff.py - IXI/
TransMorph/ , Python, 1 linemodels/ __init__.py - IXI/
TransMorph/ , Python, 318 lines, 1 matchmodels/ configs_TransMorph.py - IXI/
TransMorph/ , Python, 57 linesmodels/ configs_TransMorph_Bayes .py - IXI/
TransMorph/ , Python, 57 linesmodels/ configs_TransMorph_bspl. py - IXI/
TransMorph/ , Python, 93 linesmodels/ configs_TransMorph_diff. py - IXI/
TransMorph/ , Python, 527 linesmodels/ finite_differences.py - IXI/
TransMorph/ , Python, 257 linesmodels/ transformation.py - IXI/
TransMorph/ , Python, 214 linestrain_TransMorph.py - IXI/
TransMorph/ , Python, 222 linestrain_TransMorph_Bayes.p y - IXI/
TransMorph/ , Python, 208 linestrain_TransMorph_bspl.py - IXI/
TransMorph/ , Python, 215 linestrain_TransMorph_diff.py - IXI/
TransMorph/ , Python, 352 linesutils.py - IXI/
analysis.py , Python, 136 lines - IXI/
analysis_trans.py , Python, 128 lines - OASIS/
TransMorph/ , Python, 1 linedata/ __init__.py - OASIS/
TransMorph/ , Python, 66 linesdata/ data_utils.py - OASIS/
TransMorph/ , Python, 70 linesdata/ datasets.py - OASIS/
TransMorph/ , Python, 27 linesdata/ rand.py - OASIS/
TransMorph/ , Python, 536 linesdata/ trans.py - OASIS/
TransMorph/ , Python, 578 lineslosses.py - OASIS/
TransMorph/ , Python, 883 linesmodels/ TransMorph.py - OASIS/
TransMorph/ , Python, 1 linemodels/ __init__.py - OASIS/
TransMorph/ , Python, 318 linesmodels/ configs_TransMorph.py - OASIS/
TransMorph/ , Python, 66 linessubmit_TransMorph.py - OASIS/
TransMorph/ , Python, 252 linestrain_TransMorph.py - OASIS/
TransMorph/ , Python, 357 linesutils.py - OASIS/
evaluation.py , Python, 292 lines - OASIS/
surface_distance/ , Python, 17 lines__init__.py - OASIS/
surface_distance/ , Python, 400 lineslookup_tables.py - OASIS/
surface_distance/ , Python, 443 linesmetrics.py - RaFD/
TransMorph2D/ , Python, 1 linedata/ __init__.py - RaFD/
TransMorph2D/ , Python, 66 linesdata/ data_utils.py - RaFD/
TransMorph2D/ , Python, 57 linesdata/ datasets.py - RaFD/
TransMorph2D/ , Python, 27 linesdata/ rand.py - RaFD/
TransMorph2D/ , Python, 539 linesdata/ trans.py - RaFD/
TransMorph2D/ , Python, 103 linesinfer_TransMorph.py - RaFD/
TransMorph2D/ , Python, 715 lineslosses.py - RaFD/
TransMorph2D/ , Python, 848 linesmodels/ TransMorph.py - RaFD/
TransMorph2D/ , Python, 871 linesmodels/ TransMorph_Bayes.py - RaFD/
TransMorph2D/ , Python, 170 linesmodels/ TransMorph_bspl.py - RaFD/
TransMorph2D/ , Python, 601 linesmodels/ TransMorph_diff.py - RaFD/
TransMorph2D/ , Python, 1 linemodels/ __init__.py - RaFD/
TransMorph2D/ , Python, 317 linesmodels/ configs_TransMorph.py - RaFD/
TransMorph2D/ , Python, 55 linesmodels/ configs_TransMorph_Bayes .py - RaFD/
TransMorph2D/ , Python, 57 linesmodels/ configs_TransMorph_bspl. py - RaFD/
TransMorph2D/ , Python, 61 linesmodels/ configs_TransMorph_diff. py - RaFD/
TransMorph2D/ , Python, 527 linesmodels/ finite_differences.py - RaFD/
TransMorph2D/ , Python, 257 linesmodels/ transformation.py - RaFD/
TransMorph2D/ , Python, 248 linestrain_TransMorph.py - RaFD/
TransMorph2D/ , Python, 362 linesutils.py - TransMorph/
data/ , Python, 1 line__init__.py - TransMorph/
data/ , Python, 66 linesdata_utils.py - TransMorph/
data/ , Python, 95 linesdatasets.py - TransMorph/
data/ , Python, 27 linesrand.py - TransMorph/
data/ , Python, 536 linestrans.py - TransMorph/
infer_TransMorph.py , Python, 116 lines - TransMorph/
infer_TransMorph_Bayes.p , Python, 123 linesy - TransMorph/
infer_TransMorph_bspl.py , Python, 113 lines - TransMorph/
infer_TransMorph_diff.py , Python, 115 lines - TransMorph/
losses.py , Python, 622 lines - TransMorph/
models/ , Python, 890 linesTransMorph.py - TransMorph/
models/ , Python, 908 linesTransMorph_Bayes.py - TransMorph/
models/ , Python, 181 linesTransMorph_bspl.py - TransMorph/
models/ , Python, 605 linesTransMorph_diff.py - TransMorph/
models/ , Python, 1 line__init__.py - TransMorph/
models/ , Python, 325 linesconfigs_TransMorph.py - TransMorph/
models/ , Python, 63 linesconfigs_TransMorph_Bayes .py - TransMorph/
models/ , Python, 63 linesconfigs_TransMorph_bspl. py - TransMorph/
models/ , Python, 67 linesconfigs_TransMorph_diff. py - TransMorph/
models/ , Python, 527 linesfinite_differences.py - TransMorph/
models/ , Python, 256 linestransformation.py - TransMorph/
train_TransMorph.py , Python, 230 lines - TransMorph/
train_TransMorph_Bayes.p , Python, 237 linesy - TransMorph/
train_TransMorph_bspl.py , Python, 225 lines - TransMorph/
train_TransMorph_diff.py , Python, 226 lines - TransMorph/
utils.py , Python, 341 lines - TransMorph_affine/
data/ , Python, 1 line__init__.py - TransMorph_affine/
data/ , Python, 69 linesdata_utils.py - TransMorph_affine/
data/ , Python, 82 linesdatasets.py - TransMorph_affine/
data/ , Python, 27 linesrand.py - TransMorph_affine/
data/ , Python, 536 linestrans.py - TransMorph_affine/
losses.py , Python, 721 lines - TransMorph_affine/
models/ , Python, 953 linesTransMorph_affine.py - TransMorph_affine/
models/ , Python, 1 line__init__.py - TransMorph_affine/
models/ , Python, 372 linesconfigs_TransMorph.py - TransMorph_affine/
train_TransMorph_affine. , Python, 253 linespy - TransMorph_affine/
utils.py , Python, 352 lines - LICENSE, License, 21 lines
- README.md, Text, 123 lines
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:
- 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 384 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.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 28 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 9 authors, 4 keywords, 11 MeSH terms, 1 funder, 15 references.
Cite
This paper
Ueyama, T., Takahashi, E., Fujita, N., Suzuki, Y., Inui, S., Yasaka, K., Iwanaga, H., Abe, O., & Terada, Y. (2026). Image Distortion Correction for Diffusion MR Imaging Using a Transformer-based U-Net. Magnetic resonance in medical sciences : MRMS : an official journal of Japan Society of Magnetic Resonance in Medicine, 25(2), 2024-0149. https://
BibTeX
@article{ueyama2026image
author = {Ueyama, Tsuyoshi and Takahashi, Erika and Fujita, Naoto and Suzuki, Yuichi and Inui, Shohei and Yasaka, Koichiro and Iwanaga, Hideyuki and Abe, Osamu and Terada, Yasuhiko},
title = {{Image Distortion Correction for Diffusion MR Imaging Using a Transformer-based U-Net}},
journal = {Magnetic resonance in medical sciences : MRMS : an official journal of Japan Society of Magnetic Resonance in Medicine},
year = {2026},
month = may,
volume = {25},
number = {2},
pages = {2024--0149},
publisher = {Japanese Society for Magnetic Resonance in Medicine},
issn = {1347-3182},
doi = {10.2463/
url = {https://
pmid = {42128846},
pmcid = {PMC13500223}
}
RIS
TY - JOUR
AU - Ueyama, Tsuyoshi
AU - Takahashi, Erika
AU - Fujita, Naoto
AU - Suzuki, Yuichi
AU - Inui, Shohei
AU - Yasaka, Koichiro
AU - Iwanaga, Hideyuki
AU - Abe, Osamu
AU - Terada, Yasuhiko
TI - Image Distortion Correction for Diffusion MR Imaging Using a Transformer-based U-Net
T2 - Magnetic resonance in medical sciences : MRMS : an official journal of Japan Society of Magnetic Resonance in Medicine
J2 - Magn Reson Med Sci
PY - 2026
DA - 2026/
VL - 25
IS - 2
SP - 2024
EP - 0149
SN - 1347-3182
PB - Japanese Society for Magnetic Resonance in Medicine
DO - 10.2463/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.2463/
"type": "article-journal",
"title": "Image Distortion Correction for Diffusion MR Imaging Using a Transformer-based U-Net",
"container-title": "Magnetic resonance in medical sciences : MRMS : an official journal of Japan Society of Magnetic Resonance in Medicine",
"author": [
{
"family": "Ueyama",
"given": "Tsuyoshi"
},
{
"family": "Takahashi",
"given": "Erika"
},
{
"family": "Fujita",
"given": "Naoto"
},
{
"family": "Suzuki",
"given": "Yuichi"
},
{
"family": "Inui",
"given": "Shohei"
},
{
"family": "Yasaka",
"given": "Koichiro"
},
{
"family": "Iwanaga",
"given": "Hideyuki"
},
{
"family": "Abe",
"given": "Osamu"
},
{
"family": "Terada",
"given": "Yasuhiko"
}
],
"container-title-short":
"volume": "25",
"issue": "2",
"page": "2024-0149",
"DOI": "10.2463/
"PMID": "42128846",
"PMCID": "PMC13500223",
"ISSN": "1347-3182",
"publisher": "Japanese Society for Magnetic Resonance in Medicine",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
14
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1002/alz.71649 [code]
- Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.Journal: Alzheimer's & dementia : the journal of the Alzheimer's AssociationIn common: nnU-Net, ANTs, FreeSurfer, 10 other tools, structural MRI / diffusion
- [2] doi:10.3390/jimaging12070276 [code]
- Hyperelastic Regularization for Near-Diffeomorphic Transformer-Based Brain MRI Registration.Journal: Journal of imagingIn common: nnU-Net, ANTs, scikit-image, 8 other tools, structural MRI / diffusion, 2 references
- [3] doi:10.1002/epi.70296 [code]
- Fully automated three-dimensional deep learning-based magnetic resonance imaging segmentation of brain cavities in epilepsy surgery.Journal: EpilepsiaIn common: nnU-Net, ANTs, FSL, 9 other tools, structural MRI / diffusion, clinical / translational
- [4] doi:10.1162/imag.a.1262 [code]
- Frame-wise multi-echo distortion correction for superior functional MRI.Journal: Imaging neuroscience (Cambridge, Mass.)In common: Tools for NIfTI and ANALYZE image (MATLAB), FreeSurfer, FSL, 8 other tools, 2 references
- [5] doi:10.1371/journal.pbio.3003856 [code]
- Aging and metabolism contribute separately to brain-body health.Journal: PLoS biologyIn common: Tools for NIfTI and ANALYZE image (MATLAB), ANTs, FreeSurfer, 9 other tools, structural MRI / diffusion, clinical / translational
- [6] doi:10.1093/braincomms/fcag134 [code]
- Neurophysiological, imaging and neurobiological markers of central fatigue in multiple sclerosis.Journal: Brain communicationsIn common: Tools for NIfTI and ANALYZE image (MATLAB), ANTs, FreeSurfer, 8 other tools, structural MRI / diffusion, 1 reference
- [7] doi:10.1016/j.xcrm.2026.102943 [code]
- Parent-of-origin effects in Alzheimer's liability dissociate neurocognitive and cardiovascular traits in at-risk individuals.Journal: Cell reports. MedicineIn common: Tools for NIfTI and ANALYZE image (MATLAB), ANTs, FreeSurfer, 9 other tools, clinical / translational
- [8] doi:10.7554/elife.108408 [code]
- Frequency and laminar profile of feature-specific visual activity revealed by interleaved EEG-fMRI.Journal: eLifeIn common: Tools for NIfTI and ANALYZE image (MATLAB), ANTs, FreeSurfer, 8 other tools, 1 reference
- [9] doi:10.1038/s41586-026-10631-3 [code]
- A prognostic human brain network for diffuse midline glioma.Journal: NatureIn common: Tools for NIfTI and ANALYZE image (MATLAB), ANTs, FreeSurfer, 8 other tools, clinical / translational, 1 reference
- [10] doi:10.1162/imag.a.1222 [code]
- Network-based near-scalp personalized brain stimulation targets.Journal: Imaging neuroscience (Cambridge, Mass.)In common: Tools for NIfTI and ANALYZE image (MATLAB), ANTs, FreeSurfer, 8 other tools, 1 reference
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
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: 2 repositories of the authors' code, each at its verified commit and with its license, 384 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:cdab20858187e1b9…
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.
