MambaKAN: An Interpretable Framework for Alzheimer's Disease Diagnosis via Selective State Space Modeling of Dynamic Functional Connectivity.
The 17 matches
- [1] § 3. Methods › 3.5. Joint Training Strategy › 3.5.1. Two-Phase Training Protocol ↔ train_stage1.py, lines 1–30 · score 0.92 · lowest validation reconstruction, VAE Unsupervised Pre, noise robust latent, best checkpoint, pre training, Adam
- [2] § 3. Methods › 3.4. Kolmogorov–Arnold Network Classifier › 3.4.1. KAN Architecture Motivation ↔ models/kan_classifier.py, lines 1–17 · score 0.84 · Kolmogorov Arnold Networks, learnable weights, visualizable curve, intrinsically interpretable, inspired, MLPs
- [3] § 3. Methods › 3.3. Mamba Selective State Space Temporal Encoder › 3.3.3. Selective State Space Model (S6) ↔ models/mamba_encoder.py, lines 13–22 · score 0.82 · Selective State Space, depthwise convolution, selectivity mechanism, Mamba block, S6, Model
- [4] § 3. Methods › 3.5. Joint Training Strategy › 3.5.1. Two-Phase Training Protocol ↔ train_stage2.py, lines 1–44 · score 0.76 · End Joint Fine, trained jointly, MambaKAN, tuning, pipeline, accuracy
- [5] § 3. Methods › 3.4. Kolmogorov–Arnold Network Classifier › 3.4.5. MambaKAN Classifier Structure ↔ models/kan_classifier.py, lines 180–280 · score 0.73 · layer KAN, context vector, Mamba encoder, KAN classifier, learnable, weights
- [6] § 3. Methods › 3.1. Overview of MambaKAN ↔ train_stage2.py, lines 1–44 · score 0.72 · VAE encoding, MambaKAN, KAN classification, Mamba temporal, pipeline, joint
- [7] § 5. Interpretability Analysis › 5.3. Layer 3: Gradient-Based Brain Region Attribution ↔ analysis.py, lines 230–357 · score 0.67 · bar charts, attribution scores, chord, heatmap, connectivity, Gradient
- [8] § 3. Methods › 3.5. Joint Training Strategy › 3.5.2. Differential Learning Rates ↔ train_stage2.py, lines 46–74 · score 0.63 · warmup epochs, fine tuning, frozen, trained, KAN, Mamba
- [9] § 3. Methods › 3.3. Mamba Selective State Space Temporal Encoder › 3.3.1. Rationale for Mamba over LSTM and Transformer ↔ models/mamba_encoder.py, lines 152–179 · score 0.63 · parallel scan, hardware aware, efficient, matrices, Selective, Mamba
- [10] § 3. Methods › 3.4. Kolmogorov–Arnold Network Classifier › 3.4.2. Rationale for KAN over MLP ↔ models/kan_classifier.py, lines 1–17 · score 0.63 · provides intrinsic interpretability, latent dimension, MLP, mapping, linear, logits
- [11] § 5. Interpretability Analysis › 5.3. Layer 3: Gradient-Based Brain Region Attribution ↔ analysis.py, lines 230–357 · score 0.59 · chord diagram, attribution scores, connectivity, Gradient, Brain, class
- [12] § 3. Methods › 3.4. Kolmogorov–Arnold Network Classifier › 3.4.3. B-Spline Edge Activations ↔ models/kan_classifier.py, lines 29–66 · score 0.59 · spline coefficients, spline basis functions, uniform, learnable, weight
- [13] § 3. Methods › 3.3. Mamba Selective State Space Temporal Encoder › 3.3.5. Temporal Context Aggregation ↔ models/mamba_encoder.py, lines 182–235 · score 0.58 · stacked Mamba blocks, temporal context, sequence, vector
- [14] § 5. Interpretability Analysis › 5.3. Layer 3: Gradient-Based Brain Region Attribution ↔ analysis.py, lines 1–41 · score 0.55 · attribution matrix, pairwise, ROI, brain, map, Gradient
- [15] § 3. Methods › 3.2. Variational Autoencoder for Per-Window Feature Extraction › 3.2.3. VAE Loss Function ↔ models/vae.py, lines 51–55 · score 0.54 · KL divergence, reconstruction loss, VAE
- [16] § 3. Methods › 3.5. Joint Training Strategy › 3.5.3. Regularization ↔ models/mamba_encoder.py, lines 13–22 · score 0.54 · depthwise convolution, Mamba block, Dropout, space
- [17] § 4. Experiments › 4.2. Evaluation Metrics ↔ analysis.py, lines 1–41 · score 0.52 · ROC curve, MambaKAN, AUC
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 · 297 lines · 10 KB · no license · 4 matches
- """
- Kolmogorov-Arnold Network (KAN) Classifier — B-Spline implementation.
- Reference: Liu et al., "KAN: Kolmogorov-Arnold Networks" (2024)
- Inspired by efficient-kan (https://github.com/Blealtan/efficient-kan)
- Key difference from MLP:
- MLP: y = W · σ(x) — fixed activation, learnable weights
- KAN: y = Σ φ_{q,p}(x_p) — learnable spline activations per connection
- Visualizing φ curves directly reveals the non-linear mapping from each
- latent dimension to class logits, providing intrinsic interpretability.
- """
- import math
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- class KANLinear(nn.Module):
- """
- A single KAN layer: replaces Linear + fixed activation with
- per-connection learnable B-spline activations.
- Input: (batch, in_features)
- Output: (batch, out_features)
- """
- def __init__(
- self,
- in_features: int,
- out_features: int,
- grid_size: int = 5,
- spline_order: int = 3,
- scale_noise: float = 0.1,
- scale_base: float = 1.0,
- grid_range: tuple = (-1.0, 1.0),
- ):
- super().__init__()
- self.in_features = in_features
- self.out_features = out_features
- self.grid_size = grid_size
- self.spline_order = spline_order
- # Build extended B-spline grid
- h = (grid_range[1] - grid_range[0]) / grid_size
- grid = torch.linspace(
- grid_range[0] - spline_order * h,
- grid_range[1] + spline_order * h,
- grid_size + 2 * spline_order + 1,
- )
- self.register_buffer("grid", grid)
- # Number of B-spline basis functions
- n_basis = grid_size + spline_order
- self.base_weight = nn.Parameter(torch.empty(out_features, in_features))
- # Spline coefficients: (out, in, n_basis)
- self.spline_weight = nn.Parameter(torch.empty(out_features, in_features, n_basis))
- # Per-connection scaling factors (learnable)
- self.scale_base = nn.Parameter(torch.ones(out_features, in_features) * scale_base)
- self.scale_spline = nn.Parameter(torch.ones(out_features, in_features))
- nn.init.kaiming_uniform_(self.base_weight, a=math.sqrt(5))
- nn.init.normal_(self.spline_weight, mean=0.0, std=scale_noise)
- def b_splines(self, x: torch.Tensor) -> torch.Tensor:
- """
- Evaluate B-spline basis functions at x via Cox–de Boor recursion.
- Args:
- x: (batch, in_features)
- Returns:
- bases: (batch, in_features, n_basis)
- """
- assert x.dim() == 2
- x = x.unsqueeze(-1)
- grid = self.grid
- # Order-0 indicator basis
- bases = ((x >= grid[:-1]) & (x < grid[1:])).float()
- # Cox–de Boor recursion
- for k in range(1, self.spline_order + 1):
- denom_l = grid[k:-1] - grid[: -(k + 1)]
- denom_r = grid[k + 1 :] - grid[1: -k]
- # Avoid division by zero
- left = torch.where(
- denom_l != 0,
- (x - grid[: -(k + 1)]) / denom_l * bases[..., :-1],
- torch.zeros_like(bases[..., :-1]),
- )
- right = torch.where(
- denom_r != 0,
- (grid[k + 1 :] - x) / denom_r * bases[..., 1:],
- torch.zeros_like(bases[..., 1:]),
- )
- bases = left + right
- return bases # (b, in, n_basis)
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- """
- Args:
- x: (batch, in_features)
- Returns:
- out: (batch, out_features)
- """
- # Base (SiLU) branch
- base_out = F.linear(F.silu(x), self.base_weight * self.scale_base)
- # Spline branch: contract (b, in, n_basis) × (out, in, n_basis) → (b, out)
- bases = self.b_splines(x)
- spline_out = torch.einsum(
- "bik,oik->bo",
- bases,
- self.spline_weight * self.scale_spline.unsqueeze(-1),
- )
- return base_out + spline_out
- def get_activation_curve(self, dim: int, n_points: int = 200, x_range=(-3.0, 3.0)):
- """
- Evaluate the learned spline activation φ(x) for a single input dimension.
- Args:
- dim: input dimension index
- n_points: number of evaluation points
- x_range: evaluation range
- Returns:
- x_vals: (n_points,)
- y_vals: (out_features, n_points)
- """
- x_vals = torch.linspace(x_range[0], x_range[1], n_points, device=self.grid.device)
- dummy = torch.zeros(n_points, self.in_features, device=self.grid.device)
- dummy[:, dim] = x_vals
- bases = self.b_splines(dummy)
- sw = self.spline_weight[:, dim, :] * self.scale_spline[:, dim].unsqueeze(-1)
- y_vals = (bases[:, dim, :] @ sw.T).T # (out, n_pts)
- return x_vals.detach(), y_vals.detach()
- def update_grid(self, x: torch.Tensor, margin: float = 0.01):
- """
- Adapt the B-spline grid to span the activation range of x.
- Note on coefficient resampling:
- A full grid update should re-fit spline_weight onto the new basis
- (e.g. via least squares) to preserve learned activation shapes.
- This implementation omits that step because the current architecture
- uses a single shared 1-D grid for all in_features; when the grid
- shifts significantly the new and old basis spaces diverge and simple
- least-squares resampling is numerically unreliable (larger error
- than no resampling at all on large feature dimensions).
- In practice, call this method infrequently (every 50 epochs) so the
- network has enough gradient steps to recover from the small
- coefficient mismatch before the next update.
- Args:
- x: (batch, in_features) — representative input batch
- margin: fractional padding added beyond [x_min, x_max]
- """
- with torch.no_grad():
- x_min, x_max = x.min().item(), x.max().item()
- span = max(x_max - x_min, 1e-6)
- x_min -= margin * span
- x_max += margin * span
- h = (x_max - x_min) / self.grid_size
- new_grid = torch.linspace(
- x_min - self.spline_order * h,
- x_max + self.spline_order * h,
- len(self.grid),
- device=self.grid.device,
- dtype=self.grid.dtype,
- )
- self.grid.copy_(new_grid)
- class KANClassifier(nn.Module):
- """
- Two-layer KAN for classification.
- Structure: in_features → hidden_dim → num_classes
- Parameter count note:
- This implementation includes per-connection learnable scaling factors
- (scale_base and scale_spline in each KANLinear layer) for training
- stability. These add (out * in) extra parameters per layer compared
- to counting only the spline and base weights. As a result the KAN
- head has ~93 K parameters, somewhat larger than the ~43 K figure
- reported in the paper (which counted only spline_weight + base_weight).
- """
- def __init__(
- self,
- in_features: int,
- hidden_dim: int,
- num_classes: int,
- grid_size: int = 5,
- spline_order: int = 3,
- ):
- super().__init__()
- self.layer1 = KANLinear(in_features, hidden_dim, grid_size, spline_order)
- self.layer2 = KANLinear(hidden_dim, num_classes, grid_size, spline_order)
- def forward(self, x: torch.Tensor) -> torch.Tensor:
- return self.layer2(self.layer1(x))
- def update_grid(self, x: torch.Tensor):
- """
- Update B-spline grids in both layers.
- Call with a representative context batch (output of MambaEncoder)
- every few epochs during Phase-2 training.
- Args:
- x: (batch, in_features) — Mamba context vectors
- """
- self.layer1.update_grid(x)
- with torch.no_grad():
- h = self.layer1(x)
- self.layer2.update_grid(h)
- def get_input_importance(self) -> torch.Tensor:
- """
- L1 norm of spline weights in layer 1, summed over outputs and basis functions.
- Returns:
- importance: (in_features,) — higher = more influential
- """
- w = self.layer1.spline_weight # (out, in, n_basis)
- return w.abs().sum(dim=[0, 2]) # (in_features,)
- def get_class_curves(
- self,
- dim: int,
- x_mean: torch.Tensor,
- n_points: int = 200,
- x_range: tuple = (-3.0, 3.0),
- ):
- """
- Ceteris-paribus class activation curve for input dimension `dim`.
- Varies dim over x_range while holding all other dims at x_mean.
- Args:
- dim: input dimension to vary
- x_mean: (in_features,) anchor values for non-varied dims
- n_points: evaluation resolution
- x_range: range of variation for `dim`
- Returns:
- x_vals: (n_points,)
- logits: (num_classes, n_points)
- """
- device = next(self.parameters()).device
- x_vals = torch.linspace(x_range[0], x_range[1], n_points, device=device)
- probe = x_mean.unsqueeze(0).expand(n_points, -1).clone().to(device)
- probe[:, dim] = x_vals
- with torch.no_grad():
- out = self.forward(probe) # (n_points, num_classes)
- return x_vals.cpu(), out.T.cpu() # (num_classes, n_points)
- def get_top_activation_curves(self, top_k: int = 10, n_points: int = 200):
- """
- Returns activation curves for the top-k most important input dimensions.
- Returns:
- top_dims: (top_k,)
- x_vals: (n_points,)
- curves: (top_k, out_features, n_points)
- """
- importance = self.get_input_importance()
- top_dims = importance.topk(top_k).indices
- curves, x_vals = [], None
- for d in top_dims.tolist():
- xv, yv = self.layer1.get_activation_curve(d, n_points)
- x_vals = xv
- curves.append(yv)
- return top_dims, x_vals, torch.stack(curves, dim=0)
- if __name__ == "__main__":
- torch.manual_seed(0)
- b, in_f, hidden, num_cls = 16, 128, 64, 4
- model = KANClassifier(in_f, hidden, num_cls)
- x = torch.randn(b, in_f)
- logits = model(x)
- print("Logits shape:", logits.shape) # (16, 4)
- imp = model.get_input_importance()
- print("Importance shape:", imp.shape) # (128,)
- top_dims, x_vals, curves = model.get_top_activation_curves(top_k=5)
- print("Top dims:", top_dims.tolist())
- print("Curves shape:", curves.shape) # (5, 64, 200)
kan_classifier.py at commit 0ff30c0, no license · at the source
Overview
- Artificial Intelligence College, Zhejiang Industry & Trade Vocational College, Wenzhou 325000, China
- College of Computer Science and Artificial Intelligence, Wenzhou University, Wenzhou 325035, China
Abstract
Background/
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 17 matches between paragraphs and lines of code.
l1binn/MambaKAN
0ff30c0a3602f29e7eb6b90543f4ef7978b34707, 17 April 2026Availability: 1 check, the latest on 29 September 2026: the link answers
- 29 September 2026: the link answers
10 files
- analysis.py, Python, 661 lines, 4 matches
- demo.py, Python, 107 lines
- models/
__init__.py , Python, 4 lines - models/
kan_classifier.py , Python, 297 lines, 4 matches - models/
mamba_encoder.py , Python, 249 lines, 4 matches - models/
proposed.py , Python, 215 lines - models/
vae.py , Python, 65 lines, 1 match - train_stage1.py, Python, 133 lines, 1 match
- train_stage2.py, Python, 263 lines, 3 matches
- README.md, Text, 261 lines
The paper's code and data availability statement is in the Data section.
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;
- 9 scripts, each with its path and the digest of its content;
- 17 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 ADNI dataset is publicly available at https://
Reproduced under the paper's license (CC BY), from the paper cited above.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 29 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 2 authors, 10 keywords, 5 funders, 20 references.
Cite
This paper
Gao, L., & Hu, Z. (2026). MambaKAN: An Interpretable Framework for Alzheimer's Disease Diagnosis via Selective State Space Modeling of Dynamic Functional Connectivity. Brain sciences, 16(4), 421. https://
BibTeX
@article{gao2026mambakan
author = {Gao, Libin and Hu, Zhongyi},
title = {{MambaKAN: An Interpretable Framework for Alzheimer's Disease Diagnosis via Selective State Space Modeling of Dynamic Functional Connectivity}},
journal = {Brain sciences},
year = {2026},
month = apr,
volume = {16},
number = {4},
pages = {421},
publisher = {Multidisciplinary Digital Publishing Institute (MDPI)},
issn = {2076-3425},
doi = {10.3390/
url = {https://
pmid = {42041829},
pmcid = {PMC13114612}
}
RIS
TY - JOUR
AU - Gao, Libin
AU - Hu, Zhongyi
TI - MambaKAN: An Interpretable Framework for Alzheimer's Disease Diagnosis via Selective State Space Modeling of Dynamic Functional Connectivity
T2 - Brain sciences
J2 - Brain Sci
PY - 2026
DA - 2026/
VL - 16
IS - 4
SP - 421
SN - 2076-3425
PB - Multidisciplinary Digital Publishing Institute (MDPI)
DO - 10.3390/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.3390/
"type": "article-journal",
"title": "MambaKAN: An Interpretable Framework for Alzheimer's Disease Diagnosis via Selective State Space Modeling of Dynamic Functional Connectivity",
"container-title": "Brain sciences",
"author": [
{
"family": "Gao",
"given": "Libin"
},
{
"family": "Hu",
"given": "Zhongyi"
}
],
"container-title-short":
"volume": "16",
"issue": "4",
"page": "421",
"DOI": "10.3390/
"PMID": "42041829",
"PMCID": "PMC13114612",
"ISSN": "2076-3425",
"publisher": "Multidisciplinary Digital Publishing Institute (MDPI)",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
17
]
]
}
}
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.1002/hbm.70483 [code]
- Untamed: Unconstrained Tensor Decomposition and Graph Node Embedding for Cortical Parcellation.Journal: Human brain mappingIn common: PyTorch, scikit-learn, SciPy, 2 other tools, fMRI, 1 reference
- [2] doi:10.1038/s41467-026-73687-9 [code]
- Automatic selection of the best neural architecture for time series forecasting.Journal: Nature communicationsIn common: PyTorch, SciPy, Matplotlib, 1 other tool, 2 references
- [3] doi:10.1002/alz.71365 [code]
- Benchmarking speech biomarkers of Alzheimer's against cognitive and neural measures.Journal: Alzheimer's & dementia : the journal of the Alzheimer's AssociationIn common: scikit-learn, SciPy, Matplotlib, 1 other tool, fMRI, Alzheimer's / dementia, 1 reference
- [4] doi:10.1038/s41586-026-10528-1 [code]
- A critical initialization for biological neural networks.Journal: NatureIn common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference
- [5] doi:10.1162/imag.a.1220 [code]
- Brain functional network connectivity interpolation characterizes the neuropsychiatric continuum and heterogeneity.Journal: Imaging neuroscience (Cambridge, Mass.)In common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference
- [6] doi:10.1038/s41598-026-68186-2 [code]
- NeuroStream: spectral-spatio-temporal
deep learning for visual stimulus classification from EEG. Journal: Scientific reportsIn common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference - [7] doi:10.1038/s43856-026-01817-x [code]
- Visual prompt engineering for multimodal and irregularly sampled medical data.Journal: Communications medicineIn common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference
- [8] doi:10.1098/rstb.2024.0461 [code]
- Shallow recurrent decoders for neural and behavioural dynamics.Journal: Philosophical transactions of the Royal Society of London. Series B, Biological sciencesIn common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference
- [9] doi:10.3390/bioengineering13080924 [code]
- Deep Learning-Based Temporal Gait Analysis Using a Smartphone IMU in Older Adults with and Without Non-Specific Low Back Pain.Journal: Bioengineering (Basel, Switzerland)In common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference
- [10] doi:10.1038/s41467-026-75783-2 [code]
- Broadband encoding and high-speed probabilistic bit generation with integrated microwave neurons.Journal: Nature communicationsIn common: PyTorch, scikit-learn, SciPy, 2 other tools, 1 reference
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
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, 9 scripts, and 17 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:99ecccd7df2c9a88…
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.
