OSCR

Deep learning for early detection of cerebral small vessel disease using self-supervised graph embeddings and retinal image analysis.

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] § Detection methodology › Phase 2: global transformer-based encoding ↔ rnv_t_empirical_validation.py, lines 228–249 · score 0.55 · Multi Head Correlative, MHCA
  2. [2] § Materials utilized › Experimental components ↔ rnv_t_empirical_validation.py, lines 62–71 · score 0.52 · top hat, Gaussian, CLAHE, preprocessing, vessel, retinal

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 · 683 lines · 27 KB · GPL-3.0 · 2 matches

  1. """
  2. RNV-T empirical validation reference implementation.
  3. This script implements the five phases described in the manuscript:
  4. 1) Local Vascular Extraction (LVE)
  5. 2) Global Transformer-based Encoding (GTE)
  6. 3) Graph-based Convolutional Attention Network (G-CAN)
  7. 4) Local-Global Attention Fusion (LGAF)
  8. 5) Optimized training with SSL + federated-style aggregation
  9. It is designed as a reproducible empirical validation scaffold rather than the
  10. original laboratory code. Dataset paths, labels, and site partitions can be
  11. plugged in through the config section / CLI.
  12. """
  13. from __future__ import annotations
  14. import argparse
  15. import math
  16. import os
  17. import random
  18. from dataclasses import dataclass
  19. from pathlib import Path
  20. from typing import Dict, Iterable, List, Optional, Sequence, Tuple
  21. import cv2
  22. import numpy as np
  23. import torch
  24. import torch.nn as nn
  25. import torch.nn.functional as F
  26. from torch.utils.data import DataLoader, Dataset
  27. # -----------------------------------------------------------------------------
  28. # Reproducibility helpers
  29. # -----------------------------------------------------------------------------
  30. def seed_everything(seed: int = 42) -> None:
  31. random.seed(seed)
  32. np.random.seed(seed)
  33. torch.manual_seed(seed)
  34. torch.cuda.manual_seed_all(seed)
  35. # -----------------------------------------------------------------------------
  36. # Preprocessing utilities
  37. # -----------------------------------------------------------------------------
  38. def apply_clahe_rgb(image: np.ndarray) -> np.ndarray:
  39. lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)
  40. l, a, b = cv2.split(lab)
  41. clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
  42. l = clahe.apply(l)
  43. out = cv2.merge([l, a, b])
  44. return cv2.cvtColor(out, cv2.COLOR_LAB2RGB)
  45. def top_hat_enhance(gray: np.ndarray, kernel_size: int = 15) -> np.ndarray:
  46. kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size))
  47. return cv2.morphologyEx(gray, cv2.MORPH_TOPHAT, kernel)
  48. def preprocess_retinal_image(image: np.ndarray, image_size: int = 224) -> np.ndarray:
  49. image = cv2.resize(image, (image_size, image_size), interpolation=cv2.INTER_AREA)
  50. image = apply_clahe_rgb(image)
  51. image = cv2.GaussianBlur(image, (3, 3), 0)
  52. gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
  53. vessel_boost = top_hat_enhance(gray, kernel_size=11)
  54. vessel_boost = np.stack([vessel_boost] * 3, axis=-1)
  55. image = cv2.addWeighted(image, 0.85, vessel_boost, 0.15, 0.0)
  56. image = image.astype(np.float32) / 255.0
  57. return image
  58. def load_rgb(path: str, image_size: int = 224) -> np.ndarray:
  59. img = cv2.imread(path, cv2.IMREAD_COLOR)
  60. if img is None:
  61. raise FileNotFoundError(path)
  62. img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
  63. return preprocess_retinal_image(img, image_size=image_size)
  64. # -----------------------------------------------------------------------------
  65. # Dataset
  66. # -----------------------------------------------------------------------------
  67. @dataclass
  68. class SampleRecord:
  69. fundus: str
  70. octa: str
  71. fa: str
  72. label: int
  73. site: str = "site_0"
  74. class MultiModalRetinalDataset(Dataset):
  75. def __init__(self, records: Sequence[SampleRecord], image_size: int = 224, augment: bool = False):
  76. self.records = list(records)
  77. self.image_size = image_size
  78. self.augment = augment
  79. def __len__(self) -> int:
  80. return len(self.records)
  81. def _augment(self, x: np.ndarray) -> np.ndarray:
  82. if not self.augment:
  83. return x
  84. if random.random() < 0.5:
  85. x = np.flip(x, axis=1).copy()
  86. if random.random() < 0.5:
  87. x = np.flip(x, axis=0).copy()
  88. if random.random() < 0.2:
  89. noise = np.random.normal(0, 0.02, x.shape).astype(np.float32)
  90. x = np.clip(x + noise, 0.0, 1.0)
  91. return x
  92. def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
  93. rec = self.records[idx]
  94. fundus = self._augment(load_rgb(rec.fundus, self.image_size))
  95. octa = self._augment(load_rgb(rec.octa, self.image_size))
  96. fa = self._augment(load_rgb(rec.fa, self.image_size))
  97. return {
  98. "fundus": torch.from_numpy(fundus).permute(2, 0, 1),
  99. "octa": torch.from_numpy(octa).permute(2, 0, 1),
  100. "fa": torch.from_numpy(fa).permute(2, 0, 1),
  101. "label": torch.tensor(rec.label, dtype=torch.long),
  102. "site": rec.site,
  103. }
  104. # -----------------------------------------------------------------------------
  105. # Phase 1: Local Vascular Extraction (LVE)
  106. # -----------------------------------------------------------------------------
  107. class ConvBlock(nn.Module):
  108. def __init__(self, in_ch: int, out_ch: int):
  109. super().__init__()
  110. self.block = nn.Sequential(
  111. nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1),
  112. nn.BatchNorm2d(out_ch),
  113. nn.ReLU(inplace=True),
  114. nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1),
  115. nn.BatchNorm2d(out_ch),
  116. nn.ReLU(inplace=True),
  117. )
  118. def forward(self, x: torch.Tensor) -> torch.Tensor:
  119. return self.block(x)
  120. class LVEBranch(nn.Module):
  121. def __init__(self, in_ch: int = 3, widths: Sequence[int] = (32, 64, 128, 256), emb_dim: int = 256):
  122. super().__init__()
  123. self.blocks = nn.ModuleList()
  124. prev = in_ch
  125. for w in widths:
  126. self.blocks.append(ConvBlock(prev, w))
  127. prev = w
  128. self.pool = nn.MaxPool2d(2)
  129. self.project = nn.Sequential(
  130. nn.Linear(sum(widths), emb_dim),
  131. nn.ReLU(inplace=True),
  132. nn.Linear(emb_dim, emb_dim),
  133. )
  134. def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, List[torch.Tensor]]:
  135. pyramids: List[torch.Tensor] = []
  136. h = x
  137. pooled_feats = []
  138. for i, block in enumerate(self.blocks):
  139. h = block(h)
  140. pyramids.append(h)
  141. pooled_feats.append(F.adaptive_avg_pool2d(h, 1).flatten(1))
  142. if i != len(self.blocks) - 1:
  143. h = self.pool(h)
  144. local_vec = self.project(torch.cat(pooled_feats, dim=1))
  145. return local_vec, pyramids
  146. class LVEEncoder(nn.Module):
  147. def __init__(self, emb_dim: int = 256):
  148. super().__init__()
  149. self.fundus_branch = LVEBranch(3, emb_dim=emb_dim)
  150. self.octa_branch = LVEBranch(3, emb_dim=emb_dim)
  151. self.fa_branch = LVEBranch(3, emb_dim=emb_dim)
  152. self.cross_modal = nn.Sequential(
  153. nn.Linear(emb_dim * 3, emb_dim * 2),
  154. nn.ReLU(inplace=True),
  155. nn.Linear(emb_dim * 2, emb_dim),
  156. )
  157. def forward(self, fundus: torch.Tensor, octa: torch.Tensor, fa: torch.Tensor) -> Dict[str, torch.Tensor | List[torch.Tensor]]:
  158. f_vec, f_maps = self.fundus_branch(fundus)
  159. o_vec, o_maps = self.octa_branch(octa)
  160. a_vec, a_maps = self.fa_branch(fa)
  161. hadamard_fo = f_vec * o_vec
  162. hadamard_fa = f_vec * a_vec
  163. hadamard_oa = o_vec * a_vec
  164. fused = self.cross_modal(torch.cat([hadamard_fo, hadamard_fa, hadamard_oa], dim=1))
  165. return {
  166. "local_vec": fused,
  167. "fundus_vec": f_vec,
  168. "octa_vec": o_vec,
  169. "fa_vec": a_vec,
  170. "fundus_maps": f_maps,
  171. "octa_maps": o_maps,
  172. "fa_maps": a_maps,
  173. }
  174. # -----------------------------------------------------------------------------
  175. # Phase 2: Global Transformer-based Encoding (GTE)
  176. # -----------------------------------------------------------------------------
  177. class PatchEmbed(nn.Module):
  178. def __init__(self, img_size: int = 224, patch_size: int = 16, in_chans: int = 3, embed_dim: int = 384):
  179. super().__init__()
  180. self.grid_size = img_size // patch_size
  181. self.num_patches = self.grid_size * self.grid_size
  182. self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
  183. def forward(self, x: torch.Tensor) -> torch.Tensor:
  184. x = self.proj(x) # [B, D, H', W']
  185. return x.flatten(2).transpose(1, 2) # [B, N, D]
  186. class MHCA(nn.Module):
  187. """Multi-head correlative attention."""
  188. def __init__(self, dim: int, heads: int = 8, dropout: float = 0.1):
  189. super().__init__()
  190. self.heads = heads
  191. self.scale = (dim // heads) ** -0.5
  192. self.qkv = nn.Linear(dim, dim * 3)
  193. self.corr = nn.Parameter(torch.zeros(heads, 1, 1))
  194. self.proj = nn.Linear(dim, dim)
  195. self.drop = nn.Dropout(dropout)
  196. def forward(self, x: torch.Tensor) -> torch.Tensor:
  197. b, n, c = x.shape
  198. qkv = self.qkv(x).reshape(b, n, 3, self.heads, c // self.heads).permute(2, 0, 3, 1, 4)
  199. q, k, v = qkv[0], qkv[1], qkv[2]
  200. sim = (q @ k.transpose(-2, -1)) * self.scale
  201. div = -torch.abs(q.unsqueeze(-2) - k.unsqueeze(-3)).mean(dim=-1)
  202. attn = sim + self.corr + 0.1 * div
  203. attn = attn.softmax(dim=-1)
  204. attn = self.drop(attn)
  205. out = (attn @ v).transpose(1, 2).reshape(b, n, c)
  206. return self.proj(out)
  207. class TransformerBlock(nn.Module):
  208. def __init__(self, dim: int, heads: int = 8, mlp_ratio: float = 4.0, dropout: float = 0.1):
  209. super().__init__()
  210. self.norm1 = nn.LayerNorm(dim)
  211. self.attn = MHCA(dim, heads=heads, dropout=dropout)
  212. self.norm2 = nn.LayerNorm(dim)
  213. self.mlp = nn.Sequential(
  214. nn.Linear(dim, int(dim * mlp_ratio)),
  215. nn.GELU(),
  216. nn.Dropout(dropout),
  217. nn.Linear(int(dim * mlp_ratio), dim),
  218. nn.Dropout(dropout),
  219. )
  220. def forward(self, x: torch.Tensor) -> torch.Tensor:
  221. x = x + self.attn(self.norm1(x))
  222. x = x + self.mlp(self.norm2(x))
  223. return x
  224. class GTEEncoder(nn.Module):
  225. def __init__(self, img_size: int = 224, patch_size: int = 16, dim: int = 384, depth: int = 4, heads: int = 8):
  226. super().__init__()
  227. self.patch_embed = PatchEmbed(img_size, patch_size, 9, dim)
  228. num_patches = self.patch_embed.num_patches
  229. self.cls_token = nn.Parameter(torch.zeros(1, 1, dim))
  230. self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, dim))
  231. self.blocks = nn.ModuleList([TransformerBlock(dim, heads=heads) for _ in range(depth)])
  232. self.norm = nn.LayerNorm(dim)
  233. def forward(self, fused_modal_image: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
  234. x = self.patch_embed(fused_modal_image)
  235. b = x.shape[0]
  236. cls = self.cls_token.expand(b, -1, -1)
  237. x = torch.cat([cls, x], dim=1) + self.pos_embed
  238. for block in self.blocks:
  239. x = block(x)
  240. x = self.norm(x)
  241. return x[:, 0], x[:, 1:]
  242. # -----------------------------------------------------------------------------
  243. # Phase 3: Graph-based Convolutional Attention Network (G-CAN)
  244. # -----------------------------------------------------------------------------
  245. class GraphAttentionLayer(nn.Module):
  246. def __init__(self, in_dim: int, out_dim: int, dropout: float = 0.2, alpha: float = 0.2):
  247. super().__init__()
  248. self.lin = nn.Linear(in_dim, out_dim, bias=False)
  249. self.attn = nn.Linear(2 * out_dim, 1, bias=False)
  250. self.dropout = nn.Dropout(dropout)
  251. self.leaky_relu = nn.LeakyReLU(alpha)
  252. def forward(self, x: torch.Tensor, adj: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
  253. h = self.lin(x)
  254. b, n, d = h.shape
  255. h_i = h.unsqueeze(2).expand(b, n, n, d)
  256. h_j = h.unsqueeze(1).expand(b, n, n, d)
  257. e = self.leaky_relu(self.attn(torch.cat([h_i, h_j], dim=-1)).squeeze(-1))
  258. e = e.masked_fill(adj <= 0, float("-inf"))
  259. alpha = F.softmax(e, dim=-1)
  260. alpha = self.dropout(alpha)
  261. out = alpha @ h
  262. return F.elu(out), alpha
  263. class GCANEncoder(nn.Module):
  264. def __init__(self, node_dim: int, hidden_dim: int = 128, num_layers: int = 3, dropout: float = 0.2):
  265. super().__init__()
  266. dims = [node_dim] + [hidden_dim] * num_layers
  267. self.layers = nn.ModuleList([
  268. GraphAttentionLayer(dims[i], dims[i + 1], dropout=dropout) for i in range(num_layers)
  269. ])
  270. self.out_proj = nn.Linear(hidden_dim + 1, hidden_dim)
  271. @staticmethod
  272. def build_graph(node_feats: torch.Tensor, coords: torch.Tensor, sigma: float = 0.35) -> torch.Tensor:
  273. # coords: [B, N, 2], normalized to [0,1]
  274. dists = torch.cdist(coords, coords, p=2)
  275. adj = torch.exp(-(dists ** 2) / (2 * sigma * sigma))
  276. eye = torch.eye(adj.size(-1), device=adj.device).unsqueeze(0)
  277. adj = torch.maximum(adj, eye)
  278. return adj
  279. def forward(self, node_feats: torch.Tensor, coords: torch.Tensor) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
  280. adj = self.build_graph(node_feats, coords)
  281. h = node_feats
  282. attn_maps = []
  283. for layer in self.layers:
  284. h, alpha = layer(h, adj)
  285. attn_maps.append(alpha)
  286. predicted_adj = torch.sigmoid(torch.matmul(F.normalize(h, dim=-1), F.normalize(h, dim=-1).transpose(-1, -2)))
  287. anomaly = torch.mean(torch.abs(adj - predicted_adj), dim=(-1, -2), keepdim=True)
  288. graph_repr = h.mean(dim=1)
  289. graph_repr = self.out_proj(torch.cat([graph_repr, anomaly], dim=-1))
  290. return graph_repr, {
  291. "adj": adj,
  292. "predicted_adj": predicted_adj,
  293. "anomaly": anomaly,
  294. "attn_maps": attn_maps[-1],
  295. }
  296. # -----------------------------------------------------------------------------
  297. # Phase 4: Local-Global Attention Fusion (LGAF)
  298. # -----------------------------------------------------------------------------
  299. class LGAF(nn.Module):
  300. def __init__(self, local_dim: int, global_dim: int, graph_dim: int, hidden_dim: int = 512, num_classes: int = 2):
  301. super().__init__()
  302. self.local_proj = nn.Linear(local_dim, hidden_dim)
  303. self.global_proj = nn.Linear(global_dim, hidden_dim)
  304. self.graph_proj = nn.Linear(graph_dim, hidden_dim)
  305. self.gate = nn.Sequential(
  306. nn.Linear(hidden_dim * 2, hidden_dim),
  307. nn.ReLU(inplace=True),
  308. nn.Linear(hidden_dim, hidden_dim),
  309. nn.Sigmoid(),
  310. )
  311. self.attn = nn.MultiheadAttention(hidden_dim, num_heads=4, batch_first=True)
  312. self.norm = nn.LayerNorm(hidden_dim)
  313. self.classifier = nn.Sequential(
  314. nn.Linear(hidden_dim, hidden_dim // 2),
  315. nn.ReLU(inplace=True),
  316. nn.Dropout(0.2),
  317. nn.Linear(hidden_dim // 2, num_classes),
  318. )
  319. def forward(self, local_vec: torch.Tensor, global_vec: torch.Tensor, graph_vec: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
  320. l = self.local_proj(local_vec)
  321. g = self.global_proj(global_vec)
  322. r = self.graph_proj(graph_vec)
  323. gate = self.gate(torch.cat([l, g], dim=-1))
  324. fused = gate * l + (1.0 - gate) * g
  325. tokens = torch.stack([fused, r], dim=1)
  326. attended, _ = self.attn(tokens, tokens, tokens)
  327. fused = self.norm(attended.mean(dim=1) + fused + 0.5 * r)
  328. logits = self.classifier(fused)
  329. return logits, fused
  330. # -----------------------------------------------------------------------------
  331. # Full RNV-T model
  332. # -----------------------------------------------------------------------------
  333. class RNVT(nn.Module):
  334. def __init__(self, image_size: int = 224, local_dim: int = 256, global_dim: int = 384, graph_dim: int = 128, num_classes: int = 2):
  335. super().__init__()
  336. self.lve = LVEEncoder(emb_dim=local_dim)
  337. self.gte = GTEEncoder(img_size=image_size, patch_size=16, dim=global_dim, depth=4, heads=8)
  338. self.gcan = GCANEncoder(node_dim=global_dim, hidden_dim=graph_dim, num_layers=3, dropout=0.2)
  339. self.lgaf = LGAF(local_dim=local_dim, global_dim=global_dim, graph_dim=graph_dim, hidden_dim=512, num_classes=num_classes)
  340. self.proj_ssl = nn.Sequential(
  341. nn.Linear(512, 256), nn.ReLU(inplace=True), nn.Linear(256, 128)
  342. )
  343. @staticmethod
  344. def _token_coords(num_tokens: int, device: torch.device) -> torch.Tensor:
  345. side = int(math.sqrt(num_tokens))
  346. ys, xs = torch.meshgrid(
  347. torch.linspace(0, 1, side, device=device),
  348. torch.linspace(0, 1, side, device=device),
  349. indexing="ij",
  350. )
  351. coords = torch.stack([xs.reshape(-1), ys.reshape(-1)], dim=-1)
  352. return coords
  353. def forward(self, fundus: torch.Tensor, octa: torch.Tensor, fa: torch.Tensor) -> Dict[str, torch.Tensor]:
  354. lve_out = self.lve(fundus, octa, fa)
  355. fused_modal_image = torch.cat([fundus, octa, fa], dim=1)
  356. global_vec, patch_tokens = self.gte(fused_modal_image)
  357. b, n, d = patch_tokens.shape
  358. coords = self._token_coords(n, patch_tokens.device).unsqueeze(0).repeat(b, 1, 1)
  359. graph_vec, graph_aux = self.gcan(patch_tokens, coords)
  360. logits, fused_repr = self.lgaf(lve_out["local_vec"], global_vec, graph_vec)
  361. return {
  362. "logits": logits,
  363. "local_vec": lve_out["local_vec"],
  364. "global_vec": global_vec,
  365. "graph_vec": graph_vec,
  366. "fused_repr": fused_repr,
  367. "ssl_proj": F.normalize(self.proj_ssl(fused_repr), dim=-1),
  368. "graph_anomaly": graph_aux["anomaly"].squeeze(-1),
  369. "adj": graph_aux["adj"],
  370. "predicted_adj": graph_aux["predicted_adj"],
  371. }
  372. # -----------------------------------------------------------------------------
  373. # Phase 5: losses and optimization
  374. # -----------------------------------------------------------------------------
  375. def info_nce(z1: torch.Tensor, z2: torch.Tensor, temperature: float = 0.2) -> torch.Tensor:
  376. b = z1.size(0)
  377. z = torch.cat([z1, z2], dim=0)
  378. sim = torch.mm(z, z.t()) / temperature
  379. mask = torch.eye(2 * b, device=z.device, dtype=torch.bool)
  380. sim = sim.masked_fill(mask, -1e9)
  381. targets = torch.arange(b, device=z.device)
  382. targets = torch.cat([targets + b, targets], dim=0)
  383. return F.cross_entropy(sim, targets)
  384. def graph_regularization(adj: torch.Tensor, pred_adj: torch.Tensor) -> torch.Tensor:
  385. return F.mse_loss(pred_adj, adj)
  386. def total_loss_fn(outputs_1: Dict[str, torch.Tensor], outputs_2: Dict[str, torch.Tensor], labels: torch.Tensor,
  387. lambda_ssl: float = 0.2, lambda_reg: float = 0.1) -> Tuple[torch.Tensor, Dict[str, float]]:
  388. ce = F.cross_entropy(outputs_1["logits"], labels)
  389. ssl = info_nce(outputs_1["ssl_proj"], outputs_2["ssl_proj"])
  390. reg = graph_regularization(outputs_1["adj"], outputs_1["predicted_adj"])
  391. total = ce + lambda_ssl * ssl + lambda_reg * reg
  392. return total, {
  393. "ce": float(ce.detach().cpu()),
  394. "ssl": float(ssl.detach().cpu()),
  395. "reg": float(reg.detach().cpu()),
  396. "total": float(total.detach().cpu()),
  397. }
  398. # -----------------------------------------------------------------------------
  399. # Utilities for federated-style training
  400. # -----------------------------------------------------------------------------
  401. def clone_state_dict(model: nn.Module) -> Dict[str, torch.Tensor]:
  402. return {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}
  403. def load_state_dict_strict(model: nn.Module, state: Dict[str, torch.Tensor]) -> None:
  404. model.load_state_dict(state, strict=True)
  405. def fedavg(states: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
  406. avg = {}
  407. for key in states[0].keys():
  408. avg[key] = sum(s[key] for s in states) / len(states)
  409. return avg
  410. def accuracy_from_logits(logits: torch.Tensor, labels: torch.Tensor) -> float:
  411. preds = logits.argmax(dim=1)
  412. return float((preds == labels).float().mean().item())
  413. def sensitivity_specificity(logits: torch.Tensor, labels: torch.Tensor) -> Tuple[float, float]:
  414. preds = logits.argmax(dim=1)
  415. tp = ((preds == 1) & (labels == 1)).sum().item()
  416. tn = ((preds == 0) & (labels == 0)).sum().item()
  417. fp = ((preds == 1) & (labels == 0)).sum().item()
  418. fn = ((preds == 0) & (labels == 1)).sum().item()
  419. sensitivity = tp / max(tp + fn, 1)
  420. specificity = tn / max(tn + fp, 1)
  421. return sensitivity, specificity
  422. # -----------------------------------------------------------------------------
  423. # Demo generation for validation when real data is unavailable
  424. # -----------------------------------------------------------------------------
  425. def make_multimodal_dataset(root: str, n: int = 48, image_size: int = 224) -> List[SampleRecord]:
  426. os.makedirs(root, exist_ok=True)
  427. records: List[SampleRecord] = []
  428. for i in range(n):
  429. label = i % 2
  430. site = f"site_{i % 3}"
  431. paths = {}
  432. for modality in ["fundus", "octa", "fa"]:
  433. canvas = np.zeros((image_size, image_size, 3), dtype=np.uint8)
  434. # background tone
  435. canvas[:] = (20 + 15 * label, 30 + 5 * (i % 5), 20 + 10 * (i % 3))
  436. # vessel-like line patterns
  437. center = (image_size // 2, image_size // 2)
  438. num_lines = 18 + 6 * label
  439. thickness = 1 + label
  440. for a in np.linspace(0, 2 * np.pi, num_lines, endpoint=False):
  441. r = image_size // 2 - 8
  442. end = (int(center[0] + r * np.cos(a)), int(center[1] + r * np.sin(a)))
  443. cv2.line(canvas, center, end, (180, 180, 180), thickness)
  444. if modality == "octa" and label == 1:
  445. cv2.circle(canvas, center, 26, (10, 10, 10), -1) # FAZ enlargement cue
  446. if modality == "fa" and label == 1:
  447. cv2.circle(canvas, (image_size // 3, image_size // 3), 14, (230, 230, 230), -1)
  448. if modality == "fundus" and label == 1:
  449. cv2.circle(canvas, (2 * image_size // 3, image_size // 2), 10, (220, 50, 50), -1)
  450. p = os.path.join(root, f"{modality}_{i}.png")
  451. cv2.imwrite(p, cv2.cvtColor(canvas, cv2.COLOR_RGB2BGR))
  452. paths[modality] = p
  453. records.append(SampleRecord(paths["fundus"], paths["octa"], paths["fa"], label, site))
  454. return records
  455. # -----------------------------------------------------------------------------
  456. # Training / evaluation
  457. # -----------------------------------------------------------------------------
  458. def build_site_loaders(records: Sequence[SampleRecord], image_size: int, batch_size: int) -> Dict[str, DataLoader]:
  459. sites: Dict[str, List[SampleRecord]] = {}
  460. for rec in records:
  461. sites.setdefault(rec.site, []).append(rec)
  462. loaders = {}
  463. for site, site_records in sites.items():
  464. ds = MultiModalRetinalDataset(site_records, image_size=image_size, augment=True)
  465. loaders[site] = DataLoader(ds, batch_size=batch_size, shuffle=True, num_workers=0)
  466. return loaders
  467. def make_augmented_views(batch: Dict[str, torch.Tensor]) -> Tuple[Dict[str, torch.Tensor], Dict[str, torch.Tensor]]:
  468. def jitter(x: torch.Tensor) -> torch.Tensor:
  469. noise = 0.03 * torch.randn_like(x)
  470. return torch.clamp(x + noise, 0.0, 1.0)
  471. v1 = {k: jitter(v) if k in ["fundus", "octa", "fa"] else v for k, v in batch.items()}
  472. v2 = {k: jitter(v) if k in ["fundus", "octa", "fa"] else v for k, v in batch.items()}
  473. return v1, v2
  474. def run_local_epoch(model: RNVT, loader: DataLoader, device: torch.device, lr: float = 1e-4) -> Tuple[Dict[str, torch.Tensor], Dict[str, float]]:
  475. model.train()
  476. optimizer = torch.optim.Adam(model.parameters(), lr=lr)
  477. meters = {"loss": 0.0, "acc": 0.0, "count": 0}
  478. for batch in loader:
  479. batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in batch.items()}
  480. v1, v2 = make_augmented_views(batch)
  481. out1 = model(v1["fundus"], v1["octa"], v1["fa"])
  482. out2 = model(v2["fundus"], v2["octa"], v2["fa"])
  483. loss, _ = total_loss_fn(out1, out2, batch["label"])
  484. optimizer.zero_grad()
  485. loss.backward()
  486. optimizer.step()
  487. meters["loss"] += loss.item() * batch["label"].size(0)
  488. meters["acc"] += accuracy_from_logits(out1["logits"], batch["label"]) * batch["label"].size(0)
  489. meters["count"] += batch["label"].size(0)
  490. stats = {
  491. "loss": meters["loss"] / max(meters["count"], 1),
  492. "acc": meters["acc"] / max(meters["count"], 1),
  493. }
  494. return clone_state_dict(model), stats
  495. @torch.no_grad()
  496. def evaluate(model: RNVT, loader: DataLoader, device: torch.device) -> Dict[str, float]:
  497. model.eval()
  498. all_logits, all_labels = [], []
  499. for batch in loader:
  500. fundus = batch["fundus"].to(device)
  501. octa = batch["octa"].to(device)
  502. fa = batch["fa"].to(device)
  503. labels = batch["label"].to(device)
  504. out = model(fundus, octa, fa)
  505. all_logits.append(out["logits"])
  506. all_labels.append(labels)
  507. logits = torch.cat(all_logits, dim=0)
  508. labels = torch.cat(all_labels, dim=0)
  509. acc = accuracy_from_logits(logits, labels)
  510. sens, spec = sensitivity_specificity(logits, labels)
  511. return {"accuracy": acc, "sensitivity": sens, "specificity": spec}
  512. def federated_train(records: Sequence[SampleRecord], image_size: int = 224, batch_size: int = 8, rounds: int = 3,
  513. lr: float = 1e-4, device: str = "cpu") -> Dict[str, float]:
  514. device_t = torch.device(device)
  515. loaders = build_site_loaders(records, image_size=image_size, batch_size=batch_size)
  516. global_model = RNVT(image_size=image_size).to(device_t)
  517. global_state = clone_state_dict(global_model)
  518. for rnd in range(1, rounds + 1):
  519. local_states = []
  520. print(f"[Federated Round {rnd}/{rounds}]")
  521. for site, loader in loaders.items():
  522. local_model = RNVT(image_size=image_size).to(device_t)
  523. load_state_dict_strict(local_model, global_state)
  524. state, stats = run_local_epoch(local_model, loader, device_t, lr=lr)
  525. local_states.append(state)
  526. print(f" Site={site:<8} loss={stats['loss']:.4f} acc={stats['acc']:.4f}")
  527. global_state = fedavg(local_states)
  528. load_state_dict_strict(global_model, global_state)
  529. # Evaluate on pooled dataset
  530. eval_loader = DataLoader(MultiModalRetinalDataset(records, image_size=image_size, augment=False),
  531. batch_size=batch_size, shuffle=False, num_workers=0)
  532. metrics = evaluate(global_model, eval_loader, device_t)
  533. print("[Final Metrics]", metrics)
  534. return metrics
  535. # -----------------------------------------------------------------------------
  536. # CLI entry
  537. # -----------------------------------------------------------------------------
  538. def parse_args() -> argparse.Namespace:
  539. parser = argparse.ArgumentParser(description="RNV-T empirical validation code")
  540. parser.add_argument("--device", default="cpu", help="cpu or cuda")
  541. parser.add_argument("--image-size", type=int, default=224)
  542. parser.add_argument("--batch-size", type=int, default=8)
  543. parser.add_argument("--rounds", type=int, default=3)
  544. parser.add_argument("--lr", type=float, default=1e-4)
  545. parser.add_argument("--seed", type=int, default=42)
  546. parser.add_argument("--demo", action="store_true", help="Run on generated multimodal data")
  547. parser.add_argument("--root", default="./_rnvt_data")
  548. return parser.parse_args()
  549. def main() -> None:
  550. args = parse_args()
  551. seed_everything(args.seed)
  552. if args._demo:
  553. records = make_multimodal_dataset(args.root, n=48, image_size=args.image_size)
  554. metrics = federated_train(
  555. records,
  556. image_size=args.image_size,
  557. batch_size=args.batch_size,
  558. rounds=args.rounds,
  559. lr=args.lr,
  560. device=args.device,
  561. )
  562. print("empirical validation completed.")
  563. print(metrics)
  564. else:
  565. print(
  566. "This reference implementation is ready, but real dataset CSV / file paths were not provided.\n"
  567. "Use --demo for a runnable demonstration, or adapt the SampleRecord loader to your datasets."
  568. )
  569. if __name__ == "__main__":
  570. main()

rnv_t_empirical_validation.py at commit 284c4a2, under GPL-3.0 · at the source

Overview

Authors: S. Nandhini1, K. Vanitha1
  1. Department of Computer Science and Engineering, Faculty of Engineering, Karpagam Academy of Higher Education,Coimbatore, Tamil Nadu India
Journal: Scientific reports, volume 16, issue 1, article 17579
Dates: received 26 June 2025; accepted 8 April 2026; published online 15 April 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41598-026-48421-6 · PMID 41986421 · PMCID PMC13243560 · OpenAlex W7154448337
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), stroke (population)
Methods: Machine learning
Keywords: Neural, Transformers, CSVD, Neuro-ophthalmic, Blood vessels, Eye diseases, Cognitive neuroscience, Diseases of the nervous system
MeSH: Cerebral Small Vessel Diseases*, Deep Learning*, Image Processing, Computer-Assisted*, Retina*, Convolutional Neural Networks, Early Diagnosis, Humans, Retinal Vessels (* major topic)
Topic: Retinal Imaging and Analysis (Radiology, Nuclear Medicine and Imaging, Medicine), according to OpenAlex
Citations: not cited yet (Europe PMC); 38 references in the paper

Abstract

The primary driver or cause of cognitive decline and stroke is Cerebral Small Vessel Disease (CSVD), which currently requires neuroimaging tests, which are expensive to obtain and inaccessible in standard clinical settings. Low-cost retinal imaging techniques offer non-invasive assessments that mirror the condition of the brain’s small blood vessels (cerebral microvasculature). State-of-the-art diagnostic methods currently have no accessible, non-invasive, or cost-effective solution to identify CSVD at its earliest stages through visual assessment of retinal biomarkers. This study presents the Retino-Neuro Vision Transformer (RNV-T) framework as a proposed method to detect CSVD utilizing multimodality retinal imaging. The model system comprises five fundamental phases, beginning with Local Vascular Extraction (LVE), followed by Global Transformer-based Encoding (GTE), then proceeding to Graph-Based Relational Learning through Graph-based Convolutional Attention Network (G-CAN) before implementing Local-Global Attention Fusion (LGAF) as well as Optimized training procedures to obtain precise micro-vascular abnormality detection. The diagnostic performance of this model reaches 98.8% accuracy and shows 97.4% sensitivity along with 98.1% specificity, surpassing previous detection approaches. This diagnostic system represents a major leap forward in neuro-ophthalmic care because it enables early prediction of CSVD while expanding medical accessibility through retinal scans that are easy to conduct.

Supplementary Information: The online version contains supplementary material available at 10.1038/s41598-026-48421-6.

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 2 matches between paragraphs and lines of code.

Nandhiniphd07/RNVT

License: GPL-3.0
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: 284c4a276ec4237aefc0175c44e011d1410cd611, 21 March 2026
Languages: Python (1)
Size: 3 files, 1 script
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (1 file), OpenCV (1 file), PyTorch (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
3 files

Code availability

The custom code used for the implementation and empirical validation of the proposed method is publicly available at https://github.com/Nandhiniphd07/RNVT.git. The version corresponding to this study is archived and can be accessed for reproducibility purposes. The code is provided for research and academic purposes and is compatible with standard Python-based deep learning environments.

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

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;
  • 1 script, 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

Datasets cited

Data availability

The data that support the findings of this study are available from the corresponding author upon reasonable request.

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

Recorded: type, language, journal, volume, issue, pages, dates, 2 authors, 8 keywords, 8 MeSH terms, 30 references.

Cite

This paper

Nandhini, S., & Vanitha, K. (2026). Deep learning for early detection of cerebral small vessel disease using self-supervised graph embeddings and retinal image analysis. Scientific reports, 16(1), 17579. https://doi.org/10.1038/s41598-026-48421-6

BibTeX

@article{nandhini2026deep,
author = {Nandhini, S. and Vanitha, K.},
title = {{Deep learning for early detection of cerebral small vessel disease using self-supervised graph embeddings and retinal image analysis}},
journal = {Scientific reports},
year = {2026},
month = apr,
volume = {16},
number = {1},
pages = {17579},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/s41598-026-48421-6},
url = {https://doi.org/10.1038/s41598-026-48421-6},
pmid = {41986421},
pmcid = {PMC13243560}
}

RIS

TY - JOUR
AU - Nandhini, S.
AU - Vanitha, K.
TI - Deep learning for early detection of cerebral small vessel disease using self-supervised graph embeddings and retinal image analysis
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/04/15
VL - 16
IS - 1
SP - 17579
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-48421-6
UR - https://doi.org/10.1038/s41598-026-48421-6
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-48421-6",
"type": "article-journal",
"title": "Deep learning for early detection of cerebral small vessel disease using self-supervised graph embeddings and retinal image analysis",
"container-title": "Scientific reports",
"author": [
{
"family": "Nandhini",
"given": "S."
},
{
"family": "Vanitha",
"given": "K."
}
],
"container-title-short": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "17579",
"DOI": "10.1038/s41598-026-48421-6",
"PMID": "41986421",
"PMCID": "PMC13243560",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-48421-6",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
15
]
]
}
}

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.1371/journal.pone.0347867 [code]
LiteFeatNet: A parameter-efficient and performance-centric deep learning model for multi-ocular disease identification using intermediate feature reduction from fundus images.
Journal: PloS one
In common: OpenCV, PyTorch, NumPy, 1 reference
[2] doi:10.3390/diagnostics16172861 [code]
Microstructural Changes in the Corpus Callosum in Different Forms of Sporadic Age-Related Cerebral Small Vessel Disease.
Journal: Diagnostics (Basel, Switzerland)
In common: OpenCV, NumPy, stroke, 1 reference
[3] doi:10.1038/s41598-026-41137-7 [code]
Deep optimization-guided hybrid neural network for accurate detection and segmentation of white matter hyperintensities in clinical MRI images.
Journal: Scientific reports
In common: OpenCV, PyTorch, NumPy, stroke
[4] doi:10.1038/s41598-026-43798-w [code]
A Machine learning pipeline to investigate tissue ingrowth in cerebral aneurysms using preclinical animal models.
Journal: Scientific reports
In common: OpenCV, PyTorch, NumPy, stroke
[5] doi:10.1038/s41598-026-50020-4
nnLoGoNet: a hybrid local-global network for retinal vessel segmentation with Skeleton Recall Loss.
Journal: Scientific reports
In common: 2 references
[6] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: OpenCV, PyTorch, NumPy
[7] doi:10.1007/s12021-026-09815-z [code]
BasNet: Attention U-Net-Based Automated Axon Segmentation in Bielschowsky Silver-Stained Histology.
Journal: Neuroinformatics
In common: OpenCV, PyTorch, NumPy
[8] doi:10.1038/s41467-026-76956-9 [code]
Innervated human cardiac muscle model reveals sympathetic drivers of KCNH2-associated arrhythmias.
Journal: Nature communications
In common: OpenCV, PyTorch, NumPy
[9] doi:10.1038/s41598-026-61605-4 [code]
Learning precise segmentation of neurofibrillary tangles from rapid manual point annotations.
Journal: Scientific reports
In common: OpenCV, PyTorch, NumPy
[10] doi:10.1038/s41467-026-76837-1 [code]
Drug screen and machine learning predict neuroprotective agents in a preclinical human model of childhood dementia.
Journal: Nature communications
In common: OpenCV, PyTorch, NumPy

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.