Deep learning for early detection of cerebral small vessel disease using self-supervised graph embeddings and retinal image analysis.
The 2 matches
- [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] § 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
- """
- RNV-T empirical validation reference implementation.
- This script implements the five phases described in the manuscript:
- 1) Local Vascular Extraction (LVE)
- 2) Global Transformer-based Encoding (GTE)
- 3) Graph-based Convolutional Attention Network (G-CAN)
- 4) Local-Global Attention Fusion (LGAF)
- 5) Optimized training with SSL + federated-style aggregation
- It is designed as a reproducible empirical validation scaffold rather than the
- original laboratory code. Dataset paths, labels, and site partitions can be
- plugged in through the config section / CLI.
- """
- from __future__ import annotations
- import argparse
- import math
- import os
- import random
- from dataclasses import dataclass
- from pathlib import Path
- from typing import Dict, Iterable, List, Optional, Sequence, Tuple
- import cv2
- import numpy as np
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- from torch.utils.data import DataLoader, Dataset
- # -----------------------------------------------------------------------------
- # Reproducibility helpers
- # -----------------------------------------------------------------------------
- def seed_everything(seed: int = 42) -> None:
- random.seed(seed)
- np.random.seed(seed)
- torch.manual_seed(seed)
- torch.cuda.manual_seed_all(seed)
- # -----------------------------------------------------------------------------
- # Preprocessing utilities
- # -----------------------------------------------------------------------------
- def apply_clahe_rgb(image: np.ndarray) -> np.ndarray:
- lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)
- l, a, b = cv2.split(lab)
- clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
- l = clahe.apply(l)
- out = cv2.merge([l, a, b])
- return cv2.cvtColor(out, cv2.COLOR_LAB2RGB)
- def top_hat_enhance(gray: np.ndarray, kernel_size: int = 15) -> np.ndarray:
- kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size))
- return cv2.morphologyEx(gray, cv2.MORPH_TOPHAT, kernel)
- def preprocess_retinal_image(image: np.ndarray, image_size: int = 224) -> np.ndarray:
- image = cv2.resize(image, (image_size, image_size), interpolation=cv2.INTER_AREA)
- image = apply_clahe_rgb(image)
- image = cv2.GaussianBlur(image, (3, 3), 0)
- gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
- vessel_boost = top_hat_enhance(gray, kernel_size=11)
- vessel_boost = np.stack([vessel_boost] * 3, axis=-1)
- image = cv2.addWeighted(image, 0.85, vessel_boost, 0.15, 0.0)
- image = image.astype(np.float32) / 255.0
- return image
- def load_rgb(path: str, image_size: int = 224) -> np.ndarray:
- img = cv2.imread(path, cv2.IMREAD_COLOR)
- if img is None:
- raise FileNotFoundError(path)
- img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
- return preprocess_retinal_image(img, image_size=image_size)
- # -----------------------------------------------------------------------------
- # Dataset
- # -----------------------------------------------------------------------------
- @dataclass
- class SampleRecord:
- fundus: str
- octa: str
- fa: str
- label: int
- site: str = "site_0"
- class MultiModalRetinalDataset(Dataset):
- def __init__(self, records: Sequence[SampleRecord], image_size: int = 224, augment: bool = False):
- self.records = list(records)
- self.image_size = image_size
- self.augment = augment
- def __len__(self) -> int:
- return len(self.records)
- def _augment(self, x: np.ndarray) -> np.ndarray:
- if not self.augment:
- return x
- if random.random() < 0.5:
- x = np.flip(x, axis=1).copy()
- if random.random() < 0.5:
- x = np.flip(x, axis=0).copy()
- if random.random() < 0.2:
- noise = np.random.normal(0, 0.02, x.shape).astype(np.float32)
- x = np.clip(x + noise, 0.0, 1.0)
- return x
- def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
- rec = self.records[idx]
- fundus = self._augment(load_rgb(rec.fundus, self.image_size))
- octa = self._augment(load_rgb(rec.octa, self.image_size))
- fa = self._augment(load_rgb(rec.fa, self.image_size))
- return {
- "fundus": torch.from_numpy(fundus).permute(2, 0, 1),
- "octa": torch.from_numpy(octa).permute(2, 0, 1),
- "fa": torch.from_numpy(fa).permute(2, 0, 1),
- "label": torch.tensor(rec.label, dtype=torch.long),
- "site": rec.site,
- }
- # -----------------------------------------------------------------------------
- # Phase 1: Local Vascular Extraction (LVE)
- # -----------------------------------------------------------------------------
- class ConvBlock(nn.Module):
- def __init__(self, in_ch: int, out_ch: int):
- super().__init__()
- self.block = nn.Sequential(
- nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1),
- nn.BatchNorm2d(out_ch),
- nn.ReLU(inplace=True),
- nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1),
- nn.BatchNorm2d(out_ch),
- nn.ReLU(inplace=True),
- )
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- return self.block(x)
- class LVEBranch(nn.Module):
- def __init__(self, in_ch: int = 3, widths: Sequence[int] = (32, 64, 128, 256), emb_dim: int = 256):
- super().__init__()
- self.blocks = nn.ModuleList()
- prev = in_ch
- for w in widths:
- self.blocks.append(ConvBlock(prev, w))
- prev = w
- self.pool = nn.MaxPool2d(2)
- self.project = nn.Sequential(
- nn.Linear(sum(widths), emb_dim),
- nn.ReLU(inplace=True),
- nn.Linear(emb_dim, emb_dim),
- )
- def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, List[torch.Tensor]]:
- pyramids: List[torch.Tensor] = []
- h = x
- pooled_feats = []
- for i, block in enumerate(self.blocks):
- h = block(h)
- pyramids.append(h)
- pooled_feats.append(F.adaptive_avg_pool2d(h, 1).flatten(1))
- if i != len(self.blocks) - 1:
- h = self.pool(h)
- local_vec = self.project(torch.cat(pooled_feats, dim=1))
- return local_vec, pyramids
- class LVEEncoder(nn.Module):
- def __init__(self, emb_dim: int = 256):
- super().__init__()
- self.fundus_branch = LVEBranch(3, emb_dim=emb_dim)
- self.octa_branch = LVEBranch(3, emb_dim=emb_dim)
- self.fa_branch = LVEBranch(3, emb_dim=emb_dim)
- self.cross_modal = nn.Sequential(
- nn.Linear(emb_dim * 3, emb_dim * 2),
- nn.ReLU(inplace=True),
- nn.Linear(emb_dim * 2, emb_dim),
- )
- def forward(self, fundus: torch.Tensor, octa: torch.Tensor, fa: torch.Tensor) -> Dict[str, torch.Tensor | List[torch.Tensor]]:
- f_vec, f_maps = self.fundus_branch(fundus)
- o_vec, o_maps = self.octa_branch(octa)
- a_vec, a_maps = self.fa_branch(fa)
- hadamard_fo = f_vec * o_vec
- hadamard_fa = f_vec * a_vec
- hadamard_oa = o_vec * a_vec
- fused = self.cross_modal(torch.cat([hadamard_fo, hadamard_fa, hadamard_oa], dim=1))
- return {
- "local_vec": fused,
- "fundus_vec": f_vec,
- "octa_vec": o_vec,
- "fa_vec": a_vec,
- "fundus_maps": f_maps,
- "octa_maps": o_maps,
- "fa_maps": a_maps,
- }
- # -----------------------------------------------------------------------------
- # Phase 2: Global Transformer-based Encoding (GTE)
- # -----------------------------------------------------------------------------
- class PatchEmbed(nn.Module):
- def __init__(self, img_size: int = 224, patch_size: int = 16, in_chans: int = 3, embed_dim: int = 384):
- super().__init__()
- self.grid_size = img_size // patch_size
- self.num_patches = self.grid_size * self.grid_size
- self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- x = self.proj(x) # [B, D, H', W']
- return x.flatten(2).transpose(1, 2) # [B, N, D]
- class MHCA(nn.Module):
- """Multi-head correlative attention."""
- def __init__(self, dim: int, heads: int = 8, dropout: float = 0.1):
- super().__init__()
- self.heads = heads
- self.scale = (dim // heads) ** -0.5
- self.qkv = nn.Linear(dim, dim * 3)
- self.corr = nn.Parameter(torch.zeros(heads, 1, 1))
- self.proj = nn.Linear(dim, dim)
- self.drop = nn.Dropout(dropout)
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- b, n, c = x.shape
- qkv = self.qkv(x).reshape(b, n, 3, self.heads, c // self.heads).permute(2, 0, 3, 1, 4)
- q, k, v = qkv[0], qkv[1], qkv[2]
- sim = (q @ k.transpose(-2, -1)) * self.scale
- div = -torch.abs(q.unsqueeze(-2) - k.unsqueeze(-3)).mean(dim=-1)
- attn = sim + self.corr + 0.1 * div
- attn = attn.softmax(dim=-1)
- attn = self.drop(attn)
- out = (attn @ v).transpose(1, 2).reshape(b, n, c)
- return self.proj(out)
- class TransformerBlock(nn.Module):
- def __init__(self, dim: int, heads: int = 8, mlp_ratio: float = 4.0, dropout: float = 0.1):
- super().__init__()
- self.norm1 = nn.LayerNorm(dim)
- self.attn = MHCA(dim, heads=heads, dropout=dropout)
- self.norm2 = nn.LayerNorm(dim)
- self.mlp = nn.Sequential(
- nn.Linear(dim, int(dim * mlp_ratio)),
- nn.GELU(),
- nn.Dropout(dropout),
- nn.Linear(int(dim * mlp_ratio), dim),
- nn.Dropout(dropout),
- )
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- x = x + self.attn(self.norm1(x))
- x = x + self.mlp(self.norm2(x))
- return x
- class GTEEncoder(nn.Module):
- def __init__(self, img_size: int = 224, patch_size: int = 16, dim: int = 384, depth: int = 4, heads: int = 8):
- super().__init__()
- self.patch_embed = PatchEmbed(img_size, patch_size, 9, dim)
- num_patches = self.patch_embed.num_patches
- self.cls_token = nn.Parameter(torch.zeros(1, 1, dim))
- self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, dim))
- self.blocks = nn.ModuleList([TransformerBlock(dim, heads=heads) for _ in range(depth)])
- self.norm = nn.LayerNorm(dim)
- def forward(self, fused_modal_image: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
- x = self.patch_embed(fused_modal_image)
- b = x.shape[0]
- cls = self.cls_token.expand(b, -1, -1)
- x = torch.cat([cls, x], dim=1) + self.pos_embed
- for block in self.blocks:
- x = block(x)
- x = self.norm(x)
- return x[:, 0], x[:, 1:]
- # -----------------------------------------------------------------------------
- # Phase 3: Graph-based Convolutional Attention Network (G-CAN)
- # -----------------------------------------------------------------------------
- class GraphAttentionLayer(nn.Module):
- def __init__(self, in_dim: int, out_dim: int, dropout: float = 0.2, alpha: float = 0.2):
- super().__init__()
- self.lin = nn.Linear(in_dim, out_dim, bias=False)
- self.attn = nn.Linear(2 * out_dim, 1, bias=False)
- self.dropout = nn.Dropout(dropout)
- self.leaky_relu = nn.LeakyReLU(alpha)
- def forward(self, x: torch.Tensor, adj: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
- h = self.lin(x)
- b, n, d = h.shape
- h_i = h.unsqueeze(2).expand(b, n, n, d)
- h_j = h.unsqueeze(1).expand(b, n, n, d)
- e = self.leaky_relu(self.attn(torch.cat([h_i, h_j], dim=-1)).squeeze(-1))
- e = e.masked_fill(adj <= 0, float("-inf"))
- alpha = F.softmax(e, dim=-1)
- alpha = self.dropout(alpha)
- out = alpha @ h
- return F.elu(out), alpha
- class GCANEncoder(nn.Module):
- def __init__(self, node_dim: int, hidden_dim: int = 128, num_layers: int = 3, dropout: float = 0.2):
- super().__init__()
- dims = [node_dim] + [hidden_dim] * num_layers
- self.layers = nn.ModuleList([
- GraphAttentionLayer(dims[i], dims[i + 1], dropout=dropout) for i in range(num_layers)
- ])
- self.out_proj = nn.Linear(hidden_dim + 1, hidden_dim)
- @staticmethod
- def build_graph(node_feats: torch.Tensor, coords: torch.Tensor, sigma: float = 0.35) -> torch.Tensor:
- # coords: [B, N, 2], normalized to [0,1]
- dists = torch.cdist(coords, coords, p=2)
- adj = torch.exp(-(dists ** 2) / (2 * sigma * sigma))
- eye = torch.eye(adj.size(-1), device=adj.device).unsqueeze(0)
- adj = torch.maximum(adj, eye)
- return adj
- def forward(self, node_feats: torch.Tensor, coords: torch.Tensor) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:
- adj = self.build_graph(node_feats, coords)
- h = node_feats
- attn_maps = []
- for layer in self.layers:
- h, alpha = layer(h, adj)
- attn_maps.append(alpha)
- predicted_adj = torch.sigmoid(torch.matmul(F.normalize(h, dim=-1), F.normalize(h, dim=-1).transpose(-1, -2)))
- anomaly = torch.mean(torch.abs(adj - predicted_adj), dim=(-1, -2), keepdim=True)
- graph_repr = h.mean(dim=1)
- graph_repr = self.out_proj(torch.cat([graph_repr, anomaly], dim=-1))
- return graph_repr, {
- "adj": adj,
- "predicted_adj": predicted_adj,
- "anomaly": anomaly,
- "attn_maps": attn_maps[-1],
- }
- # -----------------------------------------------------------------------------
- # Phase 4: Local-Global Attention Fusion (LGAF)
- # -----------------------------------------------------------------------------
- class LGAF(nn.Module):
- def __init__(self, local_dim: int, global_dim: int, graph_dim: int, hidden_dim: int = 512, num_classes: int = 2):
- super().__init__()
- self.local_proj = nn.Linear(local_dim, hidden_dim)
- self.global_proj = nn.Linear(global_dim, hidden_dim)
- self.graph_proj = nn.Linear(graph_dim, hidden_dim)
- self.gate = nn.Sequential(
- nn.Linear(hidden_dim * 2, hidden_dim),
- nn.ReLU(inplace=True),
- nn.Linear(hidden_dim, hidden_dim),
- nn.Sigmoid(),
- )
- self.attn = nn.MultiheadAttention(hidden_dim, num_heads=4, batch_first=True)
- self.norm = nn.LayerNorm(hidden_dim)
- self.classifier = nn.Sequential(
- nn.Linear(hidden_dim, hidden_dim // 2),
- nn.ReLU(inplace=True),
- nn.Dropout(0.2),
- nn.Linear(hidden_dim // 2, num_classes),
- )
- def forward(self, local_vec: torch.Tensor, global_vec: torch.Tensor, graph_vec: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
- l = self.local_proj(local_vec)
- g = self.global_proj(global_vec)
- r = self.graph_proj(graph_vec)
- gate = self.gate(torch.cat([l, g], dim=-1))
- fused = gate * l + (1.0 - gate) * g
- tokens = torch.stack([fused, r], dim=1)
- attended, _ = self.attn(tokens, tokens, tokens)
- fused = self.norm(attended.mean(dim=1) + fused + 0.5 * r)
- logits = self.classifier(fused)
- return logits, fused
- # -----------------------------------------------------------------------------
- # Full RNV-T model
- # -----------------------------------------------------------------------------
- class RNVT(nn.Module):
- def __init__(self, image_size: int = 224, local_dim: int = 256, global_dim: int = 384, graph_dim: int = 128, num_classes: int = 2):
- super().__init__()
- self.lve = LVEEncoder(emb_dim=local_dim)
- self.gte = GTEEncoder(img_size=image_size, patch_size=16, dim=global_dim, depth=4, heads=8)
- self.gcan = GCANEncoder(node_dim=global_dim, hidden_dim=graph_dim, num_layers=3, dropout=0.2)
- self.lgaf = LGAF(local_dim=local_dim, global_dim=global_dim, graph_dim=graph_dim, hidden_dim=512, num_classes=num_classes)
- self.proj_ssl = nn.Sequential(
- nn.Linear(512, 256), nn.ReLU(inplace=True), nn.Linear(256, 128)
- )
- @staticmethod
- def _token_coords(num_tokens: int, device: torch.device) -> torch.Tensor:
- side = int(math.sqrt(num_tokens))
- ys, xs = torch.meshgrid(
- torch.linspace(0, 1, side, device=device),
- torch.linspace(0, 1, side, device=device),
- indexing="ij",
- )
- coords = torch.stack([xs.reshape(-1), ys.reshape(-1)], dim=-1)
- return coords
- def forward(self, fundus: torch.Tensor, octa: torch.Tensor, fa: torch.Tensor) -> Dict[str, torch.Tensor]:
- lve_out = self.lve(fundus, octa, fa)
- fused_modal_image = torch.cat([fundus, octa, fa], dim=1)
- global_vec, patch_tokens = self.gte(fused_modal_image)
- b, n, d = patch_tokens.shape
- coords = self._token_coords(n, patch_tokens.device).unsqueeze(0).repeat(b, 1, 1)
- graph_vec, graph_aux = self.gcan(patch_tokens, coords)
- logits, fused_repr = self.lgaf(lve_out["local_vec"], global_vec, graph_vec)
- return {
- "logits": logits,
- "local_vec": lve_out["local_vec"],
- "global_vec": global_vec,
- "graph_vec": graph_vec,
- "fused_repr": fused_repr,
- "ssl_proj": F.normalize(self.proj_ssl(fused_repr), dim=-1),
- "graph_anomaly": graph_aux["anomaly"].squeeze(-1),
- "adj": graph_aux["adj"],
- "predicted_adj": graph_aux["predicted_adj"],
- }
- # -----------------------------------------------------------------------------
- # Phase 5: losses and optimization
- # -----------------------------------------------------------------------------
- def info_nce(z1: torch.Tensor, z2: torch.Tensor, temperature: float = 0.2) -> torch.Tensor:
- b = z1.size(0)
- z = torch.cat([z1, z2], dim=0)
- sim = torch.mm(z, z.t()) / temperature
- mask = torch.eye(2 * b, device=z.device, dtype=torch.bool)
- sim = sim.masked_fill(mask, -1e9)
- targets = torch.arange(b, device=z.device)
- targets = torch.cat([targets + b, targets], dim=0)
- return F.cross_entropy(sim, targets)
- def graph_regularization(adj: torch.Tensor, pred_adj: torch.Tensor) -> torch.Tensor:
- return F.mse_loss(pred_adj, adj)
- def total_loss_fn(outputs_1: Dict[str, torch.Tensor], outputs_2: Dict[str, torch.Tensor], labels: torch.Tensor,
- lambda_ssl: float = 0.2, lambda_reg: float = 0.1) -> Tuple[torch.Tensor, Dict[str, float]]:
- ce = F.cross_entropy(outputs_1["logits"], labels)
- ssl = info_nce(outputs_1["ssl_proj"], outputs_2["ssl_proj"])
- reg = graph_regularization(outputs_1["adj"], outputs_1["predicted_adj"])
- total = ce + lambda_ssl * ssl + lambda_reg * reg
- return total, {
- "ce": float(ce.detach().cpu()),
- "ssl": float(ssl.detach().cpu()),
- "reg": float(reg.detach().cpu()),
- "total": float(total.detach().cpu()),
- }
- # -----------------------------------------------------------------------------
- # Utilities for federated-style training
- # -----------------------------------------------------------------------------
- def clone_state_dict(model: nn.Module) -> Dict[str, torch.Tensor]:
- return {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}
- def load_state_dict_strict(model: nn.Module, state: Dict[str, torch.Tensor]) -> None:
- model.load_state_dict(state, strict=True)
- def fedavg(states: List[Dict[str, torch.Tensor]]) -> Dict[str, torch.Tensor]:
- avg = {}
- for key in states[0].keys():
- avg[key] = sum(s[key] for s in states) / len(states)
- return avg
- def accuracy_from_logits(logits: torch.Tensor, labels: torch.Tensor) -> float:
- preds = logits.argmax(dim=1)
- return float((preds == labels).float().mean().item())
- def sensitivity_specificity(logits: torch.Tensor, labels: torch.Tensor) -> Tuple[float, float]:
- preds = logits.argmax(dim=1)
- tp = ((preds == 1) & (labels == 1)).sum().item()
- tn = ((preds == 0) & (labels == 0)).sum().item()
- fp = ((preds == 1) & (labels == 0)).sum().item()
- fn = ((preds == 0) & (labels == 1)).sum().item()
- sensitivity = tp / max(tp + fn, 1)
- specificity = tn / max(tn + fp, 1)
- return sensitivity, specificity
- # -----------------------------------------------------------------------------
- # Demo generation for validation when real data is unavailable
- # -----------------------------------------------------------------------------
- def make_multimodal_dataset(root: str, n: int = 48, image_size: int = 224) -> List[SampleRecord]:
- os.makedirs(root, exist_ok=True)
- records: List[SampleRecord] = []
- for i in range(n):
- label = i % 2
- site = f"site_{i % 3}"
- paths = {}
- for modality in ["fundus", "octa", "fa"]:
- canvas = np.zeros((image_size, image_size, 3), dtype=np.uint8)
- # background tone
- canvas[:] = (20 + 15 * label, 30 + 5 * (i % 5), 20 + 10 * (i % 3))
- # vessel-like line patterns
- center = (image_size // 2, image_size // 2)
- num_lines = 18 + 6 * label
- thickness = 1 + label
- for a in np.linspace(0, 2 * np.pi, num_lines, endpoint=False):
- r = image_size // 2 - 8
- end = (int(center[0] + r * np.cos(a)), int(center[1] + r * np.sin(a)))
- cv2.line(canvas, center, end, (180, 180, 180), thickness)
- if modality == "octa" and label == 1:
- cv2.circle(canvas, center, 26, (10, 10, 10), -1) # FAZ enlargement cue
- if modality == "fa" and label == 1:
- cv2.circle(canvas, (image_size // 3, image_size // 3), 14, (230, 230, 230), -1)
- if modality == "fundus" and label == 1:
- cv2.circle(canvas, (2 * image_size // 3, image_size // 2), 10, (220, 50, 50), -1)
- p = os.path.join(root, f"{modality}_{i}.png")
- cv2.imwrite(p, cv2.cvtColor(canvas, cv2.COLOR_RGB2BGR))
- paths[modality] = p
- records.append(SampleRecord(paths["fundus"], paths["octa"], paths["fa"], label, site))
- return records
- # -----------------------------------------------------------------------------
- # Training / evaluation
- # -----------------------------------------------------------------------------
- def build_site_loaders(records: Sequence[SampleRecord], image_size: int, batch_size: int) -> Dict[str, DataLoader]:
- sites: Dict[str, List[SampleRecord]] = {}
- for rec in records:
- sites.setdefault(rec.site, []).append(rec)
- loaders = {}
- for site, site_records in sites.items():
- ds = MultiModalRetinalDataset(site_records, image_size=image_size, augment=True)
- loaders[site] = DataLoader(ds, batch_size=batch_size, shuffle=True, num_workers=0)
- return loaders
- def make_augmented_views(batch: Dict[str, torch.Tensor]) -> Tuple[Dict[str, torch.Tensor], Dict[str, torch.Tensor]]:
- def jitter(x: torch.Tensor) -> torch.Tensor:
- noise = 0.03 * torch.randn_like(x)
- return torch.clamp(x + noise, 0.0, 1.0)
- v1 = {k: jitter(v) if k in ["fundus", "octa", "fa"] else v for k, v in batch.items()}
- v2 = {k: jitter(v) if k in ["fundus", "octa", "fa"] else v for k, v in batch.items()}
- return v1, v2
- def run_local_epoch(model: RNVT, loader: DataLoader, device: torch.device, lr: float = 1e-4) -> Tuple[Dict[str, torch.Tensor], Dict[str, float]]:
- model.train()
- optimizer = torch.optim.Adam(model.parameters(), lr=lr)
- meters = {"loss": 0.0, "acc": 0.0, "count": 0}
- for batch in loader:
- batch = {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in batch.items()}
- v1, v2 = make_augmented_views(batch)
- out1 = model(v1["fundus"], v1["octa"], v1["fa"])
- out2 = model(v2["fundus"], v2["octa"], v2["fa"])
- loss, _ = total_loss_fn(out1, out2, batch["label"])
- optimizer.zero_grad()
- loss.backward()
- optimizer.step()
- meters["loss"] += loss.item() * batch["label"].size(0)
- meters["acc"] += accuracy_from_logits(out1["logits"], batch["label"]) * batch["label"].size(0)
- meters["count"] += batch["label"].size(0)
- stats = {
- "loss": meters["loss"] / max(meters["count"], 1),
- "acc": meters["acc"] / max(meters["count"], 1),
- }
- return clone_state_dict(model), stats
- @torch.no_grad()
- def evaluate(model: RNVT, loader: DataLoader, device: torch.device) -> Dict[str, float]:
- model.eval()
- all_logits, all_labels = [], []
- for batch in loader:
- fundus = batch["fundus"].to(device)
- octa = batch["octa"].to(device)
- fa = batch["fa"].to(device)
- labels = batch["label"].to(device)
- out = model(fundus, octa, fa)
- all_logits.append(out["logits"])
- all_labels.append(labels)
- logits = torch.cat(all_logits, dim=0)
- labels = torch.cat(all_labels, dim=0)
- acc = accuracy_from_logits(logits, labels)
- sens, spec = sensitivity_specificity(logits, labels)
- return {"accuracy": acc, "sensitivity": sens, "specificity": spec}
- def federated_train(records: Sequence[SampleRecord], image_size: int = 224, batch_size: int = 8, rounds: int = 3,
- lr: float = 1e-4, device: str = "cpu") -> Dict[str, float]:
- device_t = torch.device(device)
- loaders = build_site_loaders(records, image_size=image_size, batch_size=batch_size)
- global_model = RNVT(image_size=image_size).to(device_t)
- global_state = clone_state_dict(global_model)
- for rnd in range(1, rounds + 1):
- local_states = []
- print(f"[Federated Round {rnd}/{rounds}]")
- for site, loader in loaders.items():
- local_model = RNVT(image_size=image_size).to(device_t)
- load_state_dict_strict(local_model, global_state)
- state, stats = run_local_epoch(local_model, loader, device_t, lr=lr)
- local_states.append(state)
- print(f" Site={site:<8} loss={stats['loss']:.4f} acc={stats['acc']:.4f}")
- global_state = fedavg(local_states)
- load_state_dict_strict(global_model, global_state)
- # Evaluate on pooled dataset
- eval_loader = DataLoader(MultiModalRetinalDataset(records, image_size=image_size, augment=False),
- batch_size=batch_size, shuffle=False, num_workers=0)
- metrics = evaluate(global_model, eval_loader, device_t)
- print("[Final Metrics]", metrics)
- return metrics
- # -----------------------------------------------------------------------------
- # CLI entry
- # -----------------------------------------------------------------------------
- def parse_args() -> argparse.Namespace:
- parser = argparse.ArgumentParser(description="RNV-T empirical validation code")
- parser.add_argument("--device", default="cpu", help="cpu or cuda")
- parser.add_argument("--image-size", type=int, default=224)
- parser.add_argument("--batch-size", type=int, default=8)
- parser.add_argument("--rounds", type=int, default=3)
- parser.add_argument("--lr", type=float, default=1e-4)
- parser.add_argument("--seed", type=int, default=42)
- parser.add_argument("--demo", action="store_true", help="Run on generated multimodal data")
- parser.add_argument("--root", default="./_rnvt_data")
- return parser.parse_args()
- def main() -> None:
- args = parse_args()
- seed_everything(args.seed)
- if args._demo:
- records = make_multimodal_dataset(args.root, n=48, image_size=args.image_size)
- metrics = federated_train(
- records,
- image_size=args.image_size,
- batch_size=args.batch_size,
- rounds=args.rounds,
- lr=args.lr,
- device=args.device,
- )
- print("empirical validation completed.")
- print(metrics)
- else:
- print(
- "This reference implementation is ready, but real dataset CSV / file paths were not provided.\n"
- "Use --demo for a runnable demonstration, or adapt the SampleRecord loader to your datasets."
- )
- if __name__ == "__main__":
- main()
rnv_t_empirical_validation.py at commit 284c4a2, under GPL-3.0 · at the source
Overview
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/
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
284c4a276ec4237aefc0175c44e011d1410cd611, 21 March 2026Availability: 1 check, the latest on 29 September 2026: the link answers
- 29 September 2026: the link answers
3 files
- rnv_t_empirical_validati
on.py , Python, 683 lines, 2 matches - LICENSE, License, 674 lines
- README.md, Text, 79 lines
Code availability
The custom code used for the implementation and empirical validation of the proposed method is publicly available at https://
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
- zenodo:12775880, at Zenodo; found in the references
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://
BibTeX
@article{nandhini2026dee
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/
url = {https://
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/
VL - 16
IS - 1
SP - 17579
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"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":
"volume": "16",
"issue": "1",
"page": "17579",
"DOI": "10.1038/
"PMID": "41986421",
"PMCID": "PMC13243560",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://
"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 oneIn 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 reportsIn 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 reportsIn 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 reportsIn common: 2 references
- [6] doi:10.1038/s42003-026-10957-8 [code]
- Brain defence by the extracellular matrix protein Cochlin.Journal: Communications biologyIn 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: NeuroinformaticsIn 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 communicationsIn 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 reportsIn 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 communicationsIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 1 script, and 2 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:37698ed4882f24d2…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
