OSCR

TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation.

Code ↔ Paper

7 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 7 matches
  1. [1] § Materials and methods ↔ code/custom/architectures/tokenunet.py, lines 415–529 · score 0.78 · skip connections, MLP Mixer, Transformer encoder, token processing, TokenUNet, bottlenecks
  2. [2] § Materials and methods › Architectures and modules › TokenLearner and TokenFuser. ↔ code/custom/architectures/tokenunet.py, lines 209–268 · score 0.76 · Linear layer, spatial attention masks, feature map, logit, sigmoid, soft
  3. [3] § Materials and methods › Training ↔ code/scripts/profiling_archs.py, lines 333–413 · score 0.74 · deep supervision, Nesterov, mirroring, momentum, nnUNet, SGD
  4. [4] § Materials and methods › Training ↔ code/custom/nnUNetTokenUNetTrainer.py, lines 175–236 · score 0.65 · deep supervision, nnUNet, schedule, SGD, augmented, batch
  5. [5] § Materials and methods › Architectures and modules › UNet. ↔ code/custom/architectures/tokenunet.py, lines 53–84 · score 0.64 · LeakyReLU, trilinear, CONV, nonlinear, volumes, upsampling
  6. [6] § Results › What TokenLearner learns ↔ code/scripts/dice_attention.py, lines 1–24 · score 0.62 · Dice score, spatial attention maps, TokenFuser, TokenLearner, alignment, fold
  7. [7] § Materials and methods › Data ↔ code/scripts/convert_fets_to_nnunet.py, lines 60–85 · score 0.57 · FeTS, BraTS, FLAIR, modalities

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 · 529 lines · 19 KB · no license · 3 matches

  1. import torch
  2. import torch.nn as nn
  3. import math
  4. from .tokenmixers import *
  5. __all__ = [
  6. "Block",
  7. "BlockUp",
  8. "Stage",
  9. "StageUp",
  10. "CNNEnc",
  11. "CNNDec",
  12. "SpatialAttentionMaskMaker3d",
  13. "TokenLearner3d",
  14. "TokenFuser3d",
  15. "TokenUNet"
  16. ]
  17. class Block(nn.Module):
  18. def __init__(
  19. self,
  20. in_channels,
  21. out_channels,
  22. kernel_size,
  23. p=0.0,
  24. downsample=False
  25. ):
  26. super().__init__()
  27. self.block = nn.Sequential(
  28. nn.InstanceNorm3d(in_channels),
  29. nn.Conv3d(in_channels, in_channels*2, kernel_size, stride=2 if downsample else 1, padding=1),
  30. nn.LeakyReLU(0.1),
  31. nn.InstanceNorm3d(in_channels*2),
  32. nn.Conv3d(in_channels*2, out_channels, kernel_size=1,padding="same"),
  33. )
  34. nn.init.kaiming_normal_(self.block[1].weight, mode='fan_in', nonlinearity='relu')
  35. nn.init.kaiming_normal_(self.block[-1].weight, mode='fan_in', nonlinearity='relu')
  36. self.res_alpha = nn.Parameter(torch.tensor([-2.5]))
  37. adjust_volume = nn.AvgPool3d(2,2) if downsample else nn.Identity()
  38. adjust_channels = nn.Conv3d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity()
  39. self.adjust = nn.Sequential(adjust_volume, adjust_channels)
  40. def forward(self, x):
  41. res = self.block(x)
  42. alpha = torch.sigmoid(self.res_alpha)
  43. x = self.adjust(x)
  44. x = (1.-alpha)*x + alpha*res
  45. return x
  46. class BlockUp(nn.Module):
  47. def __init__(
  48. self,
  49. in_channels,
  50. out_channels,
  51. kernel_size,
  52. p=0.0,
  53. upsample=False
  54. ):
  55. super().__init__()
  56. self.block = nn.Sequential(
  57. nn.InstanceNorm3d(in_channels),
  58. nn.Conv3d(in_channels, in_channels*2, kernel_size, padding=1 ),
  59. nn.LeakyReLU(0.1),
  60. nn.InstanceNorm3d(in_channels*2),
  61. nn.ConvTranspose3d(in_channels*2, out_channels, kernel_size=2, stride=2) if upsample else nn.Conv3d(in_channels*2, out_channels, kernel_size=1, stride=1),
  62. )
  63. nn.init.kaiming_normal_(self.block[1].weight, mode='fan_in', nonlinearity='relu')
  64. nn.init.kaiming_normal_(self.block[-1].weight, mode='fan_in', nonlinearity='relu')
  65. self.res_alpha = nn.Parameter(torch.tensor([-2.5]))
  66. adjust_volume = nn.Upsample(scale_factor=2, mode="trilinear") if upsample else nn.Identity()
  67. adjust_channels = nn.Conv3d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity()
  68. self.adjust = nn.Sequential(adjust_volume, adjust_channels)
  69. def forward(self, x):
  70. res = self.block(x)
  71. alpha = torch.sigmoid(self.res_alpha)
  72. x = self.adjust(x)
  73. x = (1.-alpha)*x + alpha*res
  74. return x
  75. class Stage(nn.Module):
  76. def __init__(
  77. self,
  78. in_channels,
  79. out_channels,
  80. kernel_size,
  81. n_blocks,
  82. downsample=True
  83. ):
  84. super().__init__()
  85. self.blocks = nn.ModuleList(
  86. [Block(in_channels, in_channels, kernel_size) for _ in range(n_blocks-1)]+
  87. [Block(in_channels, out_channels, kernel_size, downsample=downsample)]
  88. )
  89. def forward(self, x):
  90. for block in self.blocks:
  91. x = block(x)
  92. return x
  93. class StageUp(nn.Module):
  94. def __init__(
  95. self,
  96. in_channels,
  97. out_channels,
  98. kernel_size,
  99. n_blocks,
  100. upsample=True
  101. ):
  102. super().__init__()
  103. self.blocks = nn.ModuleList(
  104. [BlockUp(in_channels, in_channels, kernel_size) for _ in range(n_blocks-1)]+
  105. [BlockUp(in_channels, out_channels, kernel_size, upsample=upsample)]
  106. )
  107. def forward(self, x):
  108. for block in self.blocks:
  109. x = block(x)
  110. return x
  111. class CNNEnc(nn.Module):
  112. def __init__(
  113. self,
  114. in_channels,
  115. stage_channels,
  116. kernel_size,
  117. blocks_per_stage
  118. ):
  119. super().__init__()
  120. n_stages = len(stage_channels)
  121. # Normalise blocks_per_stage: accept a single int or a per-stage list
  122. if isinstance(blocks_per_stage, int):
  123. blocks_per_stage = [blocks_per_stage] * n_stages
  124. assert len(blocks_per_stage) == n_stages, (
  125. f"blocks_per_stage length ({len(blocks_per_stage)}) "
  126. f"must match the number of stages ({n_stages})"
  127. )
  128. self.norm = nn.InstanceNorm3d(in_channels)
  129. # Build stages; the output of stage i is the input of stage i+1
  130. in_ch = in_channels
  131. self.stages = nn.ModuleList()
  132. first = 0
  133. for out_ch, n_blocks in zip(stage_channels, blocks_per_stage):
  134. self.stages.append(
  135. Stage(in_ch, out_ch, kernel_size, n_blocks, downsample=first>0)#last!=n_stages)
  136. )
  137. first += 1
  138. in_ch = out_ch
  139. def forward(self, x):
  140. x = self.norm(x)
  141. maps_to_decoder = []
  142. for stage in self.stages:
  143. x = stage(x)
  144. maps_to_decoder.append(x)
  145. #print(x.shape)
  146. return maps_to_decoder
  147. class CNNDec(nn.Module):
  148. def __init__(
  149. self,
  150. in_channels,
  151. stage_channels,
  152. kernel_size,
  153. blocks_per_stage
  154. ):
  155. super().__init__()
  156. n_stages = len(stage_channels)
  157. # Normalise blocks_per_stage: accept a single int or a per-stage list
  158. if isinstance(blocks_per_stage, int):
  159. blocks_per_stage = [blocks_per_stage] * n_stages
  160. assert len(blocks_per_stage) == n_stages, (
  161. f"blocks_per_stage length ({len(blocks_per_stage)}) "
  162. f"must match the number of stages ({n_stages})"
  163. )
  164. #self.norm = nn.InstanceNorm3d(in_channels)
  165. # Build stages; the output of stage i is the input of stage i+1
  166. in_ch = in_channels
  167. self.stages = nn.ModuleList()
  168. #last = 0
  169. for out_ch, n_blocks in zip(stage_channels, blocks_per_stage):
  170. self.stages.append(
  171. StageUp(in_ch, out_ch, kernel_size, n_blocks, upsample=True)
  172. )
  173. in_ch = out_ch
  174. #last += 1
  175. def forward(self, maps_to_decoder):
  176. x = 0.0
  177. for idx,stage in enumerate(self.stages):
  178. x = stage(maps_to_decoder[idx]+x)
  179. #print(maps_to_decoder[idx].shape)
  180. return x
  181. class SpatialAttentionMaskMaker3d(nn.Module):
  182. """
  183. For each spatial location in the feature map, learns a direct linear mapping
  184. from the local feature vector to N token assignment scores (one per token).
  185. A sigmoid then turns each score into an independent soft weight in (0, 1).
  186. Input: feature map of shape (B, C, H, W, D)
  187. Output: attention masks of shape (B, n_tokens, H, W, D), values in (0, 1)
  188. """
  189. def __init__(self, in_channels: int, n_tokens: int = 8, bias: bool = True):
  190. super().__init__()
  191. self.in_channels = in_channels
  192. self.n_tokens = n_tokens
  193. #self.norm = nn.LayerNorm(in_channels, elementwise_affine=False)
  194. # Single linear layer: maps each voxel's C-dim feature vector directly
  195. # to n_tokens scalar logits. Acts like a learned 1x1x1 convolution.
  196. # (B, n_voxels, C) -> (B, n_voxels, n_tokens)
  197. self.linear = nn.Linear(in_features=in_channels, out_features=n_tokens, bias=bias)
  198. def norm(self, x):
  199. return torch.nn.functional.normalize(x, p=2, dim=-1) * math.sqrt(x.shape[-1])
  200. def forward(self, feature_map_bchwd: torch.Tensor):
  201. """
  202. Args:
  203. feature_map_bchwd: (B, C, H, W, D)
  204. Returns:
  205. attention_masks_bnhwd: (B, n_tokens, H, W, D)
  206. """
  207. B, C, H, W, D = feature_map_bchwd.shape
  208. # assert C == self.in_channels, f"Expected in_channels={self.in_channels}, got {C}"
  209. # Flatten spatial dims to treat all voxels uniformly.
  210. # (B, C, H, W, D) -> (B, C, n_voxels)
  211. n_voxels = H * W * D
  212. feature_map_bcv = feature_map_bchwd.reshape(B, C, n_voxels)
  213. # Move channels last so Linear operates on the feature dimension.
  214. # (B, C, n_voxels) -> (B, n_voxels, C)
  215. feature_map_bvc = feature_map_bcv.permute(0, 2, 1)
  216. feature_map_bvc = self.norm(feature_map_bvc)
  217. # Direct linear projection: each voxel's feature vector -> n_tokens scores.
  218. # (B, n_voxels, C) -> (B, n_voxels, n_tokens)
  219. token_scores_bvn = self.linear(feature_map_bvc)
  220. # Sigmoid: independent soft weight per (voxel, token) pair, no competition across voxels.
  221. # (B, n_voxels, n_tokens) -> same shape, values in (0, 1)
  222. attention_weights_bvn = torch.sigmoid(token_scores_bvn)
  223. # Restore spatial structure and move token dim before spatial dims.
  224. # (B, n_voxels, n_tokens) -> (B, n_tokens, n_voxels) -> (B, n_tokens, H, W, D)
  225. attention_masks_bnhwd = attention_weights_bvn.permute(0, 2, 1).reshape(
  226. B, self.n_tokens, H, W, D
  227. )
  228. return attention_masks_bnhwd
  229. class TokenLearner3d(nn.Module):
  230. """
  231. TokenLearner: compresses a volumetric feature map into a small set of learned tokens.
  232. The key idea:
  233. 1. SpatialAttentionMaskMaker produces N soft spatial masks (B, N, H, W, D)
  234. 2. Each mask is element-wise multiplied with the feature map (B, C, H, W, D)
  235. 3. The masked feature map is mean-pooled over (H, W, D) -> one C-dim token per mask
  236. 4. Result: set of N tokens, shape (B, N, C), a compact set-like representation
  237. This replaces the global average pool (which treats all voxels equally) with
  238. N learned, content-adaptive pooling operations.
  239. Args:
  240. in_channels: number of feature channels C of the input
  241. n_tokens: number of output tokens N (the compression factor)
  242. bias: whether Linear layers use bias
  243. out_channels: if set, project each token from C to out_channels via a Linear layer;
  244. otherwise tokens keep their original C channels (Identity)
  245. """
  246. def __init__(
  247. self,
  248. in_channels: int,
  249. n_tokens: int = 8,
  250. bias: bool = True,
  251. out_channels: int = None,
  252. ):
  253. super().__init__()
  254. self.in_channels = in_channels
  255. self.n_tokens = n_tokens
  256. #self.norm = nn.LayerNorm(in_channels, elementwise_affine=False)
  257. self.mask_maker = SpatialAttentionMaskMaker3d(
  258. in_channels=in_channels,
  259. n_tokens=n_tokens,
  260. bias=bias,
  261. )
  262. # Optional linear projection applied independently to each token vector
  263. self.token_projector = (
  264. nn.Linear(in_channels, out_channels) if out_channels else nn.Identity()
  265. )
  266. def norm(self, x):
  267. return torch.nn.functional.normalize(x, p=2, dim=-1) * math.sqrt(x.shape[-1])
  268. def forward(self, feature_map_bchwd: torch.Tensor):
  269. """
  270. Args:
  271. feature_map_bchwd: (B, C, H, W, D) — volumetric feature map
  272. Returns:
  273. tokens_bnc: (B, n_tokens, C_out) — the learned token set
  274. attention_masks_bnhwd: (B, n_tokens, H, W, D) — masks for inspection / auxiliary loss
  275. """
  276. B, C, H, W, D = feature_map_bchwd.shape
  277. n_voxels = H * W * D
  278. # --- Step 1: build one spatial attention mask per token ---
  279. # Each mask encodes which voxels are relevant for that token.
  280. # (B, C, H, W, D) -> (B, n_tokens, H, W, D)
  281. attention_masks_bnhwd = self.mask_maker(feature_map_bchwd)
  282. attention_masks_flat_bnv = attention_masks_bnhwd.reshape(B, self.n_tokens, n_voxels)
  283. # --- Step 2: soft-mask the feature map for each token ---
  284. # Expand feature map along the token dimension so we can broadcast the mask.
  285. # (B, C, H, W, D) -> (B, 1, C, H, W, D)
  286. # feature_map_b1chwd = feature_map_bchwd.unsqueeze(1)
  287. feature_map_bcv = feature_map_bchwd.reshape(B, C, n_voxels)
  288. # (B, n_tokens, H, W, D) -> (B, n_tokens, 1, H, W, D) for channel broadcast
  289. #attention_masks_bn1hwd = attention_masks_bnhwd.unsqueeze(2)
  290. # Element-wise product: each voxel's feature vector is scaled by its token weight.
  291. # (B, n_tokens, 1, H, W, D) * (B, 1, C, H, W, D) -> (B, n_tokens, C, H, W, D)
  292. #masked_features_bnchwd = attention_masks_bn1hwd * feature_map_b1chwd
  293. # --- Step 3: spatial mean-pooling -> one token vector per mask ---
  294. # Average over the three spatial dimensions (H, W, D).
  295. # (B, n_tokens, C, H, W, D) -> (B, n_tokens, C)
  296. #tokens_bnc = masked_features_bnchwd.mean(dim=(3, 4, 5))
  297. # --- True step 2-3 ---
  298. # We do not materialize the full (B, n_tokens, C, H, W, D) tensor!
  299. tokens_bnc = torch.bmm(
  300. attention_masks_flat_bnv / (attention_masks_flat_bnv.sum(dim=2, keepdims=True) + 1e-8),
  301. #self.norm(feature_map_bcv.transpose(1,2))
  302. feature_map_bcv.transpose(1,2)
  303. ) #/ n_voxels
  304. # --- Step 4 (optional): project token channels ---
  305. # (B, n_tokens, C) -> (B, n_tokens, C_out)
  306. tokens_bnc = self.token_projector(tokens_bnc)
  307. return tokens_bnc, attention_masks_bnhwd
  308. class TokenFuser3d(nn.Module):
  309. """
  310. This module weighted-averages N tokens and computes N spatial pertinence masks,
  311. then broadcasts the tokens over the masks to update feature maps.
  312. """
  313. def __init__(self,
  314. n_tokens,
  315. conv_dim,
  316. token_dim,
  317. bias=True
  318. ):
  319. super().__init__()
  320. self.M = nn.Linear(n_tokens, n_tokens)
  321. self.Beta = SpatialAttentionMaskMaker3d(
  322. in_channels=conv_dim,
  323. n_tokens=n_tokens,
  324. bias=bias,
  325. )
  326. self.n_tokens = n_tokens
  327. self.conv_dim = conv_dim
  328. self.token_dim = token_dim
  329. if (token_dim != conv_dim):
  330. self.C = nn.Linear(token_dim, conv_dim)
  331. else:
  332. self.C = nn.Identity()
  333. def forward(self, tokens, feat_maps):
  334. B, C, H, W, D = feat_maps.shape # B,Cc,H,[W,[D]]
  335. n_voxels = H * W * D
  336. # Mix the tokens, and eventually map them to channel dimension of convolutional features
  337. mixed_tokens = self.C(self.M(tokens.transpose(1,2)).transpose(1,2)) # B, N, Ct -> B, N, Cc
  338. # Determine how much each voxel needs each token
  339. pertinence_masks_bnhwd = self.Beta(feat_maps) # B,Cc,H,[W,[D]]
  340. pertinence_masks_flat_bnv = pertinence_masks_bnhwd.reshape(B, self.n_tokens, n_voxels) # B, N, V=HWD
  341. token_broadcast_bcv = torch.bmm(mixed_tokens.transpose(1,2), pertinence_masks_flat_bnv) # B, (Cc, N) @ (N, V) -> (B,Ct,V)
  342. token_broadcast_bchwd = token_broadcast_bcv.reshape(B,C,H,W,D)
  343. out_maps = feat_maps + token_broadcast_bchwd
  344. return out_maps
  345. class TokenUNet(nn.Module):
  346. def __init__(
  347. self,
  348. in_channels,
  349. enc_stage_channels,
  350. kernel_size,
  351. blocks_per_stage,
  352. num_classes,
  353. n_tokens=8,
  354. token_dim=None,
  355. tokenize=True,
  356. attention=False,
  357. process_tokens=True,
  358. token_blocks=2,
  359. bias=True
  360. ):
  361. super().__init__()
  362. self.tokenize = tokenize
  363. self.process_tokens = process_tokens
  364. # Ensure token_dim is set (defaults to the channel size of the bottleneck)
  365. bottleneck_channels = enc_stage_channels[-1]
  366. dec_stage_channels = enc_stage_channels[::-1][1:]
  367. self.token_dim = token_dim if token_dim is not None else bottleneck_channels
  368. # 1. Encoder
  369. self.encoder = CNNEnc(
  370. in_channels=in_channels,
  371. stage_channels=enc_stage_channels,
  372. kernel_size=kernel_size,
  373. blocks_per_stage=blocks_per_stage
  374. )
  375. # 2. Tokenizer (Optional)
  376. if self.tokenize:
  377. self.token_learner = TokenLearner3d(
  378. in_channels=bottleneck_channels,
  379. n_tokens=n_tokens,
  380. bias=bias,
  381. out_channels=self.token_dim
  382. )
  383. # 3. Token Processor (Transformer / MLP Mixer) - Optional
  384. if self.process_tokens:
  385. if attention:
  386. self.token_processor = MyTransformerEncoder(
  387. d_model=self.token_dim,
  388. nhead=4,
  389. dim_feedforward=self.token_dim*2,
  390. dropout=0.0,
  391. num_layers=token_blocks
  392. )
  393. else:
  394. self.token_processor = MyMLPMixer(
  395. n_tokens=n_tokens,
  396. d_model=self.token_dim,
  397. dim_feedforward=self.token_dim*2,
  398. n_blocks=token_blocks,
  399. dropout=0.0
  400. )
  401. else:
  402. self.token_processor = nn.Identity()
  403. # 4. Token Fuser
  404. self.token_fuser = TokenFuser3d(
  405. n_tokens=n_tokens,
  406. conv_dim=bottleneck_channels,
  407. token_dim=self.token_dim,
  408. bias=bias
  409. )
  410. # 5. Decoder
  411. self.decoder = CNNDec(
  412. in_channels=bottleneck_channels, # Starts from the bottleneck size
  413. stage_channels=dec_stage_channels,
  414. kernel_size=kernel_size,
  415. blocks_per_stage=[1,]*len(dec_stage_channels)
  416. )
  417. # 6. Segmentation Head
  418. self.segmentation_head = nn.Conv3d(
  419. in_channels=dec_stage_channels[-1],
  420. out_channels=num_classes,
  421. kernel_size=1
  422. )
  423. def forward(self, x):
  424. # Forward through Encoder
  425. maps_to_decoder = self.encoder(x)
  426. # The bottleneck is the lowest resolution feature map (the last one)
  427. bottleneck = maps_to_decoder[-1]
  428. if self.tokenize:
  429. # Extract Tokens
  430. tokens, attention_masks = self.token_learner(bottleneck)
  431. # Process Tokens (Transformer/Mixer)
  432. if self.process_tokens:
  433. tokens = self.token_processor(tokens)
  434. # Fuse Tokens back into the bottleneck feature map
  435. bottleneck = self.token_fuser(tokens, bottleneck)
  436. # Update the bottleneck in our skip-connection list
  437. maps_to_decoder[-1] = bottleneck
  438. # CRITICAL: Reverse the maps so they match the Decoder's expected ascending resolutions
  439. # Example: if Enc outputs [128x128, 64x64, 32x32], Decoder needs [32x32, 64x64, 128x128]
  440. maps_to_decoder = maps_to_decoder[::-1]
  441. # Forward through Decoder
  442. out = self.segmentation_head(self.decoder(maps_to_decoder))
  443. return out

tokenunet.py at commit 0bc16c2, no license · at the source

Overview

Authors: Louis Fabrice Tshimanga1,2,3, Andrea Zanola1,2, Federico Del Pup2,3, Manfredo Atzori1,2,3,4
  1. Department of Neuroscience, University of Padua, Padua, Italy
  2. Padova Neuroscience Center, University of Padua, Padua, Italy
  3. Department of Information Engineering, University of Padua, Padua, Italy
  4. Information Systems Institute, University of Applied Sciences Western Switzerland (HES-SO Valais), Sierre, Switzerland
Journal: PloS one, volume 21, issue 8, article e0354511
Dates: received 4 March 2026; accepted 9 July 2026; published online 5 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pone.0354511 · PMID 42555631 · PMCID PMC13440844 · OpenAlex W7172538079
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), other condition (population), methods / tools (subfield)
Methods: Connectivity, Machine learning
MeSH: Brain*, Imaging, Three-Dimensional*, Neuroimaging*, Algorithms, Convolutional Neural Networks, Humans, Image Processing, Computer-Assisted (* major topic)
Journal subjects: Medicine and Health Sciences, Oncology, Cancers and Neoplasms, Research and Analysis Methods, Mathematical and Statistical Techniques, Mathematical Functions, Convolution, Computer and Information Sciences, Computer Architecture, Artificial Intelligence, Machine Learning, Deep Learning, Neural Networks, Biology and Life Sciences, Neuroscience, Cognitive Science, Cognition, Memory, Memory Recall, Learning and Memory, Diagnostic Medicine, Diagnostic Radiology, Magnetic Resonance Imaging, Imaging Techniques, Radiology and Imaging
Topic: Advanced Neural Network Applications (Computer Vision and Pattern Recognition, Computer Science), according to OpenAlex
Citations: not cited yet (Europe PMC); 45 references in the paper

Abstract

We present TokenUNet, adopting the TokenLearner and TokenFuser modules to encase Transformers into UNets. While Transformers enable expressive global interactions among input elements in medical imaging, computational challenges hinder their deployment on common hardware. Models like (Swin)UNETR exemplify the integration of (Swin)Transformer encoders into UNets, tokenizing inputs into small subvolumes (83 voxels). The Transformer attention mechanism scales quadratically with the number of tokens, which is tied to the cubic scaling of 3D input resolution. This work reconsiders the role of convolution and attention, introducing TokenUNets, a family of 3D segmentation models better suited to constrained computational environments and time frames. To mitigate computational demands, our approach maintains the convolutional encoder of UNet-like models, and applies TokenLearner to 3D feature maps. This module pools a preset number of tokens from local and global structures, decoupling token number and input size. Our results on the BraTS challenge dataset for glioma segmentation show this tokenization effectively encodes task-relevant information, yielding naturally interpretable attention maps. The memory footprint, computation times at inference, and parameter counts of our heaviest model are reduced to 38%, 10%, and 17% of the SwinUNETR values, with statistically equivalent Dice score performance, for nnunetv2 5-fold cross-validation. This work opens the way to more efficient training in computationally restrained contexts, such as 3D medical imaging. Easing model optimization, fine-tuning, and transfer-learning in limited hardware settings can accelerate and diversify the development of approaches, for the benefit of the research community.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 7 matches between paragraphs and lines of code.

MedMaxLab/tokenunet

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 0bc16c2e27df7e6da0fed35952c2ed951f46e4e8, 17 August 2026
Languages: Python (16), Jupyter (1)
Size: 67 files, 17 scripts
Software Heritage: not archived
Found in: the text, “Training”
Holds: README, 1 notebook
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (9 files), PyTorch (8 files), Matplotlib (6 files), pandas (5 files), nnU-Net (4 files), NiBabel (3 files), MONAI (2 files), seaborn (2 files), Nilearn (1 file), scikit-learn (1 file), SciPy (1 file), Hugging Face Transformers (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
18 files

Tracing map

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

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 17 scripts, each with its path and the digest of its content;
  • 7 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

Datasets cited

Data Availability

Data is publicly accessible upon registration at https://www.synapse.org/Synapse:syn28546456/wiki/633440.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 7 MeSH terms, 1 funder, 17 references.

Cite

This paper

Tshimanga, L. F., Zanola, A., Del Pup, F., & Atzori, M. (2026). TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation. PloS one, 21(8), e0354511. https://doi.org/10.1371/journal.pone.0354511

BibTeX

@article{tshimanga2026tokenunet,
author = {Tshimanga, Louis Fabrice and Zanola, Andrea and Del Pup, Federico and Atzori, Manfredo},
title = {{TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation}},
journal = {PloS one},
year = {2026},
month = aug,
volume = {21},
number = {8},
pages = {e0354511},
publisher = {PLOS},
issn = {1932-6203},
doi = {10.1371/journal.pone.0354511},
url = {https://doi.org/10.1371/journal.pone.0354511},
pmid = {42555631},
pmcid = {PMC13440844}
}

RIS

TY - JOUR
AU - Tshimanga, Louis Fabrice
AU - Zanola, Andrea
AU - Del Pup, Federico
AU - Atzori, Manfredo
TI - TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation
T2 - PloS one
J2 - PLoS One
PY - 2026
DA - 2026/08/05
VL - 21
IS - 8
SP - e0354511
SN - 1932-6203
PB - PLOS
DO - 10.1371/journal.pone.0354511
UR - https://doi.org/10.1371/journal.pone.0354511
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pone.0354511",
"type": "article-journal",
"title": "TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation",
"container-title": "PloS one",
"author": [
{
"family": "Tshimanga",
"given": "Louis Fabrice"
},
{
"family": "Zanola",
"given": "Andrea"
},
{
"family": "Del Pup",
"given": "Federico"
},
{
"family": "Atzori",
"given": "Manfredo"
}
],
"container-title-short": "PLoS One",
"volume": "21",
"issue": "8",
"page": "e0354511",
"DOI": "10.1371/journal.pone.0354511",
"PMID": "42555631",
"PMCID": "PMC13440844",
"ISSN": "1932-6203",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pone.0354511",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
5
]
]
}
}

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.21037/qims-2026-0792 [code]
An nnU-Net-based framework with adaptive feature representation for 3D brain tumor segmentation.
Journal: Quantitative imaging in medicine and surgery
In common: nnU-Net, NiBabel, PyTorch, 6 other tools, methods / tools, other condition, 3 references
[2] doi:10.3389/fnins.2026.1870124 [code]
An end-to-end pipeline for automated fetal brain segmentation and biometry from 3D SSFP MRI.
Journal: Frontiers in neuroscience
In common: nnU-Net, MONAI, NiBabel, 7 other tools, 1 reference
[3] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: nnU-Net, MONAI, NiBabel, 7 other tools, methods / tools
[4] doi:10.1002/hipo.70124 [code]
Association Between Anterior Hippocampal Gyrification and Episodic Memory Performance in Neurotypical Young Adults.
Journal: Hippocampus
In common: nnU-Net, Nilearn, NiBabel, 7 other tools, 1 reference
[5] doi:10.1371/journal.pcbi.1014555 [code]
Body surface potential driven personalisation of electrophysiological digital twins in hypertrophic cardiomyopathy.
Journal: PLoS computational biology
In common: nnU-Net, MONAI, NiBabel, 7 other tools
[6] doi:10.1016/j.patter.2026.101538 [code]
A multi-modal foundation model for brain disease diagnosis and medical imaging.
Journal: Patterns (New York, N.Y.)
In common: MONAI, Hugging Face Transformers, NiBabel, 6 other tools, 1 reference
[7] doi:10.1038/s41467-026-72253-7 [code]
Spurious alignment between large language models and brains can emerge from non-robust methods and overlooked confounds.
Journal: Nature communications
In common: Hugging Face Transformers, Nilearn, NiBabel, 7 other tools, methods / tools
[8] doi:10.1186/s13244-026-02296-3 [code]
A pre-trained foundation model framework for multiplanar MRI classification of extramural vascular invasion and mesorectal fascia invasion in rectal cancer.
Journal: Insights into imaging
In common: nnU-Net, MONAI, NiBabel, 6 other tools, other condition
[9] doi:10.1038/s41467-026-73996-z [code]
Genetic architecture of white matter microstructure captured by unsupervised deep representation learning of fractional anisotropy maps.
Journal: Nature communications
In common: MONAI, Nilearn, NiBabel, 7 other tools
[10] doi:10.3389/fmed.2026.1875760 [code]
Adaptive multi-stage domain unlearning for white-matter lesion segmentation.
Journal: Frontiers in medicine
In common: nnU-Net, NiBabel, PyTorch, 6 other tools, methods / 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.