OSCR

Image Distortion Correction for Diffusion MR Imaging Using a Transformer-based U-Net.

Code ↔ Paper

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

The 6 matches
  1. [1] § 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. [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. [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. [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. [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. [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

  1. from einops import rearrange
  2. from copy import deepcopy
  3. from nnformer.utilities.nd_softmax import softmax_helper
  4. from torch import nn
  5. import torch
  6. import numpy as np
  7. from nnformer.network_architecture.initialization import InitWeights_He
  8. from nnformer.network_architecture.neural_network import SegmentationNetwork
  9. import torch.nn.functional
  10. import torch.nn.functional as F
  11. import torch.utils.checkpoint as checkpoint
  12. from timm.models.layers import DropPath, to_3tuple, trunc_normal_
  13. class Mlp(nn.Module):
  14. """ Multilayer perceptron."""
  15. def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
  16. super().__init__()
  17. out_features = out_features or in_features
  18. hidden_features = hidden_features or in_features
  19. self.fc1 = nn.Linear(in_features, hidden_features)
  20. self.act = act_layer()
  21. self.fc2 = nn.Linear(hidden_features, out_features)
  22. self.drop = nn.Dropout(drop)
  23. def forward(self, x):
  24. x = self.fc1(x)
  25. x = self.act(x)
  26. x = self.drop(x)
  27. x = self.fc2(x)
  28. x = self.drop(x)
  29. return x
  30. def window_partition(x, window_size):
  31. B, S, H, W, C = x.shape
  32. 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)
  33. windows = x.permute(0, 1, 3, 5, 2, 4, 6, 7).contiguous().view(-1, window_size[0], window_size[1], window_size[2], C)
  34. return windows
  35. def window_reverse(windows, window_size, S, H, W):
  36. B = int(windows.shape[0] / (S * H * W / window_size[0] / window_size[1] / window_size[2]))
  37. 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)
  38. x = x.permute(0, 1, 4, 2, 5, 3, 6, 7).contiguous().view(B, S, H, W, -1)
  39. return x
  40. class WindowAttention(nn.Module):
  41. """ Window based multi-head self attention (W-MSA) module with relative position bias.
  42. It supports both of shifted and non-shifted window.
  43. Args:
  44. dim (int): Number of input channels.
  45. window_size (tuple[int]): The height and width of the window.
  46. num_heads (int): Number of attention heads.
  47. qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
  48. qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
  49. attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
  50. proj_drop (float, optional): Dropout ratio of output. Default: 0.0
  51. """
  52. def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):
  53. super().__init__()
  54. self.dim = dim
  55. self.window_size = window_size # Wh, Ww
  56. self.num_heads = num_heads
  57. head_dim = dim // num_heads
  58. self.scale = qk_scale or head_dim ** -0.5
  59. # define a parameter table of relative position bias
  60. self.relative_position_bias_table = nn.Parameter(
  61. torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1) * (2 * window_size[2] - 1),
  62. num_heads))
  63. # get pair-wise relative position index for each token inside the window
  64. coords_s = torch.arange(self.window_size[0])
  65. coords_h = torch.arange(self.window_size[1])
  66. coords_w = torch.arange(self.window_size[2])
  67. coords = torch.stack(torch.meshgrid([coords_s, coords_h, coords_w]))
  68. coords_flatten = torch.flatten(coords, 1)
  69. relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :]
  70. relative_coords = relative_coords.permute(1, 2, 0).contiguous()
  71. relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
  72. relative_coords[:, :, 1] += self.window_size[1] - 1
  73. relative_coords[:, :, 2] += self.window_size[2] - 1
  74. relative_coords[:, :, 0] *= 3 * self.window_size[1] - 1
  75. relative_coords[:, :, 1] *= 2 * self.window_size[1] - 1
  76. relative_position_index = relative_coords.sum(-1)
  77. self.register_buffer("relative_position_index", relative_position_index)
  78. self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
  79. self.attn_drop = nn.Dropout(attn_drop)
  80. self.proj = nn.Linear(dim, dim)
  81. self.proj_drop = nn.Dropout(proj_drop)
  82. trunc_normal_(self.relative_position_bias_table, std=.02)
  83. self.softmax = nn.Softmax(dim=-1)
  84. def forward(self, x, mask=None):
  85. B_, N, C = x.shape
  86. qkv = self.qkv(x)
  87. qkv=qkv.reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
  88. q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
  89. q = q * self.scale
  90. attn = (q @ k.transpose(-2, -1))
  91. relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
  92. self.window_size[0] * self.window_size[1] * self.window_size[2],
  93. self.window_size[0] * self.window_size[1] * self.window_size[2], -1)
  94. relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous()
  95. attn = attn + relative_position_bias.unsqueeze(0)
  96. if mask is not None:
  97. nW = mask.shape[0]
  98. attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
  99. attn = attn.view(-1, self.num_heads, N, N)
  100. attn = self.softmax(attn)
  101. else:
  102. attn = self.softmax(attn)
  103. attn = self.attn_drop(attn)
  104. x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
  105. x = self.proj(x)
  106. x = self.proj_drop(x)
  107. return x
  108. class SwinTransformerBlock(nn.Module):
  109. """ Swin Transformer Block.
  110. Args:
  111. dim (int): Number of input channels.
  112. num_heads (int): Number of attention heads.
  113. window_size (int): Window size.
  114. shift_size (int): Shift size for SW-MSA.
  115. mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
  116. qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
  117. qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
  118. drop (float, optional): Dropout rate. Default: 0.0
  119. attn_drop (float, optional): Attention dropout rate. Default: 0.0
  120. drop_path (float, optional): Stochastic depth rate. Default: 0.0
  121. act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
  122. norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
  123. """
  124. def __init__(self, dim, input_resolution, num_heads, window_size=7, shift_size=0,
  125. mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
  126. act_layer=nn.GELU, norm_layer=nn.LayerNorm):
  127. super().__init__()
  128. self.dim = dim
  129. self.input_resolution = input_resolution
  130. self.num_heads = num_heads
  131. self.window_size = window_size
  132. self.shift_size = shift_size
  133. self.mlp_ratio = mlp_ratio
  134. if tuple(self.input_resolution) == tuple(self.window_size):
  135. # if window size is larger than input resolution, we don't partition windows
  136. self.shift_size = [0,0,0]
  137. #self.window_size = min(self.input_resolution)
  138. #assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
  139. self.norm1 = norm_layer(dim)
  140. self.attn = WindowAttention(
  141. dim, window_size=self.window_size, num_heads=num_heads,
  142. qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
  143. self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
  144. self.norm2 = norm_layer(dim)
  145. mlp_hidden_dim = int(dim * mlp_ratio)
  146. self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
  147. def forward(self, x, mask_matrix):
  148. B, L, C = x.shape
  149. S, H, W = self.input_resolution
  150. assert L == S * H * W, "input feature has wrong size"
  151. shortcut = x
  152. x = self.norm1(x)
  153. x = x.view(B, S, H, W, C)
  154. # pad feature maps to multiples of window size
  155. pad_r = (self.window_size[2] - W % self.window_size[2]) % self.window_size[2]
  156. pad_b = (self.window_size[1] - H % self.window_size[1]) % self.window_size[1]
  157. pad_g = (self.window_size[0] - S % self.window_size[0]) % self.window_size[0]
  158. x = F.pad(x, (0, 0, 0, pad_r, 0, pad_b, 0, pad_g))
  159. _, Sp, Hp, Wp, _ = x.shape
  160. # cyclic shift
  161. if min(self.shift_size) > 0:
  162. shifted_x = torch.roll(x, shifts=(-self.shift_size[0], -self.shift_size[1],-self.shift_size[2]), dims=(1, 2,3))
  163. attn_mask = mask_matrix
  164. else:
  165. shifted_x = x
  166. attn_mask = None
  167. # partition windows
  168. x_windows = window_partition(shifted_x, self.window_size)
  169. x_windows = x_windows.view(-1, self.window_size[0] * self.window_size[1] * self.window_size[2],
  170. C)
  171. # W-MSA/SW-MSA
  172. attn_windows = self.attn(x_windows, mask=attn_mask)
  173. # merge windows
  174. attn_windows = attn_windows.view(-1, self.window_size[0], self.window_size[1], self.window_size[2], C)
  175. shifted_x = window_reverse(attn_windows, self.window_size, Sp, Hp, Wp)
  176. # reverse cyclic shift
  177. if min(self.shift_size) > 0:
  178. x = torch.roll(shifted_x, shifts=(self.shift_size[0], self.shift_size[1], self.shift_size[2]), dims=(1, 2, 3))
  179. else:
  180. x = shifted_x
  181. if pad_r > 0 or pad_b > 0 or pad_g > 0:
  182. x = x[:, :S, :H, :W, :].contiguous()
  183. x = x.view(B, S * H * W, C)
  184. # FFN
  185. x = shortcut + self.drop_path(x)
  186. x = x + self.drop_path(self.mlp(self.norm2(x)))
  187. return x
  188. class PatchMerging(nn.Module):
  189. """ Patch Merging Layer
  190. Args:
  191. dim (int): Number of input channels.
  192. norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
  193. """
  194. def __init__(self, dim, norm_layer=nn.LayerNorm,tag=None):
  195. super().__init__()
  196. self.dim = dim
  197. if tag==0:
  198. self.reduction = nn.Conv3d(dim,dim*2,kernel_size=[1,2,2],stride=[1,2,2])
  199. else:
  200. self.reduction = nn.Conv3d(dim,dim*2,kernel_size=[2,2,2],stride=[2,2,2])
  201. self.norm = norm_layer(dim)
  202. def forward(self, x, S, H, W):
  203. B, L, C = x.shape
  204. assert L == H * W * S, "input feature has wrong size"
  205. x = x.view(B, S, H, W, C)
  206. x = F.gelu(x)
  207. x = self.norm(x)
  208. x=x.permute(0,4,1,2,3)
  209. x=self.reduction(x)
  210. x=x.permute(0,2,3,4,1).view(B,-1,2*C)
  211. return x
  212. class Patch_Expanding(nn.Module):
  213. def __init__(self, dim, norm_layer=nn.LayerNorm,tag=None):
  214. super().__init__()
  215. self.dim = dim
  216. self.norm = norm_layer(dim)
  217. if tag==0:
  218. self.up=nn.ConvTranspose3d(dim,dim//2,[1,2,2],[1,2,2])
  219. elif tag==1:
  220. self.up=nn.ConvTranspose3d(dim,dim//2,[2,2,2],[2,2,2])
  221. elif tag==2:
  222. self.up=nn.ConvTranspose3d(dim,dim//2,[2,2,2],[2,2,2],output_padding=[1,0,0])
  223. def forward(self, x, S, H, W):
  224. """ Forward function.
  225. Args:
  226. x: Input feature, tensor size (B, H*W, C).
  227. H, W: Spatial resolution of the input feature.
  228. """
  229. B, L, C = x.shape
  230. assert L == H * W * S, "input feature has wrong size"
  231. x = x.view(B, S, H, W, C)
  232. x = self.norm(x)
  233. x=x.permute(0,4,1,2,3)
  234. x = self.up(x)
  235. x=x.permute(0,2,3,4,1).view(B,-1,C//2)
  236. return x
  237. class BasicLayer(nn.Module):
  238. """ A basic Swin Transformer layer for one stage.
  239. Args:
  240. dim (int): Number of feature channels
  241. depth (int): Depths of this stage.
  242. num_heads (int): Number of attention head.
  243. window_size (int): Local window size. Default: 7.
  244. mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
  245. qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
  246. qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
  247. drop (float, optional): Dropout rate. Default: 0.0
  248. attn_drop (float, optional): Attention dropout rate. Default: 0.0
  249. drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
  250. norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
  251. downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
  252. use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
  253. """
  254. def __init__(self,
  255. dim,
  256. input_resolution,
  257. depth,
  258. num_heads,
  259. window_size=7,
  260. mlp_ratio=4.,
  261. qkv_bias=True,
  262. qk_scale=None,
  263. drop=0.,
  264. attn_drop=0.,
  265. drop_path=0.,
  266. norm_layer=nn.LayerNorm,
  267. downsample=True,
  268. use_checkpoint=False,
  269. i_layer=None):
  270. super().__init__()
  271. self.window_size = window_size
  272. self.shift_size = [window_size[0] // 2,window_size[1] // 2,window_size[2] // 2]
  273. self.depth = depth
  274. self.use_checkpoint = use_checkpoint
  275. self.i_layer=i_layer
  276. # build blocks
  277. self.blocks = nn.ModuleList([
  278. SwinTransformerBlock(
  279. dim=dim,
  280. input_resolution=input_resolution,
  281. num_heads=num_heads,
  282. window_size=window_size,
  283. shift_size=[0,0,0] if (i % 2 == 0) else self.shift_size,
  284. mlp_ratio=mlp_ratio,
  285. qkv_bias=qkv_bias,
  286. qk_scale=qk_scale,
  287. drop=drop,
  288. attn_drop=attn_drop,
  289. drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path, norm_layer=norm_layer)
  290. for i in range(depth)])
  291. # patch merging layer
  292. if downsample is not None:
  293. if i_layer==1 or i_layer==2:
  294. self.downsample = downsample(dim=dim, norm_layer=norm_layer,tag=1)
  295. else:
  296. self.downsample = downsample(dim=dim, norm_layer=norm_layer,tag=0)
  297. else:
  298. self.downsample = None
  299. def forward(self, x, S, H, W):
  300. # calculate attention mask for SW-MSA
  301. Sp = int(np.ceil(S / self.window_size[0])) * self.window_size[0]
  302. Hp = int(np.ceil(H / self.window_size[1])) * self.window_size[1]
  303. Wp = int(np.ceil(W / self.window_size[2])) * self.window_size[2]
  304. img_mask = torch.zeros((1, Sp, Hp, Wp, 1), device=x.device)
  305. s_slices = (slice(0, -self.window_size[0]),
  306. slice(-self.window_size[0], -self.shift_size[0]),
  307. slice(-self.shift_size[0], None))
  308. h_slices = (slice(0, -self.window_size[1]),
  309. slice(-self.window_size[1], -self.shift_size[1]),
  310. slice(-self.shift_size[1], None))
  311. w_slices = (slice(0, -self.window_size[2]),
  312. slice(-self.window_size[2], -self.shift_size[2]),
  313. slice(-self.shift_size[2], None))
  314. cnt = 0
  315. for s in s_slices:
  316. for h in h_slices:
  317. for w in w_slices:
  318. img_mask[:, s, h, w, :] = cnt
  319. cnt += 1
  320. mask_windows = window_partition(img_mask, self.window_size)
  321. mask_windows = mask_windows.view(-1,
  322. self.window_size[0] * self.window_size[1] * self.window_size[2])
  323. attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
  324. attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
  325. for blk in self.blocks:
  326. blk.H, blk.W = H, W
  327. if self.use_checkpoint:
  328. x = checkpoint.checkpoint(blk, x, attn_mask)
  329. else:
  330. x = blk(x, attn_mask)
  331. if self.downsample is not None:
  332. x_down = self.downsample(x, S, H, W)
  333. if self.i_layer!=1 and self.i_layer!=2:
  334. Ws, Wh, Ww = S , (H + 1) // 2, (W + 1) // 2
  335. else:
  336. Ws, Wh, Ww = S//2 , (H + 1) // 2, (W + 1) // 2
  337. return x, S, H, W, x_down, Ws, Wh, Ww
  338. else:
  339. return x, S, H, W, x, S, H, W
  340. class BasicLayer_up(nn.Module):
  341. """ A basic Swin Transformer layer for one stage.
  342. Args:
  343. dim (int): Number of feature channels
  344. depth (int): Depths of this stage.
  345. num_heads (int): Number of attention head.
  346. window_size (int): Local window size. Default: 7.
  347. mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
  348. qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
  349. qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
  350. drop (float, optional): Dropout rate. Default: 0.0
  351. attn_drop (float, optional): Attention dropout rate. Default: 0.0
  352. drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
  353. norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
  354. downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
  355. use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
  356. """
  357. def __init__(self,
  358. dim,
  359. input_resolution,
  360. depth,
  361. num_heads,
  362. window_size=7,
  363. mlp_ratio=4.,
  364. qkv_bias=True,
  365. qk_scale=None,
  366. drop=0.,
  367. attn_drop=0.,
  368. drop_path=0.,
  369. norm_layer=nn.LayerNorm,
  370. upsample=True,
  371. i_layer=None
  372. ):
  373. super().__init__()
  374. self.window_size = window_size
  375. self.shift_size = [window_size[0] // 2,window_size[1] // 2,window_size[2] // 2]
  376. self.depth = depth
  377. # build blocks
  378. self.blocks = nn.ModuleList([
  379. SwinTransformerBlock(
  380. dim=dim,
  381. input_resolution=input_resolution,
  382. num_heads=num_heads,
  383. window_size=window_size,
  384. shift_size=[0,0,0] if (i % 2 == 0) else self.shift_size,
  385. mlp_ratio=mlp_ratio,
  386. qkv_bias=qkv_bias,
  387. qk_scale=qk_scale,
  388. drop=drop,
  389. attn_drop=attn_drop,
  390. drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path, norm_layer=norm_layer)
  391. for i in range(depth)])
  392. # patch merging layer
  393. self.i_layer=i_layer
  394. if i_layer==1:
  395. self.Upsample = upsample(dim=2*dim, norm_layer=norm_layer,tag=1)
  396. elif i_layer==0:
  397. self.Upsample = upsample(dim=2*dim, norm_layer=norm_layer,tag=2)
  398. else:
  399. self.Upsample = upsample(dim=2*dim, norm_layer=norm_layer,tag=0)
  400. def forward(self, x,skip, S, H, W):
  401. """ Forward function.
  402. Args:
  403. x: Input feature, tensor size (B, H*W, C).
  404. H, W: Spatial resolution of the input feature.
  405. """
  406. skip = skip.flatten(2).transpose(1, 2)
  407. x_up = self.Upsample(x, S, H, W)
  408. x_up+=skip
  409. if self.i_layer==1:
  410. S, H, W = S * 2 , H * 2, W * 2
  411. elif self.i_layer==0:
  412. S, H, W = (S * 2)+1 , H * 2, W * 2
  413. else:
  414. S, H, W = S , H * 2, W * 2
  415. # calculate attention mask for SW-MSA
  416. Sp = int(np.ceil(S / self.window_size[0])) * self.window_size[0]
  417. Hp = int(np.ceil(H / self.window_size[1])) * self.window_size[1]
  418. Wp = int(np.ceil(W / self.window_size[2])) * self.window_size[2]
  419. img_mask = torch.zeros((1, Sp, Hp, Wp, 1), device=x.device)
  420. s_slices = (slice(0, -self.window_size[0]),
  421. slice(-self.window_size[0], -self.shift_size[0]),
  422. slice(-self.shift_size[0], None))
  423. h_slices = (slice(0, -self.window_size[1]),
  424. slice(-self.window_size[1], -self.shift_size[1]),
  425. slice(-self.shift_size[1], None))
  426. w_slices = (slice(0, -self.window_size[2]),
  427. slice(-self.window_size[2], -self.shift_size[2]),
  428. slice(-self.shift_size[2], None))
  429. cnt = 0
  430. for s in s_slices:
  431. for h in h_slices:
  432. for w in w_slices:
  433. img_mask[:, s, h, w, :] = cnt
  434. cnt += 1
  435. mask_windows = window_partition(img_mask, self.window_size)
  436. mask_windows = mask_windows.view(-1,
  437. self.window_size[0] * self.window_size[1] * self.window_size[2])
  438. attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
  439. attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
  440. for blk in self.blocks:
  441. x_up = blk(x_up, attn_mask)
  442. return x_up, S, H, W
  443. # done
  444. class project(nn.Module):
  445. def __init__(self,in_dim,out_dim,stride,padding,activate,norm,last=False):
  446. super().__init__()
  447. self.out_dim=out_dim
  448. self.conv1=nn.Conv3d(in_dim,out_dim,kernel_size=3,stride=stride,padding=padding)
  449. self.conv2=nn.Conv3d(out_dim,out_dim,kernel_size=3,stride=1,padding=1)
  450. self.activate=activate()
  451. self.norm1=norm(out_dim)
  452. self.last=last
  453. if not last:
  454. self.norm2=norm(out_dim)
  455. def forward(self,x):
  456. x=self.conv1(x)
  457. x=self.activate(x)
  458. #norm1
  459. Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
  460. x = x.flatten(2).transpose(1, 2)
  461. x = self.norm1(x)
  462. x = x.transpose(1, 2).view(-1, self.out_dim, Ws, Wh, Ww)
  463. x=self.conv2(x)
  464. if not self.last:
  465. x=self.activate(x)
  466. #norm2
  467. Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
  468. x = x.flatten(2).transpose(1, 2)
  469. x = self.norm2(x)
  470. x = x.transpose(1, 2).view(-1, self.out_dim, Ws, Wh, Ww)
  471. return x
  472. class PatchEmbed(nn.Module):
  473. """ Image to Patch Embedding
  474. Args:
  475. patch_size (int): Patch token size. Default: 4.
  476. in_chans (int): Number of input image channels. Default: 3.
  477. embed_dim (int): Number of linear projection output channels. Default: 96.
  478. norm_layer (nn.Module, optional): Normalization layer. Default: None
  479. """
  480. def __init__(self, patch_size=4, in_chans=4, embed_dim=96, norm_layer=None):
  481. super().__init__()
  482. patch_size = to_3tuple(patch_size)
  483. self.patch_size = patch_size
  484. self.in_chans = in_chans
  485. self.embed_dim = embed_dim
  486. self.proj1 = project(in_chans,embed_dim//2,[1,2,2],1,nn.GELU,nn.LayerNorm,False)
  487. self.proj2 = project(embed_dim//2,embed_dim,[1,2,2],1,nn.GELU,nn.LayerNorm,True)
  488. if norm_layer is not None:
  489. self.norm = norm_layer(embed_dim)
  490. else:
  491. self.norm = None
  492. def forward(self, x):
  493. """Forward function."""
  494. # padding
  495. _, _, S, H, W = x.size()
  496. if W % self.patch_size[2] != 0:
  497. x = F.pad(x, (0, self.patch_size[2] - W % self.patch_size[2]))
  498. if H % self.patch_size[1] != 0:
  499. x = F.pad(x, (0, 0, 0, self.patch_size[1] - H % self.patch_size[1]))
  500. if S % self.patch_size[0] != 0:
  501. x = F.pad(x, (0, 0, 0, 0, 0, self.patch_size[0] - S % self.patch_size[0]))
  502. x = self.proj1(x)
  503. x = self.proj2(x)
  504. if self.norm is not None:
  505. Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
  506. x = x.flatten(2).transpose(1, 2)
  507. x = self.norm(x)
  508. x = x.transpose(1, 2).view(-1, self.embed_dim, Ws, Wh, Ww)
  509. return x
  510. class SwinTransformer(nn.Module):
  511. """ Swin Transformer backbone.
  512. A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
  513. https://arxiv.org/pdf/2103.14030
  514. Args:
  515. pretrain_img_size (int): Input image size for training the pretrained model,
  516. used in absolute postion embedding. Default 224.
  517. patch_size (int | tuple(int)): Patch size. Default: 4.
  518. in_chans (int): Number of input image channels. Default: 3.
  519. embed_dim (int): Number of linear projection output channels. Default: 96.
  520. depths (tuple[int]): Depths of each Swin Transformer stage.
  521. num_heads (tuple[int]): Number of attention head of each stage.
  522. window_size (int): Window size. Default: 7.
  523. mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
  524. qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
  525. qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
  526. drop_rate (float): Dropout rate.
  527. attn_drop_rate (float): Attention dropout rate. Default: 0.
  528. drop_path_rate (float): Stochastic depth rate. Default: 0.2.
  529. norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
  530. ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
  531. patch_norm (bool): If True, add normalization after patch embedding. Default: True.
  532. out_indices (Sequence[int]): Output from which stages.
  533. frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
  534. -1 means not freezing any parameters.
  535. use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
  536. """
  537. def __init__(self,
  538. pretrain_img_size=224,
  539. patch_size=4,
  540. in_chans=1 ,
  541. embed_dim=96,
  542. depths=[2, 2, 2, 2],
  543. num_heads=[4, 8, 16, 32],
  544. window_size=7,
  545. mlp_ratio=4.,
  546. qkv_bias=True,
  547. qk_scale=None,
  548. drop_rate=0.,
  549. attn_drop_rate=0.,
  550. drop_path_rate=0.2,
  551. norm_layer=nn.LayerNorm,
  552. ape=False,
  553. patch_norm=True,
  554. out_indices=(0, 1, 2, 3),
  555. frozen_stages=-1,
  556. use_checkpoint=False):
  557. super().__init__()
  558. self.pretrain_img_size = pretrain_img_size
  559. self.num_layers = len(depths)
  560. self.embed_dim = embed_dim
  561. self.ape = ape
  562. self.patch_norm = patch_norm
  563. self.out_indices = out_indices
  564. self.frozen_stages = frozen_stages
  565. # split image into non-overlapping patches
  566. self.patch_embed = PatchEmbed(
  567. patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim,
  568. norm_layer=norm_layer if self.patch_norm else None)
  569. # absolute position embedding
  570. if self.ape:
  571. pretrain_img_size = to_3tuple(pretrain_img_size)
  572. patch_size = to_3tuple(patch_size)
  573. patches_resolution = [pretrain_img_size[0] // patch_size[0], pretrain_img_size[1] // patch_size[1],
  574. pretrain_img_size[2] // patch_size[2]]
  575. self.absolute_pos_embed = nn.Parameter(
  576. torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1], patches_resolution[2]))
  577. trunc_normal_(self.absolute_pos_embed, std=.02)
  578. self.pos_drop = nn.Dropout(p=drop_rate)
  579. # stochastic depth
  580. dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
  581. down_size=[[1,4,4],[1,8,8],[2,16,16],[4,32,32]]
  582. # build layers
  583. self.layers = nn.ModuleList()
  584. for i_layer in range(self.num_layers):
  585. layer = BasicLayer(
  586. dim=int(embed_dim * 2 ** i_layer),
  587. input_resolution=(
  588. pretrain_img_size[0] // down_size[i_layer][0], pretrain_img_size[1] // down_size[i_layer][1],
  589. pretrain_img_size[2] // down_size[i_layer][2]),
  590. depth=depths[i_layer],
  591. num_heads=num_heads[i_layer],
  592. window_size=window_size,
  593. mlp_ratio=mlp_ratio,
  594. qkv_bias=qkv_bias,
  595. qk_scale=qk_scale,
  596. drop=drop_rate,
  597. attn_drop=attn_drop_rate,
  598. drop_path=dpr[sum(
  599. depths[:i_layer]):sum(depths[:i_layer + 1])],
  600. norm_layer=norm_layer,
  601. downsample=PatchMerging,
  602. use_checkpoint=use_checkpoint,
  603. i_layer=i_layer)
  604. self.layers.append(layer)
  605. num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
  606. self.num_features = num_features
  607. # add a norm layer for each output
  608. for i_layer in out_indices:
  609. layer = norm_layer(num_features[i_layer])
  610. layer_name = f'norm{i_layer}'
  611. self.add_module(layer_name, layer)
  612. self._freeze_stages()
  613. def _freeze_stages(self):
  614. if self.frozen_stages >= 0:
  615. self.patch_embed.eval()
  616. for param in self.patch_embed.parameters():
  617. param.requires_grad = False
  618. if self.frozen_stages >= 1 and self.ape:
  619. self.absolute_pos_embed.requires_grad = False
  620. if self.frozen_stages >= 2:
  621. self.pos_drop.eval()
  622. for i in range(0, self.frozen_stages - 1):
  623. m = self.layers[i]
  624. m.eval()
  625. for param in m.parameters():
  626. param.requires_grad = False
  627. def forward(self, x):
  628. """Forward function."""
  629. x = self.patch_embed(x)
  630. down=[]
  631. Ws, Wh, Ww = x.size(2), x.size(3), x.size(4)
  632. if self.ape:
  633. # interpolate the position embedding to the corresponding size
  634. absolute_pos_embed = F.interpolate(self.absolute_pos_embed, size=(Ws, Wh, Ww), align_corners=True,
  635. mode='trilinear')
  636. x = (x + absolute_pos_embed).flatten(2).transpose(1, 2)
  637. else:
  638. x = x.flatten(2).transpose(1, 2)
  639. x = self.pos_drop(x)
  640. for i in range(self.num_layers):
  641. layer = self.layers[i]
  642. x_out, S, H, W, x, Ws, Wh, Ww = layer(x, Ws, Wh, Ww)
  643. if i in self.out_indices:
  644. norm_layer = getattr(self, f'norm{i}')
  645. x_out = norm_layer(x_out)
  646. out = x_out.view(-1, S, H, W, self.num_features[i]).permute(0, 4, 1, 2, 3).contiguous()
  647. down.append(out)
  648. return down
  649. def train(self, mode=True):
  650. """Convert the model into training mode while keep layers freezed."""
  651. super(SwinTransformer, self).train(mode)
  652. self._freeze_stages()
  653. class encoder(nn.Module):
  654. def __init__(self,
  655. pretrain_img_size,
  656. embed_dim,
  657. patch_size=4,
  658. depths=[2,2,2],
  659. num_heads=[24,12,6],
  660. window_size=4,
  661. mlp_ratio=4.,
  662. qkv_bias=True,
  663. qk_scale=None,
  664. drop_rate=0.,
  665. attn_drop_rate=0.,
  666. drop_path_rate=0.2,
  667. norm_layer=nn.LayerNorm
  668. ):
  669. super().__init__()
  670. self.num_layers = len(depths)
  671. self.pos_drop = nn.Dropout(p=drop_rate)
  672. # stochastic depth
  673. dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
  674. up_size=[[2,16,16],[1,8,8],[1,4,4]]
  675. # build layers
  676. self.layers = nn.ModuleList()
  677. for i_layer in range(self.num_layers)[::-1]:
  678. layer = BasicLayer_up(
  679. dim=int(embed_dim * 2 ** (len(depths)-i_layer-1)),
  680. input_resolution=(
  681. pretrain_img_size[0] // up_size[i_layer][0], pretrain_img_size[1] // up_size[i_layer][1],
  682. pretrain_img_size[2] // up_size[i_layer][2]),
  683. depth=depths[i_layer],
  684. num_heads=num_heads[i_layer],
  685. window_size=window_size,
  686. mlp_ratio=mlp_ratio,
  687. qkv_bias=qkv_bias,
  688. qk_scale=qk_scale,
  689. drop=drop_rate,
  690. attn_drop=attn_drop_rate,
  691. drop_path=dpr[sum(
  692. depths[:i_layer]):sum(depths[:i_layer + 1])],
  693. norm_layer=norm_layer,
  694. upsample=Patch_Expanding,
  695. i_layer=i_layer
  696. )
  697. self.layers.append(layer)
  698. self.num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
  699. def forward(self,x,skips):
  700. outs=[]
  701. S, H, W = x.size(2), x.size(3), x.size(4)
  702. x = x.flatten(2).transpose(1, 2)
  703. x = self.pos_drop(x)
  704. for i in range(self.num_layers)[::-1]:
  705. layer = self.layers[i]
  706. x, S, H, W, = layer(x,skips[i], S, H, W)
  707. out = x.view(-1, S, H, W, self.num_features[i])
  708. outs.append(out)
  709. return outs
  710. class final_patch_expanding(nn.Module):
  711. def __init__(self,dim,num_class,patch_size):
  712. super().__init__()
  713. self.up=nn.ConvTranspose3d(dim,num_class,patch_size,patch_size)
  714. def forward(self,x):
  715. x=x.permute(0,4,1,2,3)
  716. x=self.up(x)
  717. return x
  718. class swintransformer(SegmentationNetwork):
  719. def __init__(self, input_channels, base_num_features, num_classes, num_pool, num_conv_per_stage=2,
  720. feat_map_mul_on_downscale=2, conv_op=nn.Conv2d,
  721. norm_op=nn.BatchNorm2d, norm_op_kwargs=None,
  722. dropout_op=nn.Dropout2d, dropout_op_kwargs=None,
  723. nonlin=nn.LeakyReLU, nonlin_kwargs=None, deep_supervision=True, dropout_in_localization=False,
  724. final_nonlin=softmax_helper, weightInitializer=InitWeights_He(1e-2), pool_op_kernel_sizes=None,
  725. conv_kernel_sizes=None,
  726. upscale_logits=False, convolutional_pooling=False, convolutional_upsampling=False,
  727. max_num_features=None, basic_block=None,
  728. seg_output_use_bias=False):
  729. super(swintransformer, self).__init__()
  730. self._deep_supervision = deep_supervision
  731. self.do_ds = deep_supervision
  732. self.num_classes=num_classes
  733. self.conv_op=conv_op
  734. self.upscale_logits_ops = []
  735. self.upscale_logits_ops.append(lambda x: x)
  736. embed_dim=96
  737. depths=[2, 2, 2, 2]
  738. num_heads=[3, 6, 12, 24]
  739. patch_size=[1,4,4]
  740. 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)
  741. 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])
  742. self.final=[]
  743. for i in range(len(depths)-1):
  744. self.final.append(final_patch_expanding(embed_dim*2**i,self.num_classes,patch_size=patch_size))
  745. self.final=nn.ModuleList(self.final)
  746. def forward(self, x):
  747. seg_outputs=[]
  748. skips = self.model_down(x)
  749. neck=skips[-1]
  750. out=self.encoder(neck,skips)
  751. for i in range(len(out)):
  752. seg_outputs.append(self.final[-(i+1)](out[i]))
  753. if self._deep_supervision and self.do_ds:
  754. return tuple([seg_outputs[-1]] + [i(j) for i, j in
  755. zip(list(self.upscale_logits_ops)[::-1], seg_outputs[:-1][::-1])])
  756. else:
  757. return seg_outputs[-1]
  758. @staticmethod
  759. def compute_approx_vram_consumption(patch_size, num_pool_per_axis, base_num_features, max_num_features,
  760. num_modalities, num_classes, pool_op_kernel_sizes, deep_supervision=False,
  761. conv_per_stage=2):
  762. """
  763. This only applies for num_conv_per_stage and convolutional_upsampling=True
  764. not real vram consumption. just a constant term to which the vram consumption will be approx proportional
  765. (+ offset for parameter storage)
  766. :param deep_supervision:
  767. :param patch_size:
  768. :param num_pool_per_axis:
  769. :param base_num_features:
  770. :param max_num_features:
  771. :param num_modalities:
  772. :param num_classes:
  773. :param pool_op_kernel_sizes:
  774. :return:
  775. """
  776. if not isinstance(num_pool_per_axis, np.ndarray):
  777. num_pool_per_axis = np.array(num_pool_per_axis)
  778. npool = len(pool_op_kernel_sizes)
  779. map_size = np.array(patch_size)
  780. tmp = np.int64((conv_per_stage * 2 + 1) * np.prod(map_size, dtype=np.int64) * base_num_features +
  781. num_modalities * np.prod(map_size, dtype=np.int64) +
  782. num_classes * np.prod(map_size, dtype=np.int64))
  783. num_feat = base_num_features
  784. for p in range(npool):
  785. for pi in range(len(num_pool_per_axis)):
  786. map_size[pi] /= pool_op_kernel_sizes[p][pi]
  787. num_feat = min(num_feat * 2, max_num_features)
  788. 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
  789. tmp += num_blocks * np.prod(map_size, dtype=np.int64) * num_feat
  790. if deep_supervision and p < (npool - 2):
  791. tmp += np.prod(map_size, dtype=np.int64) * num_classes
  792. # print(p, map_size, num_feat, tmp)
  793. return tmp

Swin_Unet_s_ACDC_2laterdown.py at commit 6357a1d, under MIT · at the source

Overview

Authors: Tsuyoshi Ueyama1,2, Erika Takahashi1, Naoto Fujita1, Yuichi Suzuki2, Shohei Inui3, Koichiro Yasaka3, Hideyuki Iwanaga2, Osamu Abe3, Yasuhiko Terada1
ORCID iDs: Koichiro Yasaka
  1. Institute of Pure and Applied Sciences, University of Tsukuba, Tsukuba, Ibaraki, Japan
  2. Department of Radiology, The University of Tokyo Hospital, Tokyo, Japan
  3. Department of Radiology, The University of Tokyo, Tokyo, Japan
Dates: received 8 October 2024; accepted 3 April 2026; published online 14 May 2026; in print May 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.2463/mrms.mp.2024-0149 · PMID 42128846 · PMCID PMC13500223 · OpenAlex W7160981048
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), stroke (population), clinical / translational (subfield)
Methods: Connectivity, Statistics, Machine learning, fMRI & imaging, Physiology & signal measures
Keywords: deep learning, diffusion-weighted image, distortion correction, image registration
MeSH: Diffusion Magnetic Resonance Imaging*, Image Interpretation, Computer-Assisted*, Image Processing, Computer-Assisted*, Brain, Brain Neoplasms, Cerebrovascular Disorders, Convolutional Neural Networks, Diffusion Tensor Imaging, Female, Humans, Male (* major topic)
Journal subjects: Major Paper
Topic: Advanced MRI Techniques and Applications (Radiology, Nuclear Medicine and Imaging, Medicine), according to OpenAlex
Funding: Japan Society for the Promotion of Science (JP24K00891)
Citations: not cited yet (Europe PMC); 26 references in the paper

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

License: none: the authors keep all their rights
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: d813f9435cffb5b663f15b52d55ed0a27181f0b2, 31 July 2025
Languages: MATLAB (6), Python (6), Shell (5)
Size: 61 files, 17 scripts
Software Heritage: not archived
Found in: the text, “Training details”
Holds: README, environment (Singularity)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: FSL (3 files), NiBabel (3 files), NumPy (3 files), PyTorch (3 files), Tools for NIfTI and ANALYZE image (MATLAB) (2 files), ANTs (1 file), FreeSurfer (1 file), Image Processing Toolbox (1 file), SciPy (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
18 files

junyuchen245/transmorph_transformer_for_medical_image_registration

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 6357a1d7fc44c36db9b1d1ccaa372409253142cf, 22 May 2025
Languages: Python (359), Shell (8)
Size: 463 files, 367 scripts
Software Heritage: not archived
Found in: the text, “Training details”
Holds: README, license file, environment (requirements.txt, Docker/TransMorph_build_Docker/Dockerfile, Docker/TransMorph_build_Docker/requirements.txt)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (252 files), NumPy (243 files), Matplotlib (84 files), SciPy (63 files), scikit-image (23 files), nnU-Net (14 files), NiBabel (13 files), ANTs (5 files), scikit-learn (4 files), Pillow (2 files), pandas (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
369 files

Tracing map

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

What the map holds:

  • 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://doi.org/10.2463/mrms.mp.2024-0149

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/mrms.mp.2024-0149},
url = {https://doi.org/10.2463/mrms.mp.2024-0149},
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/05/14
VL - 25
IS - 2
SP - 2024
EP - 0149
SN - 1347-3182
PB - Japanese Society for Magnetic Resonance in Medicine
DO - 10.2463/mrms.mp.2024-0149
UR - https://doi.org/10.2463/mrms.mp.2024-0149
LA - en
ER -

CSL-JSON

{
"id": "10.2463/mrms.mp.2024-0149",
"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": "Magn Reson Med Sci",
"volume": "25",
"issue": "2",
"page": "2024-0149",
"DOI": "10.2463/mrms.mp.2024-0149",
"PMID": "42128846",
"PMCID": "PMC13500223",
"ISSN": "1347-3182",
"publisher": "Japanese Society for Magnetic Resonance in Medicine",
"URL": "https://doi.org/10.2463/mrms.mp.2024-0149",
"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 Association
In 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 imaging
In 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: Epilepsia
In 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 biology
In 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 communications
In 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. Medicine
In 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: eLife
In 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: Nature
In 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.

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.