TokenUNet: A new case for transformers integration in efficient and interpretable 3D UNets for brain imaging segmentation.
The 7 matches
- [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] § 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] § Materials and methods › Training ↔ code/scripts/profiling_archs.py, lines 333–413 · score 0.74 · deep supervision, Nesterov, mirroring, momentum, nnUNet, SGD
- [4] § Materials and methods › Training ↔ code/custom/nnUNetTokenUNetTrainer.py, lines 175–236 · score 0.65 · deep supervision, nnUNet, schedule, SGD, augmented, batch
- [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] § 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] § 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
- import torch
- import torch.nn as nn
- import math
- from .tokenmixers import *
- __all__ = [
- "Block",
- "BlockUp",
- "Stage",
- "StageUp",
- "CNNEnc",
- "CNNDec",
- "SpatialAttentionMaskMaker3d",
- "TokenLearner3d",
- "TokenFuser3d",
- "TokenUNet"
- ]
- class Block(nn.Module):
- def __init__(
- self,
- in_channels,
- out_channels,
- kernel_size,
- p=0.0,
- downsample=False
- ):
- super().__init__()
- self.block = nn.Sequential(
- nn.InstanceNorm3d(in_channels),
- nn.Conv3d(in_channels, in_channels*2, kernel_size, stride=2 if downsample else 1, padding=1),
- nn.LeakyReLU(0.1),
- nn.InstanceNorm3d(in_channels*2),
- nn.Conv3d(in_channels*2, out_channels, kernel_size=1,padding="same"),
- )
- nn.init.kaiming_normal_(self.block[1].weight, mode='fan_in', nonlinearity='relu')
- nn.init.kaiming_normal_(self.block[-1].weight, mode='fan_in', nonlinearity='relu')
- self.res_alpha = nn.Parameter(torch.tensor([-2.5]))
- adjust_volume = nn.AvgPool3d(2,2) if downsample else nn.Identity()
- adjust_channels = nn.Conv3d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity()
- self.adjust = nn.Sequential(adjust_volume, adjust_channels)
- def forward(self, x):
- res = self.block(x)
- alpha = torch.sigmoid(self.res_alpha)
- x = self.adjust(x)
- x = (1.-alpha)*x + alpha*res
- return x
- class BlockUp(nn.Module):
- def __init__(
- self,
- in_channels,
- out_channels,
- kernel_size,
- p=0.0,
- upsample=False
- ):
- super().__init__()
- self.block = nn.Sequential(
- nn.InstanceNorm3d(in_channels),
- nn.Conv3d(in_channels, in_channels*2, kernel_size, padding=1 ),
- nn.LeakyReLU(0.1),
- nn.InstanceNorm3d(in_channels*2),
- 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),
- )
- nn.init.kaiming_normal_(self.block[1].weight, mode='fan_in', nonlinearity='relu')
- nn.init.kaiming_normal_(self.block[-1].weight, mode='fan_in', nonlinearity='relu')
- self.res_alpha = nn.Parameter(torch.tensor([-2.5]))
- adjust_volume = nn.Upsample(scale_factor=2, mode="trilinear") if upsample else nn.Identity()
- adjust_channels = nn.Conv3d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity()
- self.adjust = nn.Sequential(adjust_volume, adjust_channels)
- def forward(self, x):
- res = self.block(x)
- alpha = torch.sigmoid(self.res_alpha)
- x = self.adjust(x)
- x = (1.-alpha)*x + alpha*res
- return x
- class Stage(nn.Module):
- def __init__(
- self,
- in_channels,
- out_channels,
- kernel_size,
- n_blocks,
- downsample=True
- ):
- super().__init__()
- self.blocks = nn.ModuleList(
- [Block(in_channels, in_channels, kernel_size) for _ in range(n_blocks-1)]+
- [Block(in_channels, out_channels, kernel_size, downsample=downsample)]
- )
- def forward(self, x):
- for block in self.blocks:
- x = block(x)
- return x
- class StageUp(nn.Module):
- def __init__(
- self,
- in_channels,
- out_channels,
- kernel_size,
- n_blocks,
- upsample=True
- ):
- super().__init__()
- self.blocks = nn.ModuleList(
- [BlockUp(in_channels, in_channels, kernel_size) for _ in range(n_blocks-1)]+
- [BlockUp(in_channels, out_channels, kernel_size, upsample=upsample)]
- )
- def forward(self, x):
- for block in self.blocks:
- x = block(x)
- return x
- class CNNEnc(nn.Module):
- def __init__(
- self,
- in_channels,
- stage_channels,
- kernel_size,
- blocks_per_stage
- ):
- super().__init__()
- n_stages = len(stage_channels)
- # Normalise blocks_per_stage: accept a single int or a per-stage list
- if isinstance(blocks_per_stage, int):
- blocks_per_stage = [blocks_per_stage] * n_stages
- assert len(blocks_per_stage) == n_stages, (
- f"blocks_per_stage length ({len(blocks_per_stage)}) "
- f"must match the number of stages ({n_stages})"
- )
- self.norm = nn.InstanceNorm3d(in_channels)
- # Build stages; the output of stage i is the input of stage i+1
- in_ch = in_channels
- self.stages = nn.ModuleList()
- first = 0
- for out_ch, n_blocks in zip(stage_channels, blocks_per_stage):
- self.stages.append(
- Stage(in_ch, out_ch, kernel_size, n_blocks, downsample=first>0)#last!=n_stages)
- )
- first += 1
- in_ch = out_ch
- def forward(self, x):
- x = self.norm(x)
- maps_to_decoder = []
- for stage in self.stages:
- x = stage(x)
- maps_to_decoder.append(x)
- #print(x.shape)
- return maps_to_decoder
- class CNNDec(nn.Module):
- def __init__(
- self,
- in_channels,
- stage_channels,
- kernel_size,
- blocks_per_stage
- ):
- super().__init__()
- n_stages = len(stage_channels)
- # Normalise blocks_per_stage: accept a single int or a per-stage list
- if isinstance(blocks_per_stage, int):
- blocks_per_stage = [blocks_per_stage] * n_stages
- assert len(blocks_per_stage) == n_stages, (
- f"blocks_per_stage length ({len(blocks_per_stage)}) "
- f"must match the number of stages ({n_stages})"
- )
- #self.norm = nn.InstanceNorm3d(in_channels)
- # Build stages; the output of stage i is the input of stage i+1
- in_ch = in_channels
- self.stages = nn.ModuleList()
- #last = 0
- for out_ch, n_blocks in zip(stage_channels, blocks_per_stage):
- self.stages.append(
- StageUp(in_ch, out_ch, kernel_size, n_blocks, upsample=True)
- )
- in_ch = out_ch
- #last += 1
- def forward(self, maps_to_decoder):
- x = 0.0
- for idx,stage in enumerate(self.stages):
- x = stage(maps_to_decoder[idx]+x)
- #print(maps_to_decoder[idx].shape)
- return x
- class SpatialAttentionMaskMaker3d(nn.Module):
- """
- For each spatial location in the feature map, learns a direct linear mapping
- from the local feature vector to N token assignment scores (one per token).
- A sigmoid then turns each score into an independent soft weight in (0, 1).
- Input: feature map of shape (B, C, H, W, D)
- Output: attention masks of shape (B, n_tokens, H, W, D), values in (0, 1)
- """
- def __init__(self, in_channels: int, n_tokens: int = 8, bias: bool = True):
- super().__init__()
- self.in_channels = in_channels
- self.n_tokens = n_tokens
- #self.norm = nn.LayerNorm(in_channels, elementwise_affine=False)
- # Single linear layer: maps each voxel's C-dim feature vector directly
- # to n_tokens scalar logits. Acts like a learned 1x1x1 convolution.
- # (B, n_voxels, C) -> (B, n_voxels, n_tokens)
- self.linear = nn.Linear(in_features=in_channels, out_features=n_tokens, bias=bias)
- def norm(self, x):
- return torch.nn.functional.normalize(x, p=2, dim=-1) * math.sqrt(x.shape[-1])
- def forward(self, feature_map_bchwd: torch.Tensor):
- """
- Args:
- feature_map_bchwd: (B, C, H, W, D)
- Returns:
- attention_masks_bnhwd: (B, n_tokens, H, W, D)
- """
- B, C, H, W, D = feature_map_bchwd.shape
- # assert C == self.in_channels, f"Expected in_channels={self.in_channels}, got {C}"
- # Flatten spatial dims to treat all voxels uniformly.
- # (B, C, H, W, D) -> (B, C, n_voxels)
- n_voxels = H * W * D
- feature_map_bcv = feature_map_bchwd.reshape(B, C, n_voxels)
- # Move channels last so Linear operates on the feature dimension.
- # (B, C, n_voxels) -> (B, n_voxels, C)
- feature_map_bvc = feature_map_bcv.permute(0, 2, 1)
- feature_map_bvc = self.norm(feature_map_bvc)
- # Direct linear projection: each voxel's feature vector -> n_tokens scores.
- # (B, n_voxels, C) -> (B, n_voxels, n_tokens)
- token_scores_bvn = self.linear(feature_map_bvc)
- # Sigmoid: independent soft weight per (voxel, token) pair, no competition across voxels.
- # (B, n_voxels, n_tokens) -> same shape, values in (0, 1)
- attention_weights_bvn = torch.sigmoid(token_scores_bvn)
- # Restore spatial structure and move token dim before spatial dims.
- # (B, n_voxels, n_tokens) -> (B, n_tokens, n_voxels) -> (B, n_tokens, H, W, D)
- attention_masks_bnhwd = attention_weights_bvn.permute(0, 2, 1).reshape(
- B, self.n_tokens, H, W, D
- )
- return attention_masks_bnhwd
- class TokenLearner3d(nn.Module):
- """
- TokenLearner: compresses a volumetric feature map into a small set of learned tokens.
- The key idea:
- 1. SpatialAttentionMaskMaker produces N soft spatial masks (B, N, H, W, D)
- 2. Each mask is element-wise multiplied with the feature map (B, C, H, W, D)
- 3. The masked feature map is mean-pooled over (H, W, D) -> one C-dim token per mask
- 4. Result: set of N tokens, shape (B, N, C), a compact set-like representation
- This replaces the global average pool (which treats all voxels equally) with
- N learned, content-adaptive pooling operations.
- Args:
- in_channels: number of feature channels C of the input
- n_tokens: number of output tokens N (the compression factor)
- bias: whether Linear layers use bias
- out_channels: if set, project each token from C to out_channels via a Linear layer;
- otherwise tokens keep their original C channels (Identity)
- """
- def __init__(
- self,
- in_channels: int,
- n_tokens: int = 8,
- bias: bool = True,
- out_channels: int = None,
- ):
- super().__init__()
- self.in_channels = in_channels
- self.n_tokens = n_tokens
- #self.norm = nn.LayerNorm(in_channels, elementwise_affine=False)
- self.mask_maker = SpatialAttentionMaskMaker3d(
- in_channels=in_channels,
- n_tokens=n_tokens,
- bias=bias,
- )
- # Optional linear projection applied independently to each token vector
- self.token_projector = (
- nn.Linear(in_channels, out_channels) if out_channels else nn.Identity()
- )
- def norm(self, x):
- return torch.nn.functional.normalize(x, p=2, dim=-1) * math.sqrt(x.shape[-1])
- def forward(self, feature_map_bchwd: torch.Tensor):
- """
- Args:
- feature_map_bchwd: (B, C, H, W, D) — volumetric feature map
- Returns:
- tokens_bnc: (B, n_tokens, C_out) — the learned token set
- attention_masks_bnhwd: (B, n_tokens, H, W, D) — masks for inspection / auxiliary loss
- """
- B, C, H, W, D = feature_map_bchwd.shape
- n_voxels = H * W * D
- # --- Step 1: build one spatial attention mask per token ---
- # Each mask encodes which voxels are relevant for that token.
- # (B, C, H, W, D) -> (B, n_tokens, H, W, D)
- attention_masks_bnhwd = self.mask_maker(feature_map_bchwd)
- attention_masks_flat_bnv = attention_masks_bnhwd.reshape(B, self.n_tokens, n_voxels)
- # --- Step 2: soft-mask the feature map for each token ---
- # Expand feature map along the token dimension so we can broadcast the mask.
- # (B, C, H, W, D) -> (B, 1, C, H, W, D)
- # feature_map_b1chwd = feature_map_bchwd.unsqueeze(1)
- feature_map_bcv = feature_map_bchwd.reshape(B, C, n_voxels)
- # (B, n_tokens, H, W, D) -> (B, n_tokens, 1, H, W, D) for channel broadcast
- #attention_masks_bn1hwd = attention_masks_bnhwd.unsqueeze(2)
- # Element-wise product: each voxel's feature vector is scaled by its token weight.
- # (B, n_tokens, 1, H, W, D) * (B, 1, C, H, W, D) -> (B, n_tokens, C, H, W, D)
- #masked_features_bnchwd = attention_masks_bn1hwd * feature_map_b1chwd
- # --- Step 3: spatial mean-pooling -> one token vector per mask ---
- # Average over the three spatial dimensions (H, W, D).
- # (B, n_tokens, C, H, W, D) -> (B, n_tokens, C)
- #tokens_bnc = masked_features_bnchwd.mean(dim=(3, 4, 5))
- # --- True step 2-3 ---
- # We do not materialize the full (B, n_tokens, C, H, W, D) tensor!
- tokens_bnc = torch.bmm(
- attention_masks_flat_bnv / (attention_masks_flat_bnv.sum(dim=2, keepdims=True) + 1e-8),
- #self.norm(feature_map_bcv.transpose(1,2))
- feature_map_bcv.transpose(1,2)
- ) #/ n_voxels
- # --- Step 4 (optional): project token channels ---
- # (B, n_tokens, C) -> (B, n_tokens, C_out)
- tokens_bnc = self.token_projector(tokens_bnc)
- return tokens_bnc, attention_masks_bnhwd
- class TokenFuser3d(nn.Module):
- """
- This module weighted-averages N tokens and computes N spatial pertinence masks,
- then broadcasts the tokens over the masks to update feature maps.
- """
- def __init__(self,
- n_tokens,
- conv_dim,
- token_dim,
- bias=True
- ):
- super().__init__()
- self.M = nn.Linear(n_tokens, n_tokens)
- self.Beta = SpatialAttentionMaskMaker3d(
- in_channels=conv_dim,
- n_tokens=n_tokens,
- bias=bias,
- )
- self.n_tokens = n_tokens
- self.conv_dim = conv_dim
- self.token_dim = token_dim
- if (token_dim != conv_dim):
- self.C = nn.Linear(token_dim, conv_dim)
- else:
- self.C = nn.Identity()
- def forward(self, tokens, feat_maps):
- B, C, H, W, D = feat_maps.shape # B,Cc,H,[W,[D]]
- n_voxels = H * W * D
- # Mix the tokens, and eventually map them to channel dimension of convolutional features
- mixed_tokens = self.C(self.M(tokens.transpose(1,2)).transpose(1,2)) # B, N, Ct -> B, N, Cc
- # Determine how much each voxel needs each token
- pertinence_masks_bnhwd = self.Beta(feat_maps) # B,Cc,H,[W,[D]]
- pertinence_masks_flat_bnv = pertinence_masks_bnhwd.reshape(B, self.n_tokens, n_voxels) # B, N, V=HWD
- token_broadcast_bcv = torch.bmm(mixed_tokens.transpose(1,2), pertinence_masks_flat_bnv) # B, (Cc, N) @ (N, V) -> (B,Ct,V)
- token_broadcast_bchwd = token_broadcast_bcv.reshape(B,C,H,W,D)
- out_maps = feat_maps + token_broadcast_bchwd
- return out_maps
- class TokenUNet(nn.Module):
- def __init__(
- self,
- in_channels,
- enc_stage_channels,
- kernel_size,
- blocks_per_stage,
- num_classes,
- n_tokens=8,
- token_dim=None,
- tokenize=True,
- attention=False,
- process_tokens=True,
- token_blocks=2,
- bias=True
- ):
- super().__init__()
- self.tokenize = tokenize
- self.process_tokens = process_tokens
- # Ensure token_dim is set (defaults to the channel size of the bottleneck)
- bottleneck_channels = enc_stage_channels[-1]
- dec_stage_channels = enc_stage_channels[::-1][1:]
- self.token_dim = token_dim if token_dim is not None else bottleneck_channels
- # 1. Encoder
- self.encoder = CNNEnc(
- in_channels=in_channels,
- stage_channels=enc_stage_channels,
- kernel_size=kernel_size,
- blocks_per_stage=blocks_per_stage
- )
- # 2. Tokenizer (Optional)
- if self.tokenize:
- self.token_learner = TokenLearner3d(
- in_channels=bottleneck_channels,
- n_tokens=n_tokens,
- bias=bias,
- out_channels=self.token_dim
- )
- # 3. Token Processor (Transformer / MLP Mixer) - Optional
- if self.process_tokens:
- if attention:
- self.token_processor = MyTransformerEncoder(
- d_model=self.token_dim,
- nhead=4,
- dim_feedforward=self.token_dim*2,
- dropout=0.0,
- num_layers=token_blocks
- )
- else:
- self.token_processor = MyMLPMixer(
- n_tokens=n_tokens,
- d_model=self.token_dim,
- dim_feedforward=self.token_dim*2,
- n_blocks=token_blocks,
- dropout=0.0
- )
- else:
- self.token_processor = nn.Identity()
- # 4. Token Fuser
- self.token_fuser = TokenFuser3d(
- n_tokens=n_tokens,
- conv_dim=bottleneck_channels,
- token_dim=self.token_dim,
- bias=bias
- )
- # 5. Decoder
- self.decoder = CNNDec(
- in_channels=bottleneck_channels, # Starts from the bottleneck size
- stage_channels=dec_stage_channels,
- kernel_size=kernel_size,
- blocks_per_stage=[1,]*len(dec_stage_channels)
- )
- # 6. Segmentation Head
- self.segmentation_head = nn.Conv3d(
- in_channels=dec_stage_channels[-1],
- out_channels=num_classes,
- kernel_size=1
- )
- def forward(self, x):
- # Forward through Encoder
- maps_to_decoder = self.encoder(x)
- # The bottleneck is the lowest resolution feature map (the last one)
- bottleneck = maps_to_decoder[-1]
- if self.tokenize:
- # Extract Tokens
- tokens, attention_masks = self.token_learner(bottleneck)
- # Process Tokens (Transformer/Mixer)
- if self.process_tokens:
- tokens = self.token_processor(tokens)
- # Fuse Tokens back into the bottleneck feature map
- bottleneck = self.token_fuser(tokens, bottleneck)
- # Update the bottleneck in our skip-connection list
- maps_to_decoder[-1] = bottleneck
- # CRITICAL: Reverse the maps so they match the Decoder's expected ascending resolutions
- # Example: if Enc outputs [128x128, 64x64, 32x32], Decoder needs [32x32, 64x64, 128x128]
- maps_to_decoder = maps_to_decoder[::-1]
- # Forward through Decoder
- out = self.segmentation_head(self.decoder(maps_to_decoder))
- return out
tokenunet.py at commit 0bc16c2, no license · at the source
Overview
- Department of Neuroscience, University of Padua, Padua, Italy
- Padova Neuroscience Center, University of Padua, Padua, Italy
- Department of Information Engineering, University of Padua, Padua, Italy
- Information Systems Institute, University of Applied Sciences Western Switzerland (HES-SO Valais), Sierre, Switzerland
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
0bc16c2e27df7e6da0fed35952c2ed951f46e4e8, 17 August 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
18 files
- code/
custom/ , Python, 1 line__init__.py - code/
custom/ , Python, 2 linesarchitectures/ __init__.py - code/
custom/ , Python, 1,118 linesarchitectures/ swinunetr_fix.py - code/
custom/ , Python, 162 linesarchitectures/ tokenmixers.py - code/
custom/ , Python, 529 lines, 3 matchesarchitectures/ tokenunet.py - code/
custom/ , Python, 593 lines, 1 matchnnUNetTokenUNetTrainer.p y - code/
notebooks/ , Jupyter, 228 linesboxplots.ipynb - code/
scripts/ , Python, 74 linesalignment_statplots.py - code/
scripts/ , Python, 391 lines, 1 matchconvert_fets_to_nnunet.p y - code/
scripts/ , Python, 93 linescount_pms.py - code/
scripts/ , Python, 460 lines, 1 matchdice_attention.py - code/
scripts/ , Python, 149 linesefficiency_landscape.py - code/
scripts/ , Python, 29 lineshierarchical_labels.py - code/
scripts/ , Python, 611 lines, 1 matchprofiling_archs.py - code/
scripts/ , Python, 19 linesremap_labels.py - code/
scripts/ , Python, 424 linesstats_testing.py - code/
scripts/ , Python, 1,192 linesswinunetr_fix.py - README.md, Text, 4 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:
- 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
- synapse.org/
synapse:syn27046444/ , at Synapse; found in the text, “Data”wiki - synapse.org/
synapse:syn28546456/ , at Synapse; found in “Data Availability”wiki
Data Availability
Data is publicly accessible upon registration at https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, 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://
BibTeX
@article{tshimanga2026to
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/
url = {https://
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/
VL - 21
IS - 8
SP - e0354511
SN - 1932-6203
PB - PLOS
DO - 10.1371/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1371/
"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":
"volume": "21",
"issue": "8",
"page": "e0354511",
"DOI": "10.1371/
"PMID": "42555631",
"PMCID": "PMC13440844",
"ISSN": "1932-6203",
"publisher": "PLOS",
"URL": "https://
"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 surgeryIn 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 neuroscienceIn 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 intelligenceIn 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: HippocampusIn 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 biologyIn 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 communicationsIn common: Hugging Face Transformers, Nilearn, NiBabel, 7 other tools, methods / tools
- [8] doi:10.1002/hbm.70469 [code]
- VarCoNet: A Variability-Aware Self-Supervised Framework for Functional Connectome Extraction From Resting-State fMRI.Journal: Human brain mappingIn common: Hugging Face Transformers, Nilearn, NiBabel, 7 other tools, methods / tools
- [9] 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 imagingIn common: nnU-Net, MONAI, NiBabel, 6 other tools, other condition
- [10] 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 communicationsIn common: MONAI, Nilearn, NiBabel, 7 other tools
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 17 scripts, and 7 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:54dc39be938a59eb…
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.
