OSCR

AI Augmented Confocal Laser Endomicroscopy for Rapid Intraoperative Diagnosis of Brain Tumors.

Code ↔ Paper

2 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 2 matches
  1. [1] § Methods › Development of AI diagnostic model ↔ gradientCAM/swin_agcam/models/swin_transformer.py, lines 586–706 · score 0.52 · Swin Transformer, patch embedding, dimensions, classifier, model
  2. [2] § Methods › Development of AI diagnostic model ↔ gradientCAM/swin_agcam/models/swin_transformer_v2.py, lines 593–707 · score 0.52 · Swin Transformer, patch embedding, dimensions, classifier, model

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,085 lines · 43 KB · no license · 1 match

  1. """ Swin Transformer
  2. A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows`
  3. - https://arxiv.org/pdf/2103.14030
  4. Code/weights from https://github.com/microsoft/Swin-Transformer, original copyright/license info below
  5. S3 (AutoFormerV2, https://arxiv.org/abs/2111.14725) Swin weights from
  6. - https://github.com/microsoft/Cream/tree/main/AutoFormerV2
  7. Modifications and additions for timm hacked together by / Copyright 2021, Ross Wightman
  8. """
  9. # --------------------------------------------------------
  10. # Swin Transformer
  11. # Copyright (c) 2021 Microsoft
  12. # Licensed under The MIT License [see LICENSE for details]
  13. # Written by Ze Liu
  14. # --------------------------------------------------------
  15. import logging
  16. import math
  17. from typing import Callable, List, Optional, Tuple, Union
  18. import torch
  19. import torch.nn as nn
  20. from ..data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD
  21. from ..layers import PatchEmbed, Mlp, DropPath, ClassifierHead, to_2tuple, to_ntuple, trunc_normal_, \
  22. _assert, use_fused_attn, resize_rel_pos_bias_table, resample_patch_embed, ndgrid
  23. from ._builder import build_model_with_cfg
  24. from ._features import feature_take_indices
  25. from ._features_fx import register_notrace_function
  26. from ._manipulate import checkpoint_seq, named_apply
  27. from ._registry import generate_default_cfgs, register_model, register_model_deprecations
  28. from .vision_transformer import get_init_weights_vit
  29. __all__ = ['SwinTransformer'] # model_registry will add each entrypoint fn to this
  30. _logger = logging.getLogger(__name__)
  31. _int_or_tuple_2_t = Union[int, Tuple[int, int]]
  32. def window_partition(
  33. x: torch.Tensor,
  34. window_size: Tuple[int, int],
  35. ) -> torch.Tensor:
  36. """
  37. Partition into non-overlapping windows with padding if needed.
  38. Args:
  39. x (tensor): input tokens with [B, H, W, C].
  40. window_size (int): window size.
  41. Returns:
  42. windows: windows after partition with [B * num_windows, window_size, window_size, C].
  43. (Hp, Wp): padded height and width before partition
  44. """
  45. B, H, W, C = x.shape
  46. x = x.view(B, H // window_size[0], window_size[0], W // window_size[1], window_size[1], C)
  47. windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size[0], window_size[1], C)
  48. return windows
  49. @register_notrace_function # reason: int argument is a Proxy
  50. def window_reverse(windows, window_size: Tuple[int, int], H: int, W: int):
  51. """
  52. Args:
  53. windows: (num_windows*B, window_size, window_size, C)
  54. window_size (int): Window size
  55. H (int): Height of image
  56. W (int): Width of image
  57. Returns:
  58. x: (B, H, W, C)
  59. """
  60. C = windows.shape[-1]
  61. x = windows.view(-1, H // window_size[0], W // window_size[1], window_size[0], window_size[1], C)
  62. x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, H, W, C)
  63. return x
  64. def get_relative_position_index(win_h: int, win_w: int):
  65. # get pair-wise relative position index for each token inside the window
  66. coords = torch.stack(ndgrid(torch.arange(win_h), torch.arange(win_w))) # 2, Wh, Ww
  67. coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
  68. relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
  69. relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
  70. relative_coords[:, :, 0] += win_h - 1 # shift to start from 0
  71. relative_coords[:, :, 1] += win_w - 1
  72. relative_coords[:, :, 0] *= 2 * win_w - 1
  73. return relative_coords.sum(-1) # Wh*Ww, Wh*Ww
  74. class WindowAttention(nn.Module):
  75. """ Window based multi-head self attention (W-MSA) module with relative position bias.
  76. It supports shifted and non-shifted windows.
  77. """
  78. fused_attn: torch.jit.Final[bool]
  79. def __init__(
  80. self,
  81. dim: int,
  82. num_heads: int,
  83. head_dim: Optional[int] = None,
  84. window_size: _int_or_tuple_2_t = 7,
  85. qkv_bias: bool = True,
  86. attn_drop: float = 0.,
  87. proj_drop: float = 0.,
  88. ):
  89. """
  90. Args:
  91. dim: Number of input channels.
  92. num_heads: Number of attention heads.
  93. head_dim: Number of channels per head (dim // num_heads if not set)
  94. window_size: The height and width of the window.
  95. qkv_bias: If True, add a learnable bias to query, key, value.
  96. attn_drop: Dropout ratio of attention weight.
  97. proj_drop: Dropout ratio of output.
  98. """
  99. super().__init__()
  100. self.dim = dim
  101. self.window_size = to_2tuple(window_size) # Wh, Ww
  102. win_h, win_w = self.window_size
  103. self.window_area = win_h * win_w
  104. self.num_heads = num_heads
  105. head_dim = head_dim or dim // num_heads
  106. attn_dim = head_dim * num_heads
  107. self.scale = head_dim ** -0.5
  108. self.fused_attn = use_fused_attn(experimental=True) # NOTE not tested for prime-time yet
  109. # define a parameter table of relative position bias, shape: 2*Wh-1 * 2*Ww-1, nH
  110. self.relative_position_bias_table = nn.Parameter(torch.zeros((2 * win_h - 1) * (2 * win_w - 1), num_heads))
  111. # get pair-wise relative position index for each token inside the window
  112. self.register_buffer("relative_position_index", get_relative_position_index(win_h, win_w), persistent=False)
  113. self.qkv = nn.Linear(dim, attn_dim * 3, bias=qkv_bias)
  114. self.attn_drop = nn.Dropout(attn_drop)
  115. self.proj = nn.Linear(attn_dim, dim)
  116. self.proj_drop = nn.Dropout(proj_drop)
  117. trunc_normal_(self.relative_position_bias_table, std=.02)
  118. self.softmax = nn.Softmax(dim=-1)
  119. # Identity layers for pytorch hook
  120. self.forward_hook_before_softmax = nn.Identity()
  121. self.backward_hook_after_softmax = nn.Identity()
  122. def set_window_size(self, window_size: Tuple[int, int]) -> None:
  123. """Update window size & interpolate position embeddings
  124. Args:
  125. window_size (int): New window size
  126. """
  127. window_size = to_2tuple(window_size)
  128. if window_size == self.window_size:
  129. return
  130. self.window_size = window_size
  131. win_h, win_w = self.window_size
  132. self.window_area = win_h * win_w
  133. with torch.no_grad():
  134. new_bias_shape = (2 * win_h - 1) * (2 * win_w - 1), self.num_heads
  135. self.relative_position_bias_table = nn.Parameter(
  136. resize_rel_pos_bias_table(
  137. self.relative_position_bias_table,
  138. new_window_size=self.window_size,
  139. new_bias_shape=new_bias_shape,
  140. ))
  141. self.register_buffer("relative_position_index", get_relative_position_index(win_h, win_w), persistent=False)
  142. def _get_rel_pos_bias(self) -> torch.Tensor:
  143. relative_position_bias = self.relative_position_bias_table[
  144. self.relative_position_index.view(-1)].view(self.window_area, self.window_area, -1) # Wh*Ww,Wh*Ww,nH
  145. relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
  146. return relative_position_bias.unsqueeze(0)
  147. def forward(self, x, mask: Optional[torch.Tensor] = None):
  148. """
  149. Args:
  150. x: input features with shape of (num_windows*B, N, C)
  151. mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
  152. """
  153. B_, N, C = x.shape
  154. qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)
  155. q, k, v = qkv.unbind(0)
  156. if self.fused_attn:
  157. attn_mask = self._get_rel_pos_bias()
  158. if mask is not None:
  159. num_win = mask.shape[0]
  160. mask = mask.view(1, num_win, 1, N, N).expand(B_ // num_win, -1, self.num_heads, -1, -1)
  161. attn_mask = attn_mask + mask.reshape(-1, self.num_heads, N, N)
  162. x = torch.nn.functional.scaled_dot_product_attention(
  163. q, k, v,
  164. attn_mask=attn_mask,
  165. dropout_p=self.attn_drop.p if self.training else 0.,
  166. )
  167. else:
  168. q = q * self.scale
  169. attn = q @ k.transpose(-2, -1)
  170. attn = self.forward_hook_before_softmax(attn)
  171. attn = self.backward_hook_after_softmax(attn) # 왜 옮겼더니 된거지?
  172. attn = attn + self._get_rel_pos_bias()
  173. if mask is not None:
  174. num_win = mask.shape[0]
  175. attn = attn.view(-1, num_win, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
  176. attn = attn.view(-1, self.num_heads, N, N)
  177. #forward hook
  178. attn = self.softmax(attn)
  179. attn = self.attn_drop(attn)
  180. #backward hook
  181. x = attn @ v
  182. x = x.transpose(1, 2).reshape(B_, N, -1)
  183. x = self.proj(x)
  184. x = self.proj_drop(x)
  185. return x
  186. class SwinTransformerBlock(nn.Module):
  187. """ Swin Transformer Block.
  188. """
  189. def __init__(
  190. self,
  191. dim: int,
  192. input_resolution: _int_or_tuple_2_t,
  193. num_heads: int = 4,
  194. head_dim: Optional[int] = None,
  195. window_size: _int_or_tuple_2_t = 7,
  196. shift_size: int = 0,
  197. always_partition: bool = False,
  198. dynamic_mask: bool = False,
  199. mlp_ratio: float = 4.,
  200. qkv_bias: bool = True,
  201. proj_drop: float = 0.,
  202. attn_drop: float = 0.,
  203. drop_path: float = 0.,
  204. act_layer: Callable = nn.GELU,
  205. norm_layer: Callable = nn.LayerNorm,
  206. ):
  207. """
  208. Args:
  209. dim: Number of input channels.
  210. input_resolution: Input resolution.
  211. window_size: Window size.
  212. num_heads: Number of attention heads.
  213. head_dim: Enforce the number of channels per head
  214. shift_size: Shift size for SW-MSA.
  215. always_partition: Always partition into full windows and shift
  216. mlp_ratio: Ratio of mlp hidden dim to embedding dim.
  217. qkv_bias: If True, add a learnable bias to query, key, value.
  218. proj_drop: Dropout rate.
  219. attn_drop: Attention dropout rate.
  220. drop_path: Stochastic depth rate.
  221. act_layer: Activation layer.
  222. norm_layer: Normalization layer.
  223. """
  224. super().__init__()
  225. self.dim = dim
  226. self.input_resolution = input_resolution
  227. self.target_shift_size = to_2tuple(shift_size) # store for later resize
  228. self.always_partition = always_partition
  229. self.dynamic_mask = dynamic_mask
  230. self.window_size, self.shift_size = self._calc_window_shift(window_size, shift_size)
  231. self.window_area = self.window_size[0] * self.window_size[1]
  232. self.mlp_ratio = mlp_ratio
  233. self.norm1 = norm_layer(dim)
  234. self.attn = WindowAttention(
  235. dim,
  236. num_heads=num_heads,
  237. head_dim=head_dim,
  238. window_size=self.window_size,
  239. qkv_bias=qkv_bias,
  240. attn_drop=attn_drop,
  241. proj_drop=proj_drop,
  242. )
  243. self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()
  244. self.norm2 = norm_layer(dim)
  245. self.mlp = Mlp(
  246. in_features=dim,
  247. hidden_features=int(dim * mlp_ratio),
  248. act_layer=act_layer,
  249. drop=proj_drop,
  250. )
  251. self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()
  252. self.register_buffer(
  253. "attn_mask",
  254. None if self.dynamic_mask else self.get_attn_mask(),
  255. persistent=False,
  256. )
  257. self.backward_hook_before_mlp = nn.Identity()
  258. def get_attn_mask(self, x: Optional[torch.Tensor] = None) -> Optional[torch.Tensor]:
  259. if any(self.shift_size):
  260. # calculate attention mask for SW-MSA
  261. if x is not None:
  262. H, W = x.shape[1], x.shape[2]
  263. device = x.device
  264. dtype = x.dtype
  265. else:
  266. H, W = self.input_resolution
  267. device = None
  268. dtype = None
  269. H = math.ceil(H / self.window_size[0]) * self.window_size[0]
  270. W = math.ceil(W / self.window_size[1]) * self.window_size[1]
  271. img_mask = torch.zeros((1, H, W, 1), dtype=dtype, device=device) # 1 H W 1
  272. cnt = 0
  273. for h in (
  274. (0, -self.window_size[0]),
  275. (-self.window_size[0], -self.shift_size[0]),
  276. (-self.shift_size[0], None),
  277. ):
  278. for w in (
  279. (0, -self.window_size[1]),
  280. (-self.window_size[1], -self.shift_size[1]),
  281. (-self.shift_size[1], None),
  282. ):
  283. img_mask[:, h[0]:h[1], w[0]:w[1], :] = cnt
  284. cnt += 1
  285. mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1
  286. mask_windows = mask_windows.view(-1, self.window_area)
  287. attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
  288. attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
  289. else:
  290. attn_mask = None
  291. return attn_mask
  292. def _calc_window_shift(
  293. self,
  294. target_window_size: Union[int, Tuple[int, int]],
  295. target_shift_size: Optional[Union[int, Tuple[int, int]]] = None,
  296. ) -> Tuple[Tuple[int, int], Tuple[int, int]]:
  297. target_window_size = to_2tuple(target_window_size)
  298. if target_shift_size is None:
  299. # if passed value is None, recalculate from default window_size // 2 if it was previously non-zero
  300. target_shift_size = self.target_shift_size
  301. if any(target_shift_size):
  302. target_shift_size = (target_window_size[0] // 2, target_window_size[1] // 2)
  303. else:
  304. target_shift_size = to_2tuple(target_shift_size)
  305. if self.always_partition:
  306. return target_window_size, target_shift_size
  307. window_size = [r if r <= w else w for r, w in zip(self.input_resolution, target_window_size)]
  308. shift_size = [0 if r <= w else s for r, w, s in zip(self.input_resolution, window_size, target_shift_size)]
  309. return tuple(window_size), tuple(shift_size)
  310. def set_input_size(
  311. self,
  312. feat_size: Tuple[int, int],
  313. window_size: Tuple[int, int],
  314. always_partition: Optional[bool] = None,
  315. ):
  316. """
  317. Args:
  318. feat_size: New input resolution
  319. window_size: New window size
  320. always_partition: Change always_partition attribute if not None
  321. """
  322. self.input_resolution = feat_size
  323. if always_partition is not None:
  324. self.always_partition = always_partition
  325. self.window_size, self.shift_size = self._calc_window_shift(window_size)
  326. self.window_area = self.window_size[0] * self.window_size[1]
  327. self.attn.set_window_size(self.window_size)
  328. self.register_buffer(
  329. "attn_mask",
  330. None if self.dynamic_mask else self.get_attn_mask(),
  331. persistent=False,
  332. )
  333. def _attn(self, x):
  334. B, H, W, C = x.shape
  335. # cyclic shift
  336. has_shift = any(self.shift_size)
  337. if has_shift:
  338. shifted_x = torch.roll(x, shifts=(-self.shift_size[0], -self.shift_size[1]), dims=(1, 2))
  339. else:
  340. shifted_x = x
  341. # pad for resolution not divisible by window size
  342. pad_h = (self.window_size[0] - H % self.window_size[0]) % self.window_size[0]
  343. pad_w = (self.window_size[1] - W % self.window_size[1]) % self.window_size[1]
  344. shifted_x = torch.nn.functional.pad(shifted_x, (0, 0, 0, pad_w, 0, pad_h))
  345. _, Hp, Wp, _ = shifted_x.shape
  346. # partition windows
  347. x_windows = window_partition(shifted_x, self.window_size) # nW*B, window_size, window_size, C
  348. x_windows = x_windows.view(-1, self.window_area, C) # nW*B, window_size*window_size, C
  349. # W-MSA/SW-MSA
  350. if getattr(self, 'dynamic_mask', False):
  351. attn_mask = self.get_attn_mask(shifted_x)
  352. else:
  353. attn_mask = self.attn_mask
  354. attn_windows = self.attn(x_windows, mask=attn_mask) # nW*B, window_size*window_size, C
  355. # merge windows
  356. attn_windows = attn_windows.view(-1, self.window_size[0], self.window_size[1], C)
  357. shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp) # B H' W' C
  358. shifted_x = shifted_x[:, :H, :W, :].contiguous()
  359. # reverse cyclic shift
  360. if has_shift:
  361. x = torch.roll(shifted_x, shifts=self.shift_size, dims=(1, 2))
  362. else:
  363. x = shifted_x
  364. return x
  365. def forward(self, x):
  366. B, H, W, C = x.shape
  367. x = x + self.drop_path1(self._attn(self.norm1(x)))
  368. x = x.reshape(B, -1, C)
  369. x = x + self.drop_path2(self.mlp(self.norm2(x)))
  370. x = self.backward_hook_before_mlp(x)
  371. x = x.reshape(B, H, W, C)
  372. return x
  373. class PatchMerging(nn.Module):
  374. """ Patch Merging Layer.
  375. """
  376. def __init__(
  377. self,
  378. dim: int,
  379. out_dim: Optional[int] = None,
  380. norm_layer: Callable = nn.LayerNorm,
  381. ):
  382. """
  383. Args:
  384. dim: Number of input channels.
  385. out_dim: Number of output channels (or 2 * dim if None)
  386. norm_layer: Normalization layer.
  387. """
  388. super().__init__()
  389. self.dim = dim
  390. self.out_dim = out_dim or 2 * dim
  391. self.norm = norm_layer(4 * dim)
  392. self.reduction = nn.Linear(4 * dim, self.out_dim, bias=False)
  393. def forward(self, x):
  394. B, H, W, C = x.shape
  395. pad_values = (0, 0, 0, W % 2, 0, H % 2)
  396. x = nn.functional.pad(x, pad_values)
  397. _, H, W, _ = x.shape
  398. x = x.reshape(B, H // 2, 2, W // 2, 2, C).permute(0, 1, 3, 4, 2, 5).flatten(3)
  399. x = self.norm(x)
  400. x = self.reduction(x)
  401. return x
  402. class SwinTransformerStage(nn.Module):
  403. """ A basic Swin Transformer layer for one stage.
  404. """
  405. def __init__(
  406. self,
  407. dim: int,
  408. out_dim: int,
  409. input_resolution: Tuple[int, int],
  410. depth: int,
  411. downsample: bool = True,
  412. num_heads: int = 4,
  413. head_dim: Optional[int] = None,
  414. window_size: _int_or_tuple_2_t = 7,
  415. always_partition: bool = False,
  416. dynamic_mask: bool = False,
  417. mlp_ratio: float = 4.,
  418. qkv_bias: bool = True,
  419. proj_drop: float = 0.,
  420. attn_drop: float = 0.,
  421. drop_path: Union[List[float], float] = 0.,
  422. norm_layer: Callable = nn.LayerNorm,
  423. ):
  424. """
  425. Args:
  426. dim: Number of input channels.
  427. out_dim: Number of output channels.
  428. input_resolution: Input resolution.
  429. depth: Number of blocks.
  430. downsample: Downsample layer at the end of the layer.
  431. num_heads: Number of attention heads.
  432. head_dim: Channels per head (dim // num_heads if not set)
  433. window_size: Local window size.
  434. mlp_ratio: Ratio of mlp hidden dim to embedding dim.
  435. qkv_bias: If True, add a learnable bias to query, key, value.
  436. proj_drop: Projection dropout rate.
  437. attn_drop: Attention dropout rate.
  438. drop_path: Stochastic depth rate.
  439. norm_layer: Normalization layer.
  440. """
  441. super().__init__()
  442. self.dim = dim
  443. self.input_resolution = input_resolution
  444. self.output_resolution = tuple(i // 2 for i in input_resolution) if downsample else input_resolution
  445. self.depth = depth
  446. self.grad_checkpointing = False
  447. window_size = to_2tuple(window_size)
  448. shift_size = tuple([w // 2 for w in window_size])
  449. # patch merging layer
  450. if downsample:
  451. self.downsample = PatchMerging(
  452. dim=dim,
  453. out_dim=out_dim,
  454. norm_layer=norm_layer,
  455. )
  456. else:
  457. assert dim == out_dim
  458. self.downsample = nn.Identity()
  459. # build blocks
  460. self.blocks = nn.Sequential(*[
  461. SwinTransformerBlock(
  462. dim=out_dim,
  463. input_resolution=self.output_resolution,
  464. num_heads=num_heads,
  465. head_dim=head_dim,
  466. window_size=window_size,
  467. shift_size=0 if (i % 2 == 0) else shift_size,
  468. always_partition=always_partition,
  469. dynamic_mask=dynamic_mask,
  470. mlp_ratio=mlp_ratio,
  471. qkv_bias=qkv_bias,
  472. proj_drop=proj_drop,
  473. attn_drop=attn_drop,
  474. drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
  475. norm_layer=norm_layer,
  476. )
  477. for i in range(depth)])
  478. def set_input_size(
  479. self,
  480. feat_size: Tuple[int, int],
  481. window_size: int,
  482. always_partition: Optional[bool] = None,
  483. ):
  484. """ Updates the resolution, window size and so the pair-wise relative positions.
  485. Args:
  486. feat_size: New input (feature) resolution
  487. window_size: New window size
  488. always_partition: Always partition / shift the window
  489. """
  490. self.input_resolution = feat_size
  491. if isinstance(self.downsample, nn.Identity):
  492. self.output_resolution = feat_size
  493. else:
  494. self.output_resolution = tuple(i // 2 for i in feat_size)
  495. for block in self.blocks:
  496. block.set_input_size(
  497. feat_size=self.output_resolution,
  498. window_size=window_size,
  499. always_partition=always_partition,
  500. )
  501. def forward(self, x):
  502. x = self.downsample(x)
  503. if self.grad_checkpointing and not torch.jit.is_scripting():
  504. x = checkpoint_seq(self.blocks, x)
  505. else:
  506. x = self.blocks(x)
  507. return x
  508. class SwinTransformer(nn.Module):
  509. """ Swin Transformer
  510. A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
  511. https://arxiv.org/pdf/2103.14030
  512. """
  513. def __init__(
  514. self,
  515. img_size: _int_or_tuple_2_t = 224,
  516. patch_size: int = 4,
  517. in_chans: int = 3,
  518. num_classes: int = 1000,
  519. global_pool: str = 'avg',
  520. embed_dim: int = 96,
  521. depths: Tuple[int, ...] = (2, 2, 6, 2),
  522. num_heads: Tuple[int, ...] = (3, 6, 12, 24),
  523. head_dim: Optional[int] = None,
  524. window_size: _int_or_tuple_2_t = 7,
  525. always_partition: bool = False,
  526. strict_img_size: bool = True,
  527. mlp_ratio: float = 4.,
  528. qkv_bias: bool = True,
  529. drop_rate: float = 0.,
  530. proj_drop_rate: float = 0.,
  531. attn_drop_rate: float = 0.,
  532. drop_path_rate: float = 0.1,
  533. embed_layer: Callable = PatchEmbed,
  534. norm_layer: Union[str, Callable] = nn.LayerNorm,
  535. weight_init: str = '',
  536. **kwargs,
  537. ):
  538. """
  539. Args:
  540. img_size: Input image size.
  541. patch_size: Patch size.
  542. in_chans: Number of input image channels.
  543. num_classes: Number of classes for classification head.
  544. embed_dim: Patch embedding dimension.
  545. depths: Depth of each Swin Transformer layer.
  546. num_heads: Number of attention heads in different layers.
  547. head_dim: Dimension of self-attention heads.
  548. window_size: Window size.
  549. mlp_ratio: Ratio of mlp hidden dim to embedding dim.
  550. qkv_bias: If True, add a learnable bias to query, key, value.
  551. drop_rate: Dropout rate.
  552. attn_drop_rate (float): Attention dropout rate.
  553. drop_path_rate (float): Stochastic depth rate.
  554. embed_layer: Patch embedding layer.
  555. norm_layer (nn.Module): Normalization layer.
  556. """
  557. super().__init__()
  558. assert global_pool in ('', 'avg')
  559. self.num_classes = num_classes
  560. self.global_pool = global_pool
  561. self.output_fmt = 'NHWC'
  562. self.num_layers = len(depths)
  563. self.embed_dim = embed_dim
  564. self.num_features = self.head_hidden_size = int(embed_dim * 2 ** (self.num_layers - 1))
  565. self.feature_info = []
  566. if not isinstance(embed_dim, (tuple, list)):
  567. embed_dim = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
  568. # split image into non-overlapping patches
  569. self.patch_embed = embed_layer(
  570. img_size=img_size,
  571. patch_size=patch_size,
  572. in_chans=in_chans,
  573. embed_dim=embed_dim[0],
  574. norm_layer=norm_layer,
  575. strict_img_size=strict_img_size,
  576. output_fmt='NHWC',
  577. )
  578. patch_grid = self.patch_embed.grid_size
  579. # build layers
  580. head_dim = to_ntuple(self.num_layers)(head_dim)
  581. if not isinstance(window_size, (list, tuple)):
  582. window_size = to_ntuple(self.num_layers)(window_size)
  583. elif len(window_size) == 2:
  584. window_size = (window_size,) * self.num_layers
  585. assert len(window_size) == self.num_layers
  586. mlp_ratio = to_ntuple(self.num_layers)(mlp_ratio)
  587. dpr = [x.tolist() for x in torch.linspace(0, drop_path_rate, sum(depths)).split(depths)]
  588. layers = []
  589. in_dim = embed_dim[0]
  590. scale = 1
  591. for i in range(self.num_layers):
  592. out_dim = embed_dim[i]
  593. layers += [SwinTransformerStage(
  594. dim=in_dim,
  595. out_dim=out_dim,
  596. input_resolution=(
  597. patch_grid[0] // scale,
  598. patch_grid[1] // scale
  599. ),
  600. depth=depths[i],
  601. downsample=i > 0,
  602. num_heads=num_heads[i],
  603. head_dim=head_dim[i],
  604. window_size=window_size[i],
  605. always_partition=always_partition,
  606. dynamic_mask=not strict_img_size,
  607. mlp_ratio=mlp_ratio[i],
  608. qkv_bias=qkv_bias,
  609. proj_drop=proj_drop_rate,
  610. attn_drop=attn_drop_rate,
  611. drop_path=dpr[i],
  612. norm_layer=norm_layer,
  613. )]
  614. in_dim = out_dim
  615. if i > 0:
  616. scale *= 2
  617. self.feature_info += [dict(num_chs=out_dim, reduction=patch_size * scale, module=f'layers.{i}')]
  618. self.layers = nn.Sequential(*layers)
  619. self.norm = norm_layer(self.num_features)
  620. self.head = ClassifierHead(
  621. self.num_features,
  622. num_classes,
  623. pool_type=global_pool,
  624. drop_rate=drop_rate,
  625. input_fmt=self.output_fmt,
  626. )
  627. if weight_init != 'skip':
  628. self.init_weights(weight_init)
  629. @torch.jit.ignore
  630. def init_weights(self, mode=''):
  631. assert mode in ('jax', 'jax_nlhb', 'moco', '')
  632. head_bias = -math.log(self.num_classes) if 'nlhb' in mode else 0.
  633. named_apply(get_init_weights_vit(mode, head_bias=head_bias), self)
  634. @torch.jit.ignore
  635. def no_weight_decay(self):
  636. nwd = set()
  637. for n, _ in self.named_parameters():
  638. if 'relative_position_bias_table' in n:
  639. nwd.add(n)
  640. return nwd
  641. def set_input_size(
  642. self,
  643. img_size: Optional[Tuple[int, int]] = None,
  644. patch_size: Optional[Tuple[int, int]] = None,
  645. window_size: Optional[Tuple[int, int]] = None,
  646. window_ratio: int = 8,
  647. always_partition: Optional[bool] = None,
  648. ) -> None:
  649. """ Updates the image resolution and window size.
  650. Args:
  651. img_size: New input resolution, if None current resolution is used
  652. patch_size (Optional[Tuple[int, int]): New patch size, if None use current patch size
  653. window_size: New window size, if None based on new_img_size // window_div
  654. window_ratio: divisor for calculating window size from grid size
  655. always_partition: always partition into windows and shift (even if window size < feat size)
  656. """
  657. if img_size is not None or patch_size is not None:
  658. self.patch_embed.set_input_size(img_size=img_size, patch_size=patch_size)
  659. patch_grid = self.patch_embed.grid_size
  660. if window_size is None:
  661. window_size = tuple([pg // window_ratio for pg in patch_grid])
  662. for index, stage in enumerate(self.layers):
  663. stage_scale = 2 ** max(index - 1, 0)
  664. stage.set_input_size(
  665. feat_size=(patch_grid[0] // stage_scale, patch_grid[1] // stage_scale),
  666. window_size=window_size,
  667. always_partition=always_partition,
  668. )
  669. @torch.jit.ignore
  670. def group_matcher(self, coarse=False):
  671. return dict(
  672. stem=r'^patch_embed', # stem and embed
  673. blocks=r'^layers\.(\d+)' if coarse else [
  674. (r'^layers\.(\d+).downsample', (0,)),
  675. (r'^layers\.(\d+)\.\w+\.(\d+)', None),
  676. (r'^norm', (99999,)),
  677. ]
  678. )
  679. @torch.jit.ignore
  680. def set_grad_checkpointing(self, enable=True):
  681. for l in self.layers:
  682. l.grad_checkpointing = enable
  683. @torch.jit.ignore
  684. def get_classifier(self) -> nn.Module:
  685. return self.head.fc
  686. def reset_classifier(self, num_classes: int, global_pool: Optional[str] = None):
  687. self.num_classes = num_classes
  688. self.head.reset(num_classes, pool_type=global_pool)
  689. def forward_intermediates(
  690. self,
  691. x: torch.Tensor,
  692. indices: Optional[Union[int, List[int]]] = None,
  693. norm: bool = False,
  694. stop_early: bool = False,
  695. output_fmt: str = 'NCHW',
  696. intermediates_only: bool = False,
  697. ) -> Union[List[torch.Tensor], Tuple[torch.Tensor, List[torch.Tensor]]]:
  698. """ Forward features that returns intermediates.
  699. Args:
  700. x: Input image tensor
  701. indices: Take last n blocks if int, all if None, select matching indices if sequence
  702. norm: Apply norm layer to compatible intermediates
  703. stop_early: Stop iterating over blocks when last desired intermediate hit
  704. output_fmt: Shape of intermediate feature outputs
  705. intermediates_only: Only return intermediate features
  706. Returns:
  707. """
  708. assert output_fmt in ('NCHW',), 'Output shape must be NCHW.'
  709. intermediates = []
  710. take_indices, max_index = feature_take_indices(len(self.layers), indices)
  711. # forward pass
  712. x = self.patch_embed(x)
  713. num_stages = len(self.layers)
  714. if torch.jit.is_scripting() or not stop_early: # can't slice blocks in torchscript
  715. stages = self.layers
  716. else:
  717. stages = self.layers[:max_index + 1]
  718. for i, stage in enumerate(stages):
  719. x = stage(x)
  720. if i in take_indices:
  721. if norm and i == num_stages - 1:
  722. x_inter = self.norm(x) # applying final norm last intermediate
  723. else:
  724. x_inter = x
  725. x_inter = x_inter.permute(0, 3, 1, 2).contiguous()
  726. intermediates.append(x_inter)
  727. if intermediates_only:
  728. return intermediates
  729. x = self.norm(x)
  730. return x, intermediates
  731. def prune_intermediate_layers(
  732. self,
  733. indices: Union[int, List[int]] = 1,
  734. prune_norm: bool = False,
  735. prune_head: bool = True,
  736. ):
  737. """ Prune layers not required for specified intermediates.
  738. """
  739. take_indices, max_index = feature_take_indices(len(self.layers), indices)
  740. self.layers = self.layers[:max_index + 1] # truncate blocks
  741. if prune_norm:
  742. self.norm = nn.Identity()
  743. if prune_head:
  744. self.reset_classifier(0, '')
  745. return take_indices
  746. def forward_features(self, x):
  747. x = self.patch_embed(x)
  748. x = self.layers(x)
  749. x = self.norm(x)
  750. return x
  751. def forward_head(self, x, pre_logits: bool = False):
  752. return self.head(x, pre_logits=True) if pre_logits else self.head(x)
  753. def forward(self, x):
  754. x = self.forward_features(x)
  755. x = self.forward_head(x)
  756. return x
  757. def checkpoint_filter_fn(state_dict, model):
  758. """ convert patch embedding weight from manual patchify + linear proj to conv"""
  759. old_weights = True
  760. if 'head.fc.weight' in state_dict:
  761. old_weights = False
  762. import re
  763. out_dict = {}
  764. state_dict = state_dict.get('model', state_dict)
  765. state_dict = state_dict.get('state_dict', state_dict)
  766. for k, v in state_dict.items():
  767. if any([n in k for n in ('relative_position_index', 'attn_mask')]):
  768. continue # skip buffers that should not be persistent
  769. if 'patch_embed.proj.weight' in k:
  770. _, _, H, W = model.patch_embed.proj.weight.shape
  771. if v.shape[-2] != H or v.shape[-1] != W:
  772. v = resample_patch_embed(
  773. v,
  774. (H, W),
  775. interpolation='bicubic',
  776. antialias=True,
  777. verbose=True,
  778. )
  779. if k.endswith('relative_position_bias_table'):
  780. m = model.get_submodule(k[:-29])
  781. if v.shape != m.relative_position_bias_table.shape or m.window_size[0] != m.window_size[1]:
  782. v = resize_rel_pos_bias_table(
  783. v,
  784. new_window_size=m.window_size,
  785. new_bias_shape=m.relative_position_bias_table.shape,
  786. )
  787. if old_weights:
  788. k = re.sub(r'layers.(\d+).downsample', lambda x: f'layers.{int(x.group(1)) + 1}.downsample', k)
  789. k = k.replace('head.', 'head.fc.')
  790. out_dict[k] = v
  791. return out_dict
  792. def _create_swin_transformer(variant, pretrained=False, **kwargs):
  793. default_out_indices = tuple(i for i, _ in enumerate(kwargs.get('depths', (1, 1, 3, 1))))
  794. out_indices = kwargs.pop('out_indices', default_out_indices)
  795. model = build_model_with_cfg(
  796. SwinTransformer, variant, pretrained,
  797. pretrained_filter_fn=checkpoint_filter_fn,
  798. feature_cfg=dict(flatten_sequential=True, out_indices=out_indices),
  799. **kwargs)
  800. return model
  801. def _cfg(url='', **kwargs):
  802. return {
  803. 'url': url,
  804. 'num_classes': 1000, 'input_size': (3, 224, 224), 'pool_size': (7, 7),
  805. 'crop_pct': .9, 'interpolation': 'bicubic', 'fixed_input_size': True,
  806. 'mean': IMAGENET_DEFAULT_MEAN, 'std': IMAGENET_DEFAULT_STD,
  807. 'first_conv': 'patch_embed.proj', 'classifier': 'head.fc',
  808. 'license': 'mit', **kwargs
  809. }
  810. default_cfgs = generate_default_cfgs({
  811. 'swin_small_patch4_window7_224.ms_in22k_ft_in1k': _cfg(
  812. hf_hub_id='timm/',
  813. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_small_patch4_window7_224_22kto1k_finetune.pth', ),
  814. 'swin_base_patch4_window7_224.ms_in22k_ft_in1k': _cfg(
  815. hf_hub_id='timm/',
  816. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224_22kto1k.pth',),
  817. 'swin_base_patch4_window12_384.ms_in22k_ft_in1k': _cfg(
  818. hf_hub_id='timm/',
  819. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384_22kto1k.pth',
  820. input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0),
  821. 'swin_large_patch4_window7_224.ms_in22k_ft_in1k': _cfg(
  822. hf_hub_id='timm/',
  823. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window7_224_22kto1k.pth',),
  824. 'swin_large_patch4_window12_384.ms_in22k_ft_in1k': _cfg(
  825. hf_hub_id='timm/',
  826. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window12_384_22kto1k.pth',
  827. input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0),
  828. 'swin_tiny_patch4_window7_224.ms_in1k': _cfg(
  829. hf_hub_id='timm/',
  830. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_tiny_patch4_window7_224.pth',),
  831. 'swin_small_patch4_window7_224.ms_in1k': _cfg(
  832. hf_hub_id='timm/',
  833. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_small_patch4_window7_224.pth',),
  834. 'swin_base_patch4_window7_224.ms_in1k': _cfg(
  835. hf_hub_id='timm/',
  836. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224.pth',),
  837. 'swin_base_patch4_window12_384.ms_in1k': _cfg(
  838. hf_hub_id='timm/',
  839. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384.pth',
  840. input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0),
  841. # tiny 22k pretrain is worse than 1k, so moved after (untagged priority is based on order)
  842. 'swin_tiny_patch4_window7_224.ms_in22k_ft_in1k': _cfg(
  843. hf_hub_id='timm/',
  844. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_tiny_patch4_window7_224_22kto1k_finetune.pth',),
  845. 'swin_tiny_patch4_window7_224.ms_in22k': _cfg(
  846. hf_hub_id='timm/',
  847. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_tiny_patch4_window7_224_22k.pth',
  848. num_classes=21841),
  849. 'swin_small_patch4_window7_224.ms_in22k': _cfg(
  850. hf_hub_id='timm/',
  851. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.8/swin_small_patch4_window7_224_22k.pth',
  852. num_classes=21841),
  853. 'swin_base_patch4_window7_224.ms_in22k': _cfg(
  854. hf_hub_id='timm/',
  855. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window7_224_22k.pth',
  856. num_classes=21841),
  857. 'swin_base_patch4_window12_384.ms_in22k': _cfg(
  858. hf_hub_id='timm/',
  859. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_base_patch4_window12_384_22k.pth',
  860. input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, num_classes=21841),
  861. 'swin_large_patch4_window7_224.ms_in22k': _cfg(
  862. hf_hub_id='timm/',
  863. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window7_224_22k.pth',
  864. num_classes=21841),
  865. 'swin_large_patch4_window12_384.ms_in22k': _cfg(
  866. hf_hub_id='timm/',
  867. url='https://github.com/SwinTransformer/storage/releases/download/v1.0.0/swin_large_patch4_window12_384_22k.pth',
  868. input_size=(3, 384, 384), pool_size=(12, 12), crop_pct=1.0, num_classes=21841),
  869. 'swin_s3_tiny_224.ms_in1k': _cfg(
  870. hf_hub_id='timm/',
  871. url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/s3_t-1d53f6a8.pth'),
  872. 'swin_s3_small_224.ms_in1k': _cfg(
  873. hf_hub_id='timm/',
  874. url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/s3_s-3bb4c69d.pth'),
  875. 'swin_s3_base_224.ms_in1k': _cfg(
  876. hf_hub_id='timm/',
  877. url='https://github.com/rwightman/pytorch-image-models/releases/download/v0.1-weights/s3_b-a1e95db4.pth'),
  878. })
  879. @register_model
  880. def swin_tiny_patch4_window7_224(pretrained=False, **kwargs) -> SwinTransformer:
  881. """ Swin-T @ 224x224, trained ImageNet-1k
  882. """
  883. model_args = dict(patch_size=4, window_size=7, embed_dim=96, depths=(2, 2, 6, 2), num_heads=(3, 6, 12, 24))
  884. return _create_swin_transformer(
  885. 'swin_tiny_patch4_window7_224', pretrained=pretrained, **dict(model_args, **kwargs))
  886. @register_model
  887. def swin_small_patch4_window7_224(pretrained=False, **kwargs) -> SwinTransformer:
  888. """ Swin-S @ 224x224
  889. """
  890. model_args = dict(patch_size=4, window_size=7, embed_dim=96, depths=(2, 2, 18, 2), num_heads=(3, 6, 12, 24))
  891. return _create_swin_transformer(
  892. 'swin_small_patch4_window7_224', pretrained=pretrained, **dict(model_args, **kwargs))
  893. @register_model
  894. def swin_base_patch4_window7_224(pretrained=False, **kwargs) -> SwinTransformer:
  895. """ Swin-B @ 224x224
  896. """
  897. model_args = dict(patch_size=4, window_size=7, embed_dim=128, depths=(2, 2, 18, 2), num_heads=(4, 8, 16, 32))
  898. return _create_swin_transformer(
  899. 'swin_base_patch4_window7_224', pretrained=pretrained, **dict(model_args, **kwargs))
  900. @register_model
  901. def swin_base_patch4_window12_384(pretrained=False, **kwargs) -> SwinTransformer:
  902. """ Swin-B @ 384x384
  903. """
  904. model_args = dict(patch_size=4, window_size=12, embed_dim=128, depths=(2, 2, 18, 2), num_heads=(4, 8, 16, 32))
  905. return _create_swin_transformer(
  906. 'swin_base_patch4_window12_384', pretrained=pretrained, **dict(model_args, **kwargs))
  907. @register_model
  908. def swin_large_patch4_window7_224(pretrained=False, **kwargs) -> SwinTransformer:
  909. """ Swin-L @ 224x224
  910. """
  911. model_args = dict(patch_size=4, window_size=7, embed_dim=192, depths=(2, 2, 18, 2), num_heads=(6, 12, 24, 48))
  912. return _create_swin_transformer(
  913. 'swin_large_patch4_window7_224', pretrained=pretrained, **dict(model_args, **kwargs))
  914. @register_model
  915. def swin_large_patch4_window12_384(pretrained=False, **kwargs) -> SwinTransformer:
  916. """ Swin-L @ 384x384
  917. """
  918. model_args = dict(patch_size=4, window_size=12, embed_dim=192, depths=(2, 2, 18, 2), num_heads=(6, 12, 24, 48))
  919. return _create_swin_transformer(
  920. 'swin_large_patch4_window12_384', pretrained=pretrained, **dict(model_args, **kwargs))
  921. @register_model
  922. def swin_s3_tiny_224(pretrained=False, **kwargs) -> SwinTransformer:
  923. """ Swin-S3-T @ 224x224, https://arxiv.org/abs/2111.14725
  924. """
  925. model_args = dict(
  926. patch_size=4, window_size=(7, 7, 14, 7), embed_dim=96, depths=(2, 2, 6, 2), num_heads=(3, 6, 12, 24))
  927. return _create_swin_transformer('swin_s3_tiny_224', pretrained=pretrained, **dict(model_args, **kwargs))
  928. @register_model
  929. def swin_s3_small_224(pretrained=False, **kwargs) -> SwinTransformer:
  930. """ Swin-S3-S @ 224x224, https://arxiv.org/abs/2111.14725
  931. """
  932. model_args = dict(
  933. patch_size=4, window_size=(14, 14, 14, 7), embed_dim=96, depths=(2, 2, 18, 2), num_heads=(3, 6, 12, 24))
  934. return _create_swin_transformer('swin_s3_small_224', pretrained=pretrained, **dict(model_args, **kwargs))
  935. @register_model
  936. def swin_s3_base_224(pretrained=False, **kwargs) -> SwinTransformer:
  937. """ Swin-S3-B @ 224x224, https://arxiv.org/abs/2111.14725
  938. """
  939. model_args = dict(
  940. patch_size=4, window_size=(7, 7, 14, 7), embed_dim=96, depths=(2, 2, 30, 2), num_heads=(3, 6, 12, 24))
  941. return _create_swin_transformer('swin_s3_base_224', pretrained=pretrained, **dict(model_args, **kwargs))
  942. register_model_deprecations(__name__, {
  943. 'swin_base_patch4_window7_224_in22k': 'swin_base_patch4_window7_224.ms_in22k',
  944. 'swin_base_patch4_window12_384_in22k': 'swin_base_patch4_window12_384.ms_in22k',
  945. 'swin_large_patch4_window7_224_in22k': 'swin_large_patch4_window7_224.ms_in22k',
  946. 'swin_large_patch4_window12_384_in22k': 'swin_large_patch4_window12_384.ms_in22k',
  947. })

swin_transformer.py at commit 9cfc3f0, no license · at the source

Overview

Authors: Yoon Hwan Byun1,2, Hyunseok Seo3, Jae-Kyung Won4, Boram Lee5, Duk Hyun Hong6, Sun Mo Nam7, Jong Ha Hwang7, Min-Sung Kim7, Yong-Hwy Kim7, Jang Hun Kim6, Mi Ok Yu6, Kyung-Jae Park6, HoJoon Kim3, Sunit Das8, Doo-Sik Kong9, Chul-Kee Park7, Shin-Hyuk Kang6
  1. Department of Neurosurgery, SMG- SNU Boramae Medical Center,Seoul, Republic of Korea
  2. Department of Neurosurgery, Seoul National University College of Medicine,Seoul, Republic of Korea
  3. Department of Artificial Intelligence, Korea University College of Informatics,Seoul, Republic of Korea
  4. Department of Pathology, Seoul National University Hospital, Seoul National University College of Medicine,Seoul, Republic of Korea
  5. Department of Pathology, Samsung Medical Center, Sungkyunkwan University School of Medicine,Seoul, Republic of Korea
  6. Department of Neurosurgery, Korea University Anam Hospital, Korea University College of Medicine,Seoul, Republic of Korea
  7. Department of Neurosurgery, Seoul National University Hospital, Seoul National University College of Medicine,Seoul, Republic of Korea
  8. Division of Neurosurgery, Li Ka Shing Knowledge Institute, St. Michael’s Hospital, University of Toronto,Toronto, Canada
  9. Department of Neurosurgery, Samsung Medical Center, Sungkyunkwan University School of Medicine,Seoul, Republic of Korea
Journal: NPJ digital medicine, volume 9, issue 1, article 471
Dates: received 14 December 2025; accepted 10 April 2026; published online 18 April 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41746-026-02651-0 · PMID 42000880 · PMCID PMC13276397 · OpenAlex W7154848825
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: histology / microscopy (modality), other condition (population), clinical / translational (subfield)
Methods: Statistics, Machine learning
Keywords: Cancer, Computational biology and bioinformatics, Medical research, Oncology
Topic: Glioma Diagnosis and Treatment (Genetics, Medicine), according to OpenAlex
Funding: Ministry of Science and ICT (RS202400338025); Ministry of Health and Welfare (RS-2022-KH129293); Korea University Anam Hospital (O241376); Technology Development Program of the Ministry of SMEs and Startups (RS-2024-00487593)
Citations: not cited yet (Europe PMC); 50 references in the paper
Research resources: RRID:SCR_008567

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.

Repository

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

ImpelKorea/AI-agumented-CLE-Imaging-npj-Digital-Medicine

License: none: the authors keep all their rights
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: 9cfc3f0914b0ad7fae4a9a7ac432c671198b0cc3, 15 March 2026
Languages: Python (251)
Size: 280 files, 251 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, environment (environment.yml)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (206 files), NumPy (20 files), Pillow (9 files), TensorFlow (2 files), Matplotlib (1 file), pandas (1 file), SciPy (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
252 files

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41746-026-02651-0.

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;
  • 251 scripts, each with its path and the digest of its content;
  • 2 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data availability statement

The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

  • it says that the data are available on request

Read it in the paper: doi.org/10.1038/s41746-026-02651-0.

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, 29 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 17 authors, 4 keywords, 4 funders, 42 references, 1 RRID.

Cite

This paper

Byun, Y. H., Seo, H., Won, J.-K., Lee, B., Hong, D. H., Nam, S. M., Hwang, J. H., Kim, M.-S., Kim, Y.-H., Kim, J. H., Yu, M. O., Park, K.-J., Kim, H., Das, S., Kong, D.-S., Park, C.-K., & Kang, S.-H. (2026). AI Augmented Confocal Laser Endomicroscopy for Rapid Intraoperative Diagnosis of Brain Tumors. NPJ digital medicine, 9(1), 471. https://doi.org/10.1038/s41746-026-02651-0

BibTeX

@article{byun2026ai,
author = {Byun, Yoon Hwan and Seo, Hyunseok and Won, Jae-Kyung and Lee, Boram and Hong, Duk Hyun and Nam, Sun Mo and Hwang, Jong Ha and Kim, Min-Sung and Kim, Yong-Hwy and Kim, Jang Hun and Yu, Mi Ok and Park, Kyung-Jae and Kim, HoJoon and Das, Sunit and Kong, Doo-Sik and Park, Chul-Kee and Kang, Shin-Hyuk},
title = {{AI Augmented Confocal Laser Endomicroscopy for Rapid Intraoperative Diagnosis of Brain Tumors}},
journal = {NPJ digital medicine},
year = {2026},
month = apr,
volume = {9},
number = {1},
pages = {471},
publisher = {Nature Publishing Group},
issn = {2398-6352},
doi = {10.1038/s41746-026-02651-0},
url = {https://doi.org/10.1038/s41746-026-02651-0},
pmid = {42000880},
pmcid = {PMC13276397}
}

RIS

TY - JOUR
AU - Byun, Yoon Hwan
AU - Seo, Hyunseok
AU - Won, Jae-Kyung
AU - Lee, Boram
AU - Hong, Duk Hyun
AU - Nam, Sun Mo
AU - Hwang, Jong Ha
AU - Kim, Min-Sung
AU - Kim, Yong-Hwy
AU - Kim, Jang Hun
AU - Yu, Mi Ok
AU - Park, Kyung-Jae
AU - Kim, HoJoon
AU - Das, Sunit
AU - Kong, Doo-Sik
AU - Park, Chul-Kee
AU - Kang, Shin-Hyuk
TI - AI Augmented Confocal Laser Endomicroscopy for Rapid Intraoperative Diagnosis of Brain Tumors
T2 - NPJ digital medicine
J2 - NPJ Digit Med
PY - 2026
DA - 2026/04/18
VL - 9
IS - 1
SP - 471
SN - 2398-6352
PB - Nature Publishing Group
DO - 10.1038/s41746-026-02651-0
UR - https://doi.org/10.1038/s41746-026-02651-0
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41746-026-02651-0",
"type": "article-journal",
"title": "AI Augmented Confocal Laser Endomicroscopy for Rapid Intraoperative Diagnosis of Brain Tumors",
"container-title": "NPJ digital medicine",
"author": [
{
"family": "Byun",
"given": "Yoon Hwan"
},
{
"family": "Seo",
"given": "Hyunseok"
},
{
"family": "Won",
"given": "Jae-Kyung"
},
{
"family": "Lee",
"given": "Boram"
},
{
"family": "Hong",
"given": "Duk Hyun"
},
{
"family": "Nam",
"given": "Sun Mo"
},
{
"family": "Hwang",
"given": "Jong Ha"
},
{
"family": "Kim",
"given": "Min-Sung"
},
{
"family": "Kim",
"given": "Yong-Hwy"
},
{
"family": "Kim",
"given": "Jang Hun"
},
{
"family": "Yu",
"given": "Mi Ok"
},
{
"family": "Park",
"given": "Kyung-Jae"
},
{
"family": "Kim",
"given": "HoJoon"
},
{
"family": "Das",
"given": "Sunit"
},
{
"family": "Kong",
"given": "Doo-Sik"
},
{
"family": "Park",
"given": "Chul-Kee"
},
{
"family": "Kang",
"given": "Shin-Hyuk"
}
],
"container-title-short": "NPJ Digit Med",
"volume": "9",
"issue": "1",
"page": "471",
"DOI": "10.1038/s41746-026-02651-0",
"PMID": "42000880",
"PMCID": "PMC13276397",
"ISSN": "2398-6352",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41746-026-02651-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
18
]
]
}
}

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.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: TensorFlow, Pillow, PyTorch, 4 other tools, clinical / translational, 1 reference
[2] doi:10.1093/jnen/nlaf152 [code]
Clinical and pathologic correlations of machine learning quantification of Aβ deposits across 3 brain regions of decedents with Alzheimer disease.
Journal: Journal of neuropathology and experimental neurology
In common: TensorFlow, Pillow, PyTorch, 4 other tools, clinical / translational
[3] doi:10.7554/elife.110588 [code]
Opening the black box toward a modular approach to spike sorting.
Journal: eLife
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[4] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[5] doi:10.3390/s26175501 [code]
Evaluation of a Hybrid Neural-Polynomial Deep Q-Network for Switching-Aware Spectrum Selection in a Controlled Radio-Frequency Measurement-Replay Testbed.
Journal: Sensors (Basel, Switzerland)
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[6] doi:10.1126/sciadv.aed3650 [code]
Truthful visualizations for mass spectrometry imaging enable high-spatial-resolution interactive &lt;i&gt;m/z&lt;/i&gt; mapping and exploration.
Journal: Science advances
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[7] doi:10.1038/s41592-026-03194-8 [code]
Beyond benchmarking: an expert-guided consensus approach to spatially aware clustering.
Journal: Nature methods
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[8] doi:10.1371/journal.pcbi.1014571 [code]
SynAPSeg: A novel dataset and image analysis framework for deep learning-based synapse detection and quantification.
Journal: PLoS computational biology
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[9] doi:10.1364/boe.605322 [code]
Generalized plaque digitization framework for multi-dimensional mesoscopic images.
Journal: Biomedical optics express
In common: TensorFlow, Pillow, PyTorch, 4 other tools
[10] doi:10.1523/eneuro.0023-26.2026 [code]
Real-Time Segmentation and Classification of Birdsong Syllables for Learning Experiments.
Journal: eNeuro
In common: TensorFlow, Pillow, PyTorch, 4 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.

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.