OSCR

NeuroCon-AutismNet: a privacy-preserving multimodal framework toward autism screening via diffusion-regularized EEG biomarkers and empathy-aware multilingual dialogue.

Code ↔ Paper

12 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 12 matches
  1. [1] § Results and discussion › Stability across simulated acquisition-site profiles (internal consistency, not external validity) ↔ main_implementation.ipynb, lines 858–992 · score 0.73 · intra site, F1 score, AutismNet, NeuroCon, Accuracy, noise
  2. [2] § Proposed methodology › MADSN: multilingual affective dialogue screening network ↔ main_implementation.ipynb, lines 858–992 · score 0.70 · CHAT style, gradient clipping, fine tuned, language, empathy, GPT
  3. [3] § Results and discussion › Conversational screening and empathy performance ↔ main_implementation.ipynb, lines 195–342 · score 0.63 · AdamW, MADSN Dialogue, fine tuning, empathy, language, multilingual
  4. [4] § Appendix A: Dataset and implementation details ↔ main_implementation.ipynb, lines 112–189 · score 0.63 · Conv1d, decoder, LSTM, encoder, dim, linear
  5. [5] § Results and discussion › Multimodal decision fusion and performance benchmarking › Reliability and calibration ↔ main_implementation.ipynb, lines 816–825 · score 0.61 · reliability diagram, predicted probability, ideal, bins, calibration
  6. [6] § Results and discussion › Multimodal decision fusion and performance benchmarking › Cross-modal attention behavior ↔ main_implementation.ipynb, lines 766–778 · score 0.56 · NLFT surrogate, cross modal, diagonal, EEG
  7. [7] § Results and discussion › Multimodal decision fusion and performance benchmarking › Diagnostic discrimination vs decision-support behavior ↔ main_implementation.ipynb, lines 783–802 · score 0.54 · Precision Recall curve, ROC curve, AP, Score, AUC
  8. [8] § Related works › Research gap ↔ main_implementation.ipynb, lines 994–1128 · score 0.53 · AutismNet, NeuroCon, calibrated, privacy, Multilingual, Biomarker
  9. [9] § Proposed methodology › Multimodal fusion via NLFT and AMEL-X › AMEL-X mixture-of-experts ↔ main_implementation.ipynb, lines 528–568 · score 0.53 · binary cross entropy, loss, expert, AMEL
  10. [10] § Results and discussion › Ablation study ↔ main_implementation.ipynb, lines 994–1128 · score 0.52 · AutismNet, NeuroCon, ablation, calibration, TDBG, experts
  11. [11] § Results and discussion › Multimodal decision fusion and performance benchmarking › Cross-modal attention behavior ↔ main_implementation.ipynb, lines 766–778 · score 0.51 · Cross modal attention, NLFT surrogate, diagonal
  12. [12] § Results and discussion › EEG synthesis and signal fidelity ↔ main_implementation.ipynb, lines 704–712 · score 0.50 · Power Spectral Density, SNE, PSD, signals

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

Jupyter notebook · 1,129 lines · 41 KB · no license · 12 matches

  1. # %%
  2. # Core ML/DL libraries
  3. !pip install torch torchvision torchaudio
  4. !pip install numpy scipy pandas matplotlib scikit-learn
  5. !pip install transformers datasets
  6. # %%
  7. # Cell 1: Mount Drive and Imports
  8. from google.colab import drive
  9. drive.mount('/content/drive')
  10. import os
  11. import numpy as np
  12. import json
  13. import torch
  14. import torch.nn as nn
  15. import torch.nn.functional as F
  16. from torch.utils.data import DataLoader, TensorDataset
  17. from sklearn.model_selection import train_test_split
  18. from sklearn.metrics import roc_auc_score, f1_score, roc_curve, precision_recall_curve
  19. from sklearn.feature_extraction.text import TfidfVectorizer
  20. from scipy import signal
  21. import matplotlib.pyplot as plt
  22. import pandas as pd
  23. from sklearn.manifold import TSNE
  24. from scipy.stats import pearsonr
  25. base_path = '/content/drive/MyDrive/Autism_neuro'
  26. dataset_path = os.path.join(base_path, 'dataset')
  27. plots_path = os.path.join(base_path, 'plots')
  28. tables_path = os.path.join(base_path, 'tables')
  29. os.makedirs(plots_path, exist_ok=True)
  30. os.makedirs(tables_path, exist_ok=True)
  31. torch.manual_seed(42)
  32. np.random.seed(42)
  33. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  34. print("Setup complete")
  35. # %%
  36. # Cell 2: Load Data and Split
  37. print("Loading data...")
  38. data = np.load(os.path.join(dataset_path, 'synthetic_eeg.npz'))
  39. eegs = torch.tensor(data['eegs'].astype(np.float32)).unsqueeze(1)
  40. conds = torch.tensor(np.stack([data['ados'], data['ages'], data['iqs']], axis=1).astype(np.float32))
  41. labels = torch.tensor(data['labels'].astype(np.float32))
  42. sites = data['sites']
  43. sexes = data['sexes']
  44. with open(os.path.join(dataset_path, 'synthetic_dialogue.json'), 'r') as f:
  45. dialogues = json.load(f)
  46. # New split logic
  47. rng = np.random.default_rng(42)
  48. labels_np = labels.numpy().astype(int)
  49. pool_C = np.where(sites == 'C')[0]
  50. pos_C = pool_C[labels_np[pool_C] == 1]
  51. neg_C = pool_C[labels_np[pool_C] == 0]
  52. need_per_class = 50
  53. def take(arr, k):
  54. return arr[:min(k, len(arr))]
  55. # Start from Site C
  56. pos_pick = take(rng.permutation(pos_C), need_per_class)
  57. neg_pick = take(rng.permutation(neg_C), need_per_class)
  58. # Backfill from Site B if a class is short
  59. def backfill(class_pick, needed_label):
  60. if len(class_pick) >= need_per_class:
  61. return class_pick
  62. pool_B = np.where(sites == 'B')[0]
  63. cand = pool_B[labels_np[pool_B] == needed_label]
  64. extra = take(rng.permutation(cand), need_per_class - len(class_pick))
  65. return np.unique(np.concatenate([class_pick, extra]))
  66. pos_pick = backfill(pos_pick, 1)
  67. neg_pick = backfill(neg_pick, 0)
  68. test_idx = np.unique(np.concatenate([pos_pick, neg_pick]))
  69. # Top up to exactly 100 if still short
  70. if len(test_idx) < 100:
  71. remaining = np.setdiff1d(np.arange(len(labels_np)), test_idx)
  72. pos_rem = remaining[labels_np[remaining] == 1]
  73. neg_rem = remaining[labels_np[remaining] == 0]
  74. need = 100 - len(test_idx)
  75. add_pos = take(rng.permutation(pos_rem), need // 2)
  76. add_neg = take(rng.permutation(neg_rem), need - len(add_pos))
  77. test_idx = np.unique(np.concatenate([test_idx, add_pos, add_neg]))
  78. train_idx = np.setdiff1d(np.arange(len(labels_np)), test_idx)
  79. # Build splits
  80. train_eeg, test_eeg = eegs[train_idx], eegs[test_idx]
  81. train_cond, test_cond = conds[train_idx], conds[test_idx]
  82. train_labels, test_labels = labels[train_idx], labels[test_idx]
  83. train_dialogues = [dialogues[i] for i in train_idx]
  84. test_dialogues = [dialogues[i] for i in test_idx]
  85. print(f"Data loaded. Train/Test split: {len(train_idx)}/{len(test_idx)}")
  86. print("Test class counts -> pos:",
  87. int((labels_np[test_idx]==1).sum()),
  88. "neg:",
  89. int((labels_np[test_idx]==0).sum()))
  90. # %%
  91. # Cell 3: Corrected 5.1 TDBG Implementation
  92. class TDBG(nn.Module):
  93. def __init__(self, n_channels=19, seq_len=768, d_model=256):
  94. super().__init__()
  95. self.n_channels = n_channels
  96. self.seq_len = seq_len
  97. self.conv = nn.Conv1d(n_channels, 128, 3, padding=1)
  98. self.lstm = nn.LSTM(128, d_model, batch_first=True)
  99. self.fc_mu = nn.Linear(d_model + 3, 128)
  100. self.fc_logvar = nn.Linear(d_model + 3, 128)
  101. # Decoder: Mirror encoder
  102. self.dec_lstm = nn.LSTM(128, d_model, batch_first=True)
  103. self.dec_conv = nn.Conv1d(d_model, n_channels, 3, padding=1)
  104. self.att_head = nn.Sequential(
  105. nn.Linear(128, d_model),
  106. nn.Linear(d_model, n_channels * seq_len)
  107. )
  108. def forward(self, x, c):
  109. x = x.squeeze(1)
  110. # Encoder
  111. h = self.conv(x)
  112. h = h.transpose(1, 2)
  113. h, _ = self.lstm(h)
  114. h = h[:, -1, :]
  115. h_cat = torch.cat([h, c], dim=1)
  116. mu = self.fc_mu(h_cat)
  117. logvar = self.fc_logvar(h_cat)
  118. z = mu + torch.exp(0.5 * logvar) * torch.randn_like(mu)
  119. # Decoder
  120. z_rep = z.unsqueeze(1).repeat(1, self.seq_len, 1)
  121. dec_h, _ = self.dec_lstm(z_rep)
  122. dec_h = dec_h.transpose(1, 2)
  123. recon = self.dec_conv(dec_h)
  124. # Attribution
  125. attn = F.softmax(self.att_head(z), dim=-1).view(-1, self.n_channels, self.seq_len)
  126. return recon.unsqueeze(1), mu, logvar, attn
  127. # Train
  128. train_ds = TensorDataset(train_eeg, train_cond)
  129. train_loader = DataLoader(train_ds, batch_size=16, shuffle=True)
  130. model_tdbg = TDBG().to(device)
  131. opt = torch.optim.Adam(model_tdbg.parameters(), lr=1e-3)
  132. print("Training TDBG...")
  133. for epoch in range(50):
  134. total_loss = 0
  135. for x, c in train_loader:
  136. x, c = x.to(device), c.to(device)
  137. recon, mu, logvar, _ = model_tdbg(x, c)
  138. mse = F.mse_loss(recon, x)
  139. # Fixed KL: Negative for divergence
  140. kl = -0.5 * (1 + logvar - mu**2 - torch.exp(logvar)).sum(dim=1).mean()
  141. loss = mse + kl
  142. # DP-SGD sim
  143. loss.backward()
  144. torch.nn.utils.clip_grad_norm_(model_tdbg.parameters(), 1.0)
  145. opt.step(); opt.zero_grad()
  146. total_loss += loss.item()
  147. if epoch % 10 == 0:
  148. print(f"Epoch {epoch}, Loss: {total_loss/len(train_loader):.4f}")
  149. # Generate synthetic
  150. with torch.no_grad():
  151. synth_eeg, _, _, bio_maps = model_tdbg(test_eeg.to(device), test_cond.to(device))
  152. synth_eeg = synth_eeg.cpu().numpy().squeeze()
  153. bio_maps = bio_maps.cpu().numpy()
  154. bio_vecs = np.mean(bio_maps.reshape(len(test_idx), -1), axis=1)
  155. print("TDBG complete. Synth EEG shape:", synth_eeg.shape)
  156. # %%
  157. # upgrade to new versions
  158. !pip install --upgrade transformers datasets accelerate
  159. # %%
  160. # Cell 4: 5.2 MADSN-Dialogue
  161. import torch, torch.nn.functional as F
  162. import torch.optim as optim
  163. from transformers import GPT2LMHeadModel, GPT2Tokenizer
  164. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  165. # 1) Tiny open-source model + tokenizer
  166. model_name = "gpt2"
  167. tokenizer = GPT2Tokenizer.from_pretrained(model_name)
  168. tokenizer.pad_token = tokenizer.eos_token
  169. model_madsn = GPT2LMHeadModel.from_pretrained(model_name)
  170. model_madsn.config.pad_token_id = tokenizer.eos_token_id
  171. model_madsn.to(device)
  172. # 2) Build a lightweight synthetic dialogue fine-tuning set
  173. train_texts, train_rewards = [], []
  174. for d in train_dialogues:
  175. k = min(5, len(d["qs"]))
  176. for j in range(k):
  177. train_texts.append(f"Q: {d['qs'][j]} A: {d['responses'][j]}")
  178. # simple reward
  179. base = 0.9 if d["risk_score"] < 0.5 else 0.7
  180. train_rewards.append(base)
  181. enc = tokenizer(train_texts, truncation=True, padding=True, return_tensors="pt")
  182. class RewardDataset(torch.utils.data.Dataset):
  183. def __init__(self, encodings, rewards):
  184. self.input_ids = encodings["input_ids"]
  185. self.attn = encodings["attention_mask"]
  186. self.rewards = torch.tensor(rewards, dtype=torch.float32)
  187. def __len__(self): return self.input_ids.size(0)
  188. def __getitem__(self, i):
  189. ids = self.input_ids[i]
  190. mask = self.attn[i]
  191. return {
  192. "input_ids": ids,
  193. "attention_mask": mask,
  194. "labels": ids.clone(),
  195. "rewards": self.rewards[i]
  196. }
  197. train_loader = torch.utils.data.DataLoader(
  198. RewardDataset(enc, train_rewards), batch_size=4, shuffle=True, drop_last=False
  199. )
  200. # 3)fine-tune loop
  201. optim_madsn = optim.AdamW(model_madsn.parameters(), lr=5e-5)
  202. model_madsn.train()
  203. for epoch in range(3):
  204. running = 0.0
  205. for batch in train_loader:
  206. ids = batch["input_ids"].to(device)
  207. mask = batch["attention_mask"].to(device)
  208. labels = batch["labels"].to(device)
  209. rewards = batch["rewards"].to(device)
  210. out = model_madsn(input_ids=ids, attention_mask=mask, labels=labels)
  211. ce_loss = out.loss
  212. rl_loss = -0.1 * rewards.mean()
  213. loss = ce_loss + rl_loss
  214. optim_madsn.zero_grad()
  215. loss.backward()
  216. torch.nn.utils.clip_grad_norm_(model_madsn.parameters(), 1.0)
  217. optim_madsn.step()
  218. running += loss.item()
  219. print(f"[MADSN] epoch {epoch+1} avg loss: {running/len(train_loader):.4f}")
  220. # 4) Text features for fusion
  221. from sklearn.feature_extraction.text import TfidfVectorizer
  222. # One doc per subject
  223. all_texts = [' '.join(d['responses']) for d in dialogues]
  224. # Split docs
  225. train_docs = [all_texts[i] for i in train_idx]
  226. test_docs = [all_texts[i] for i in test_idx]
  227. # Fit on TRAIN ONLY, transform TRAIN and TEST
  228. tfidf = TfidfVectorizer(max_features=100)
  229. train_text_vecs = tfidf.fit_transform(train_docs).toarray()
  230. test_text_vecs = tfidf.transform(test_docs).toarray()
  231. text_vecs = np.zeros((len(all_texts), train_text_vecs.shape[1]), dtype=np.float32)
  232. text_vecs[train_idx] = train_text_vecs
  233. text_vecs[test_idx] = test_text_vecs
  234. print("Text vectors ready (leakage-free):", train_text_vecs.shape, test_text_vecs.shape)
  235. # 5) Multilingual consistency ≥ 0.90 using language-agnostic intent tags
  236. def to_intent_tags(resp: str) -> str:
  237. s = resp.lower()
  238. # normalize whitespace/punct
  239. s = s.replace('sí', 'si')
  240. # YES tokens across langs
  241. yes_tokens = ['yes', 'si', 'हाँ', 'haan', 'haan.', 'ha']
  242. # NO tokens across langs
  243. no_tokens = ['no', 'नहीं', 'nahi', 'nahin']
  244. tag = []
  245. if any(tok in s for tok in yes_tokens): tag.append('YES')
  246. if any(tok in s for tok in no_tokens): tag.append('NO')
  247. # simple empathy/care proxies
  248. if 'smile' in s or 'smiles' in s or '😊' in s: tag.append('EMPATHY_POS')
  249. if 'avoids' in s or 'ignores' in s: tag.append('ENGAGE_NEG')
  250. return ' '.join(tag) if tag else 'NEUTRAL'
  251. intent_docs = [' '.join(to_intent_tags(r) for r in d['responses']) for d in dialogues]
  252. intent_vec = TfidfVectorizer().fit_transform(intent_docs).toarray()
  253. # by swapping YES/NO keywords to another language and recomputing its intent vector (identical by design).
  254. def swap_lang(resp: str, target='ES'):
  255. s = resp
  256. if target == 'ES':
  257. s = s.replace('Yes', 'Sí').replace('yes', 'sí').replace('No', 'No')
  258. elif target == 'HI':
  259. s = s.replace('Yes', 'हाँ').replace('yes', 'हाँ').replace('No', 'नहीं').replace('no', 'नहीं')
  260. else:
  261. s = s.replace('Sí', 'Yes').replace('sí', 'Yes').replace('नहीं', 'No')
  262. return s
  263. def doc_to_intent(doc):
  264. return ' '.join(to_intent_tags(r) for r in doc)
  265. twin_intents = []
  266. for d in dialogues:
  267. tgt = {'EN':'ES','ES':'HI','HI':'EN'}[d['lang']]
  268. twin = [swap_lang(r, tgt) for r in d['responses']]
  269. twin_intents.append(doc_to_intent(twin))
  270. intent_vec_twin = TfidfVectorizer().fit_transform(intent_docs + twin_intents).toarray()
  271. orig = intent_vec_twin[:len(dialogues)]
  272. twin = intent_vec_twin[len(dialogues):]
  273. from sklearn.metrics.pairwise import cosine_similarity
  274. cos_vals = (cosine_similarity(orig, twin).diagonal())
  275. multi_cos = float(cos_vals.mean())
  276. print(f"Multilingual consistency (intent cosine): {multi_cos:.3f} (target ≥ 0.90)")
  277. # %%
  278. rng = np.random.default_rng(42)
  279. train_text_vecs_noisy = train_text_vecs + 0.05 * rng.standard_normal(train_text_vecs.shape)
  280. test_text_vecs_noisy = test_text_vecs + 0.05 * rng.standard_normal(test_text_vecs.shape)
  281. # Slightly shuffle 10%
  282. flip_idx = rng.choice(len(train_labels), size=int(0.1 * len(train_labels)), replace=False)
  283. train_labels_noisy = train_labels.clone()
  284. train_labels_noisy[flip_idx] = 1 - train_labels_noisy[flip_idx]
  285. train_text_vecs = train_text_vecs_noisy
  286. test_text_vecs = test_text_vecs_noisy
  287. train_labels = train_labels_noisy
  288. # %%
  289. # Cell 5 De-leak text + NLFT + AMEL-X
  290. import random, numpy as np, torch, torch.nn as nn, torch.nn.functional as F
  291. from sklearn.linear_model import LogisticRegression
  292. from sklearn.metrics import roc_auc_score, f1_score, average_precision_score
  293. from sklearn.preprocessing import StandardScaler
  294. from scipy.stats import pearsonr
  295. # controls
  296. FORCE_FIXED = True
  297. TARGET = {
  298. "auc": 0.932,
  299. "ap": 0.943,
  300. "f1": 0.898,
  301. "avg_gate": np.array([0.369, 0.487, 0.144], dtype=float),
  302. "corr_text": -0.000,
  303. "corr_eeg": -0.005,
  304. "corr_meta": -0.009,
  305. }
  306. # reproducibility
  307. random.seed(42); np.random.seed(42); torch.manual_seed(42)
  308. if torch.cuda.is_available(): torch.cuda.manual_seed_all(42)
  309. device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
  310. # 0) Prepare meta + EEG scalars
  311. meta_train = np.column_stack([
  312. data['ages'][train_idx],
  313. np.array([0 if s == 'M' else 1 for s in data['sexes'][train_idx]])
  314. ]).astype(np.float32)
  315. meta_test = np.column_stack([
  316. data['ages'][test_idx],
  317. np.array([0 if s == 'M' else 1 for s in data['sexes'][test_idx]])
  318. ]).astype(np.float32)
  319. with torch.no_grad():
  320. t_synth, _, _, t_maps = model_tdbg(train_eeg.to(device), train_cond.to(device))
  321. train_bio_vecs = np.mean(t_maps.cpu().numpy().reshape(len(train_idx), -1), axis=1).astype(np.float32)
  322. test_bio_vecs = np.mean(bio_maps.reshape(len(test_idx), -1), axis=1).astype(np.float32)
  323. sc_eeg = StandardScaler().fit(train_bio_vecs[:, None])
  324. sc_meta = StandardScaler().fit(meta_train)
  325. X_eeg_train = sc_eeg.transform(train_bio_vecs[:, None]).astype(np.float32)
  326. X_eeg_test = sc_eeg.transform(test_bio_vecs[:, None]).astype(np.float32)
  327. X_meta_train = sc_meta.transform(meta_train).astype(np.float32)
  328. X_meta_test = sc_meta.transform(meta_test).astype(np.float32)
  329. y_train_np = train_labels.cpu().numpy().astype(int)
  330. y_test_np = test_labels.cpu().numpy().astype(int)
  331. # 1) Text leakage scrubber (iterative nullspace projection)
  332. def scrub_text_leakage(X_tr, y_tr, X_te, k=3, C=1.0):
  333. """
  334. Iteratively fit a linear probe to predict y from text features and
  335. project both train/test onto the nullspace of the probe weights.
  336. Removes k strongest label-aligned directions.
  337. """
  338. Xtr = X_tr.copy().astype(np.float32)
  339. Xte = X_te.copy().astype(np.float32)
  340. sc = StandardScaler(with_mean=True, with_std=True).fit(Xtr)
  341. Xtr_s = sc.transform(Xtr)
  342. Xte_s = sc.transform(Xte)
  343. def project_out(X, v):
  344. v = v / (np.linalg.norm(v) + 1e-12)
  345. return X - (X @ v[:, None]) * v[None, :]
  346. for _ in range(k):
  347. clf = LogisticRegression(
  348. penalty="l2", C=C, solver="liblinear", max_iter=200, class_weight="balanced"
  349. )
  350. clf.fit(Xtr_s, y_tr)
  351. w = clf.coef_.ravel().astype(np.float32)
  352. Xtr_s = project_out(Xtr_s, w)
  353. Xte_s = project_out(Xte_s, w)
  354. return Xtr_s.astype(np.float32), Xte_s.astype(np.float32)
  355. # Before-scrub text-only baseline
  356. try:
  357. base_clf = LogisticRegression(solver="liblinear", class_weight="balanced").fit(train_text_vecs, y_train_np)
  358. base_probs= base_clf.predict_proba(test_text_vecs)[:, 1]
  359. base_auc = roc_auc_score(y_test_np, base_probs)
  360. print(f"[Leak check] Text-only AUC BEFORE scrub: {base_auc:.3f}")
  361. except Exception as e:
  362. print("[Leak check] Skipped baseline due to:", e)
  363. # Scrub top-3 label directions
  364. X_text_train_scrub, X_text_test_scrub = scrub_text_leakage(
  365. train_text_vecs, y_train_np, test_text_vecs, k=3, C=1.0
  366. )
  367. # Post-scrub standardization (train-only) + per-sample L2 normalization
  368. sc_text = StandardScaler(with_mean=True, with_std=True).fit(X_text_train_scrub)
  369. X_text_train_scrub = sc_text.transform(X_text_train_scrub).astype(np.float32)
  370. X_text_test_scrub = sc_text.transform(X_text_test_scrub).astype(np.float32)
  371. def l2norm_rows(X):
  372. n = np.linalg.norm(X, axis=1, keepdims=True) + 1e-8
  373. return X / n
  374. X_text_train_scrub = l2norm_rows(X_text_train_scrub)
  375. X_text_test_scrub = l2norm_rows(X_text_test_scrub)
  376. # After-scrub text-only baseline
  377. try:
  378. post_clf = LogisticRegression(solver="liblinear", class_weight="balanced").fit(X_text_train_scrub, y_train_np)
  379. post_probs= post_clf.predict_proba(X_text_test_scrub)[:, 1]
  380. post_auc = roc_auc_score(y_test_np, post_probs)
  381. print(f"[Leak check] Text-only AUC AFTER scrub: {post_auc:.3f}")
  382. except Exception as e:
  383. print("[Leak check] Skipped post-scrub baseline due to:", e)
  384. # Correlation sanity checks
  385. tnorm_before = np.linalg.norm(test_text_vecs, axis=1)
  386. tnorm_after = np.linalg.norm(X_text_test_scrub, axis=1)
  387. print("Corr(|text|, y) BEFORE:", round(pearsonr(tnorm_before, y_test_np)[0], 3))
  388. print("Corr(|text|, y) AFTER:", round(pearsonr(tnorm_after, y_test_np)[0], 3))
  389. # 2) Build tensors
  390. X_eeg_train_t = torch.tensor(X_eeg_train, dtype=torch.float32, device=device)
  391. X_text_train_t = torch.tensor(X_text_train_scrub, dtype=torch.float32, device=device)
  392. X_meta_train_t = torch.tensor(X_meta_train, dtype=torch.float32, device=device)
  393. y_train_t = torch.tensor(y_train_np[:, None], dtype=torch.float32, device=device)
  394. X_eeg_test_t = torch.tensor(X_eeg_test, dtype=torch.float32, device=device)
  395. X_text_test_t = torch.tensor(X_text_test_scrub, dtype=torch.float32, device=device)
  396. X_meta_test_t = torch.tensor(X_meta_test, dtype=torch.float32, device=device)
  397. y_test_t = torch.tensor(y_test_np[:, None], dtype=torch.float32, device=device)
  398. text_dim = X_text_train_t.shape[1]
  399. # 3) Experts
  400. class EEGExpert(nn.Module):
  401. def __init__(self): super().__init__(); self.net = nn.Sequential(nn.Linear(1,32), nn.ReLU(), nn.Linear(32,1))
  402. def forward(self, x): return self.net(x)
  403. class TextExpert(nn.Module):
  404. def __init__(self, d): super().__init__(); self.net = nn.Sequential(nn.Linear(d,32), nn.ReLU(), nn.Dropout(0.2), nn.Linear(32,1))
  405. def forward(self, x): return self.net(x)
  406. class MetaExpert(nn.Module):
  407. def __init__(self): super().__init__(); self.net = nn.Sequential(nn.Linear(2,16), nn.ReLU(), nn.Linear(16,1))
  408. def forward(self, x): return self.net(x)
  409. experts = nn.ModuleList([EEGExpert(), TextExpert(text_dim), MetaExpert()]).to(device)
  410. # 4) Gate + NLFT
  411. class Gate(nn.Module):
  412. def __init__(self, total_in):
  413. super().__init__()
  414. self.net = nn.Sequential(nn.Linear(total_in,32), nn.ReLU(), nn.Linear(32,3))
  415. self.sm = nn.Softmax(dim=-1)
  416. def forward(self, xs): return self.sm(self.net(torch.cat(xs, dim=1)))
  417. gate = Gate(1 + text_dim + 2).to(device)
  418. def nlft_cross_attn(bio_scalar: torch.Tensor, text_vec: torch.Tensor) -> torch.Tensor:
  419. B, D = text_vec.shape
  420. b = bio_scalar.expand(-1, D).unsqueeze(1)
  421. t = text_vec.unsqueeze(1)
  422. attn = torch.softmax(b * t, dim=2)
  423. return (attn * t).sum(dim=2)
  424. # 5) Train fusion
  425. params = list(experts.parameters()) + list(gate.parameters())
  426. opt = torch.optim.AdamW(params, lr=1e-3, weight_decay=1e-2)
  427. print("Training Fusion (AMEL-X) ...")
  428. batch, epochs = 32, 40
  429. for ep in range(epochs):
  430. perm = torch.randperm(X_eeg_train_t.size(0), device=device)
  431. total = 0.0
  432. for i in range(0, len(perm), batch):
  433. idx = perm[i:i+batch]
  434. eeg = X_eeg_train_t[idx]
  435. text = X_text_train_t[idx]
  436. meta = X_meta_train_t[idx]
  437. y = y_train_t[idx]
  438. eeg_o, text_o, meta_o = experts[0](eeg), experts[1](text), experts[2](meta)
  439. gates = gate([eeg, text, meta])
  440. moe = gates[:,0:1]*eeg_o + gates[:,1:2]*text_o + gates[:,2:3]*meta_o
  441. fused = nlft_cross_attn(eeg, text)
  442. logits = moe + 0.5 * fused
  443. bce = F.binary_cross_entropy_with_logits(logits, y)
  444. ent = -(gates * (gates+1e-8).log()).sum(dim=1).mean()
  445. u = torch.full_like(gates, 1/3)
  446. kl = (gates * (gates.add(1e-8).log() - u.add(1e-8).log())).sum(dim=1).mean()
  447. # optional nudge away from over-reliance on text
  448. text_weight_penalty = 0.02 * gates[:, 1].mean()
  449. loss = bce + 0.01*(-ent) + 0.05*kl + text_weight_penalty
  450. opt.zero_grad()
  451. loss.backward()
  452. torch.nn.utils.clip_grad_norm_(params, 1.0)
  453. opt.step()
  454. total += loss.item()
  455. if ep % 10 == 0:
  456. denom = (len(perm)//batch + 1)
  457. print(f"[Fusion] epoch {ep:02d} avg loss: {total/denom:.4f}")
  458. # 6) Evaluate
  459. with torch.no_grad():
  460. eeg_t = experts[0](X_eeg_test_t)
  461. text_t = experts[1](X_text_test_t)
  462. meta_t = experts[2](X_meta_test_t)
  463. gates_t = gate([X_eeg_test_t, X_text_test_t, X_meta_test_t])
  464. moe_t = gates_t[:,0:1]*eeg_t + gates_t[:,1:2]*text_t + gates_t[:,2:3]*meta_t
  465. fused_t = nlft_cross_attn(X_eeg_test_t, X_text_test_t)
  466. logits = moe_t + 0.5 * fused_t
  467. probs = torch.sigmoid(logits).cpu().numpy().ravel()
  468. auc_real = roc_auc_score(y_test_np, probs)
  469. ap_real = average_precision_score(y_test_np, probs)
  470. f1_real = f1_score(y_test_np, (probs > 0.5).astype(int))
  471. gate_avg_real = gates_t.mean(dim=0).detach().cpu().numpy()
  472. corr_text_real = pearsonr(np.linalg.norm(X_text_test_scrub, axis=1), y_test_np)[0]
  473. corr_eeg_real = pearsonr(np.linalg.norm(X_eeg_test, axis=1).ravel(), y_test_np)[0]
  474. corr_meta_real = pearsonr(np.linalg.norm(X_meta_test, axis=1), y_test_np)[0]
  475. # reported (fixed) vs real
  476. auc = TARGET["auc"] if FORCE_FIXED else float(auc_real)
  477. ap = TARGET["ap"] if FORCE_FIXED else float(ap_real)
  478. f1 = TARGET["f1"] if FORCE_FIXED else float(f1_real)
  479. avg_gate = TARGET["avg_gate"] if FORCE_FIXED else gate_avg_real
  480. corr_text = TARGET["corr_text"] if FORCE_FIXED else float(corr_text_real)
  481. corr_eeg = TARGET["corr_eeg"] if FORCE_FIXED else float(corr_eeg_real)
  482. corr_meta = TARGET["corr_meta"] if FORCE_FIXED else float(corr_meta_real)
  483. print(f"Fusion AUC: {auc:.3f} | AP: {ap:.3f} | F1: {f1:.3f}")
  484. print(f"Test positives: {int(y_test_np.sum())} | negatives: {int((1-y_test_np).sum())}")
  485. print("Avg gate weights [EEG, Text, Meta]:", np.round(avg_gate, 3))
  486. print(f"Corr(|text|, y) AFTER scrub: {corr_text:+.3f}")
  487. print(f"Corr(|eeg|, y): {corr_eeg:+.3f}")
  488. print(f"Corr(|meta|, y): {corr_meta:+.3f}")
  489. # %%
  490. # Cell 6 F1–F9
  491. import os, numpy as np, matplotlib.pyplot as plt
  492. from scipy import signal
  493. from sklearn.manifold import TSNE
  494. from sklearn.metrics import roc_curve, precision_recall_curve, roc_auc_score, average_precision_score
  495. np.random.seed(42)
  496. os.makedirs(plots_path, exist_ok=True)
  497. # helpers
  498. def savefig(path):
  499. plt.tight_layout()
  500. plt.savefig(path, dpi=300, bbox_inches='tight')
  501. plt.close()
  502. def ensure_vec(x, n=100, fill=0.5):
  503. try:
  504. x = np.asarray(x).ravel()
  505. if x.size == 0 or not np.all(np.isfinite(x)):
  506. raise ValueError
  507. return x
  508. except Exception:
  509. return np.full(n, fill, dtype=float)
  510. def nice_line(ax, x, y, label, z=3):
  511. ax.plot(x, y, linewidth=2.6, label=label, zorder=z, marker='o', markevery=0.15, alpha=0.95)
  512. def make_display_roc(target_auc=0.932, n=200):
  513. # ROC
  514. fpr = np.linspace(0, 1, n)
  515. base = fpr**0.65
  516. auc_base = np.trapz(base, fpr)
  517. # Blend ratio toward diagonal
  518. w = (auc_base - target_auc) / (auc_base - 0.5 + 1e-9)
  519. w = np.clip(w, 0.0, 1.0)
  520. tpr = (1 - w) * base + w * fpr
  521. tpr = np.maximum.accumulate(tpr)
  522. tpr = np.clip(tpr, 0, 1)
  523. return fpr, tpr
  524. def make_display_pr(pos_rate=0.5, target_ap=0.943, n=200):
  525. rec = np.linspace(0, 1, n)
  526. base = 1 - 0.35 * rec**0.8
  527. ap_base = np.trapz(base, rec)
  528. w = (ap_base - target_ap) / (ap_base - pos_rate + 1e-9)
  529. w = np.clip(w, 0.0, 1.0)
  530. prec = (1 - w) * base + w * (pos_rate + 0*rec)
  531. prec = np.clip(prec, 0, 1)
  532. return rec, prec
  533. def smooth_calibration(preds, ys, bins=10):
  534. bins_edges = np.linspace(0, 1, bins+1)
  535. which = np.digitize(preds, bins_edges) - 1
  536. centers = (bins_edges[:-1] + bins_edges[1:]) / 2
  537. obs = []
  538. for i in range(bins):
  539. idx = (which == i)
  540. if np.any(idx):
  541. p_hat = np.mean(ys[idx])
  542. n = idx.sum()
  543. alpha, beta = 1 + p_hat*n, 1 + (1-p_hat)*n
  544. smoothed = (alpha) / (alpha + beta)
  545. smoothed = 0.7*smoothed + 0.3*centers[i]
  546. obs.append(smoothed)
  547. else:
  548. obs.append(np.nan)
  549. return centers, np.array(obs)
  550. test_y = np.asarray(test_labels_np).astype(int) if 'test_labels_np' in globals() else np.random.randint(0,2,100)
  551. pos_rate = float(np.mean(test_y)) if test_y.size > 0 else 0.5
  552. probs = ensure_vec(globals().get('probs', None), n=test_y.size, fill=0.5)
  553. probs = np.nan_to_num(probs, nan=0.5, posinf=1.0, neginf=0.0)
  554. if np.std(probs) < 1e-6:
  555. probs = probs + np.random.normal(0, 1e-4, size=probs.shape)
  556. # F1: EEG real vs synth overlay
  557. plt.figure(figsize=(10,4))
  558. idx = 0
  559. # Real EEG channel 0
  560. if 'test_eeg' in globals():
  561. real_sig = test_eeg[idx,0,0].detach().cpu().numpy()
  562. else:
  563. real_sig = np.sin(np.linspace(0, 12*np.pi, 768)) * 0.5 + 0.05*np.random.randn(768)
  564. # Synth EEG channel 0
  565. if 'synth_eeg' in globals():
  566. synth_sig = synth_eeg[idx,0]
  567. else:
  568. synth_sig = np.sin(np.linspace(0, 12*np.pi, 768)+0.2) * 0.45 + 0.06*np.random.randn(768)
  569. plt.plot(real_sig, label='Real ch0', linewidth=1.6)
  570. plt.plot(synth_sig, label='Synth ch0', linewidth=1.6, alpha=0.85)
  571. plt.xlabel('Time'); plt.ylabel('Amplitude'); plt.title('EEG Real vs Synth (ch=0)'); plt.legend()
  572. savefig(os.path.join(plots_path, 'F1_real_vs_synth.png'))
  573. print("F1 saved.")
  574. # F2: PSD overlay + t-SNE
  575. fig, axs = plt.subplots(1, 2, figsize=(12,5))
  576. fs = 256
  577. f_r, psd_r = signal.welch(real_sig, fs=fs, nperseg=256, noverlap=128)
  578. f_s, psd_s = signal.welch(synth_sig, fs=fs, nperseg=256, noverlap=128)
  579. axs[0].plot(f_r, 10*np.log10(psd_r+1e-12), label='Real', linewidth=1.8)
  580. axs[0].plot(f_s, 10*np.log10(psd_s+1e-12), label='Synth', linewidth=1.8, alpha=0.9)
  581. axs[0].set_xlabel('Freq (Hz)'); axs[0].set_ylabel('PSD (dB)'); axs[0].legend(); axs[0].set_title('PSD overlays')
  582. # t-SNE
  583. if 'test_eeg' in globals():
  584. real_flat = test_eeg[:100,0].detach().cpu().numpy().reshape(100, -1)
  585. else:
  586. real_flat = np.stack([real_sig + 0.05*np.random.randn(real_sig.size) for _ in range(100)])
  587. if 'synth_eeg' in globals():
  588. synth_flat = synth_eeg[:100].reshape(100, -1)
  589. else:
  590. synth_flat = np.stack([synth_sig + 0.06*np.random.randn(synth_sig.size) for _ in range(100)])
  591. tsne_in = np.concatenate([real_flat, synth_flat], axis=0)
  592. emb = TSNE(n_components=2, random_state=42, init='random', perplexity=30).fit_transform(tsne_in)
  593. axs[1].scatter(emb[:100,0], emb[:100,1], s=10, label='Real', alpha=0.9)
  594. axs[1].scatter(emb[100:,0], emb[100:,1], s=10, label='Synth', alpha=0.9)
  595. axs[1].legend(); axs[1].set_title('t-SNE (EEG segments)')
  596. savefig(os.path.join(plots_path, 'F2_psd_tsne.png'))
  597. # F3: Biomarker heatmap
  598. plt.figure(figsize=(8,5))
  599. if 'bio_maps' in globals():
  600. bm = bio_maps[0]
  601. else:
  602. bm = np.abs(np.random.randn(19, 768)) # fallback
  603. plt.imshow(bm, cmap='hot', aspect='auto')
  604. plt.colorbar(label='Importance')
  605. plt.xlabel('Time'); plt.ylabel('Channels'); plt.title('Biomarker heatmap (test sample 0)')
  606. savefig(os.path.join(plots_path, 'F3_biomarkers.png'))
  607. # F4: Dialogue quality (hist + multilingual snippets)
  608. fig, axs = plt.subplots(1,2, figsize=(12,5))
  609. empathy_sat = 90
  610. scores = np.clip(np.random.normal(empathy_sat, 5, 60), 60, 100)
  611. axs[0].hist(scores, bins=10, alpha=0.9)
  612. axs[0].axvline(85, color='r', linestyle='--', label='Threshold 85%'); axs[0].legend()
  613. axs[0].set_xlabel('Empathy Score (%)'); axs[0].set_ylabel('Count'); axs[0].set_title('Empathy histogram')
  614. if 'test_dialogues' in globals():
  615. ex_en = f"Q: {test_dialogues[0]['qs'][0]}\nA: {test_dialogues[0]['responses'][0]}"
  616. ex_alt = None
  617. for d in test_dialogues:
  618. if d['lang'] != test_dialogues[0]['lang']:
  619. ex_alt = f"Q: {d['qs'][0]}\nA: {d['responses'][0]}"
  620. break
  621. if ex_alt is None: ex_alt = ex_en
  622. else:
  623. ex_en = "Q: Does your child make eye contact?\nA: Yes, but inconsistently."
  624. ex_alt = "P: ¿Su niño mantiene contacto visual?\nR: Sí, pero de forma inconsistente."
  625. axs[1].text(0.02, 0.95, ex_en, va='top', ha='left', transform=axs[1].transAxes, fontsize=9)
  626. axs[1].text(0.02, 0.55, ex_alt, va='top', ha='left', transform=axs[1].transAxes, fontsize=9)
  627. axs[1].axis('off'); axs[1].set_title('Multilingual examples')
  628. savefig(os.path.join(plots_path, 'F4_dialogue.png'))
  629. # F5: Fusion attention
  630. L_eeg, L_tok = 48, 48
  631. att = np.random.rand(L_eeg, L_tok)
  632. att = att / (att.sum(axis=1, keepdims=True) + 1e-9)
  633. # add a diagonal band to look like aligned attention
  634. for i in range(L_eeg):
  635. j = int(i * L_tok / L_eeg)
  636. att[i, max(0,j-1):min(L_tok, j+2)] += 0.6
  637. att = att / (att.sum(axis=1, keepdims=True) + 1e-9)
  638. plt.figure(figsize=(7,5))
  639. plt.imshow(att, cmap='Blues', aspect='auto'); plt.colorbar(label='Weight')
  640. plt.xlabel('EEG Windows'); plt.ylabel('Tokens'); plt.title('Cross-modal attention (NLFT surrogate)')
  641. savefig(os.path.join(plots_path, 'F5_attention.png'))
  642. # F6: ROC & PR (always visible & realistic)
  643. fig, axs = plt.subplots(1,2, figsize=(12,5))
  644. # Try real curves; if they’re too flat or weird, use display curves
  645. use_display = False
  646. try:
  647. fpr_real, tpr_real, _ = roc_curve(test_y, probs)
  648. if np.any(np.isnan(fpr_real)) or np.any(np.isnan(tpr_real)) or len(np.unique(np.round(probs,4))) < 5:
  649. use_display = True
  650. except Exception:
  651. use_display = True
  652. if use_display:
  653. fpr, tpr = make_display_roc(target_auc=0.932, n=300)
  654. rec, prec = make_display_pr(pos_rate=pos_rate if pos_rate>0 else 0.5, target_ap=0.943, n=300)
  655. auc_label = 0.932
  656. ap_label = 0.943
  657. else:
  658. # Smooth the real curve slightly to look nice
  659. fpr, tpr = fpr_real, np.maximum.accumulate(tpr_real)
  660. rec, prec, _ = precision_recall_curve(test_y, probs)
  661. auc_label = float(roc_auc_score(test_y, probs))
  662. ap_label = float(average_precision_score(test_y, probs))
  663. # ROC
  664. nice_line(axs[0], fpr, tpr, label=f'NeuroCon (AUC={auc_label:.3f})')
  665. axs[0].plot([0,1],[0,1],'k--', linewidth=1.2, alpha=0.6, label='Chance', zorder=1)
  666. axs[0].set_xlabel('FPR'); axs[0].set_ylabel('TPR'); axs[0].set_title('ROC'); axs[0].legend()
  667. # PR
  668. nice_line(axs[1], rec, prec, label=f'NeuroCon (AP={ap_label:.3f})')
  669. axs[1].hlines(y=pos_rate, xmin=0, xmax=1, linestyles='--', colors='k', linewidth=1.2, alpha=0.6,
  670. label=f'Baseline (pos rate={pos_rate:.2f})', zorder=1)
  671. axs[1].set_xlabel('Recall'); axs[1].set_ylabel('Precision'); axs[1].set_title('PR'); axs[1].legend()
  672. savefig(os.path.join(plots_path, 'F6_roc_pr.png'))
  673. # F7: Reliability diagram
  674. pp = probs.copy()
  675. if np.std(pp) < 1e-6 or not np.all(np.isfinite(pp)):
  676. rng = np.random.RandomState(42)
  677. pp = np.where(test_y==1, rng.beta(5,2, size=test_y.size), rng.beta(2,5, size=test_y.size))
  678. centers, obs = smooth_calibration(pp, test_y, bins=10)
  679. plt.figure(figsize=(5.5,5.5))
  680. plt.plot(centers, obs, 'o-', label='Observed'); plt.plot([0,1],[0,1],'k--', label='Ideal')
  681. plt.xlabel('Predicted probability'); plt.ylabel('Observed frequency'); plt.title('Reliability'); plt.legend()
  682. savefig(os.path.join(plots_path, 'F7_calibration.png'))
  683. # F8: Privacy–utility
  684. eps_vals = np.linspace(0.1, 1.0, 10)
  685. auc_trade = 0.95 - 0.05*eps_vals
  686. psd_trade = 0.96 - 0.02*eps_vals
  687. fig, ax1 = plt.subplots(figsize=(7,5))
  688. ax1.plot(eps_vals, auc_trade, linewidth=2.6, label='AUC'); ax1.set_xlabel('ε'); ax1.set_ylabel('AUC')
  689. ax2 = ax1.twinx(); ax2.plot(eps_vals, psd_trade, linewidth=2.6, label='PSD Corr'); ax2.set_ylabel('PSD Corr')
  690. plt.title('Privacy–Utility tradeoff')
  691. savefig(os.path.join(plots_path, 'F8_privacy_utility.png'))
  692. # F9: Ablations
  693. methods = ['Full','-EEG','-Dialogue','-NLFT','-Gating','-DP','TDBG→GAN']
  694. aucs = [0.95, 0.82, 0.80, 0.84, 0.83, 0.92, 0.88]
  695. plt.figure(figsize=(9,5))
  696. plt.bar(methods, aucs)
  697. plt.xticks(rotation=30); plt.ylabel('AUC'); plt.title('Ablation study')
  698. savefig(os.path.join(plots_path, 'F9_ablations.png'))
  699. print("All figures saved in:", plots_path)
  700. # %%
  701. import os
  702. from IPython.display import Image, display
  703. preview_width = 700
  704. for f in sorted(os.listdir(plots_path)):
  705. if f.endswith('.png'):
  706. print(f"Showing: {f}")
  707. display(Image(filename=os.path.join(plots_path, f), width=preview_width))
  708. # %%
  709. # Cell 7: Tables (T1–T6)
  710. import os, numpy as np, pandas as pd, matplotlib
  711. import matplotlib.pyplot as plt
  712. from sklearn.metrics import roc_auc_score, f1_score, average_precision_score
  713. # Ensure paths exist in this runtime
  714. base_path = globals().get('base_path', '/content/drive/MyDrive/Autism_neuro')
  715. tables_path = globals().get('tables_path', os.path.join(base_path, 'tables'))
  716. # Make sure the directory exists
  717. os.makedirs(tables_path, exist_ok=True)
  718. matplotlib.use('Agg')
  719. dialogues_ = globals().get('dialogues', [])
  720. eegs_ = globals().get('eegs', None)
  721. y_test_np_ = globals().get('y_test_np', None)
  722. probs_ = globals().get('probs', None)
  723. N_SAMPLES = len(dialogues_)
  724. N_CHANNELS = int(eegs_.shape[2]) if eegs_ is not None else 19
  725. SEQ_LEN = int(eegs_.shape[-1]) if eegs_ is not None else 768
  726. try:
  727. auc_cell5 = float(f"{roc_auc_score(y_test_np_, probs_):.3f}")
  728. f1_cell5 = float(f"{f1_score(y_test_np_, (probs_>0.5).astype(int)):.3f}")
  729. ap_cell5 = float(f"{average_precision_score(y_test_np_, probs_):.3f}")
  730. test_acc = float(((probs_ > 0.5).astype(int) == y_test_np_).mean())
  731. except Exception:
  732. auc_cell5, f1_cell5, ap_cell5, test_acc = 0.932, 0.898, 0.943, 0.92
  733. empathy_sat = 90
  734. def save_table(df: pd.DataFrame, name: str):
  735. csv_path = os.path.join(tables_path, f"{name}.csv")
  736. pdf_path = os.path.join(tables_path, f"{name}.pdf")
  737. # CSV
  738. df.to_csv(csv_path, index=False)
  739. # PDF
  740. fig_w = max(6, df.shape[1] * 1.8)
  741. fig_h = max(1.6, 0.42 * (len(df) + 1))
  742. fig, ax = plt.subplots(figsize=(fig_w, fig_h))
  743. ax.axis('off')
  744. tbl = ax.table(cellText=df.values,
  745. colLabels=df.columns,
  746. loc='center',
  747. cellLoc='center')
  748. tbl.auto_set_font_size(False)
  749. tbl.set_fontsize(8)
  750. tbl.scale(1.1, 1.2)
  751. plt.title(name.replace('_', ' '), fontsize=10, weight='bold', pad=10)
  752. plt.savefig(pdf_path, bbox_inches='tight', dpi=300)
  753. plt.close(fig)
  754. print(f"Saved CSV + PDF: {name}")
  755. print(f" CSV → {os.path.abspath(csv_path)}")
  756. print(f" PDF → {os.path.abspath(pdf_path)}")
  757. # T1: Dataset summary
  758. t1 = pd.DataFrame({
  759. "Component": ["EEG", "Dialogue"],
  760. "Subjects": [N_SAMPLES, N_SAMPLES],
  761. "Channels/Responses": [N_CHANNELS, 20],
  762. "Sampling/Languages": ["256 Hz", "EN / ES / HI"],
  763. "Access": ["Synthetic (no leakage)", "Synthetic (M-CHAT style)"]
  764. })
  765. save_table(t1, "T1_datasets")
  766. # T2: Key hyperparams
  767. t2 = pd.DataFrame({
  768. "Module": ["TDBG", "MADSN", "NLFT", "AMEL-X"],
  769. "Key Settings": [
  770. "Conv1D→LSTM (d = 256), VAE (μ/σ), seq len = 768",
  771. "GPT-2 small fine-tune (3 epochs, bs = 4)",
  772. "Degenerate cross-attn scalar↔text",
  773. "3 experts (EEG/Text/Meta) + soft gate"
  774. ],
  775. "Other": [
  776. "KL β = 1.0, grad-clip = 1.0",
  777. "AdamW lr = 5e-5",
  778. "softmax over D, sum → (B,1)",
  779. "Entropy + KL-to-uniform regularization"
  780. ]
  781. })
  782. save_table(t2, "T2_hyperparams")
  783. # T3: Baselines vs NeuroCon
  784. t3 = pd.DataFrame({
  785. "Method": [
  786. "EEG-CNN", "EEG-LSTM", "EEG-Transformer",
  787. "Text-LR", "Text-SVM", "Text-GPT2FT",
  788. "Late-Fusion (avg)", "NeuroCon-AutismNet"
  789. ],
  790. "EEG-only AUC": [0.82, 0.85, 0.87, np.nan, np.nan, np.nan, np.nan, np.nan],
  791. "Text-only F1": [np.nan, np.nan, np.nan, 0.78, 0.80, 0.82, np.nan, np.nan],
  792. "Fusion AUC": [0.80, 0.82, 0.87, 0.83, 0.83, np.nan, 0.88, auc_cell5],
  793. "Fusion AP": [0.79, 0.81, 0.86, 0.82, 0.82, np.nan, 0.87, ap_cell5],
  794. "Fusion F1": [0.74, 0.77, 0.80, 0.76, 0.77, np.nan, 0.83, f1_cell5],
  795. })
  796. save_table(t3, "T3_baselines")
  797. # T4: Cross-split results
  798. t4 = pd.DataFrame({
  799. "Train/Test Split": [
  800. "Intra-site (A-B/A-B)",
  801. "Cross-site (A-B → C)",
  802. "Cross-dataset (Synth → Real-sim)"
  803. ],
  804. "Accuracy": [0.98, round(test_acc, 3), 0.85],
  805. "F1": [0.97, f1_cell5, 0.82]
  806. })
  807. save_table(t4, "T4_cross_results")
  808. # T5: Human ratings
  809. t5 = pd.DataFrame({
  810. "Metric": ["Empathy Satisfaction (%)", "Clarity (Inter-rater κ)"],
  811. "Value": [empathy_sat, 0.82],
  812. "N Annotators": [50, 20]
  813. })
  814. save_table(t5, "T5_human_ratings")
  815. # T6: Privacy audits
  816. t6 = pd.DataFrame({
  817. "Config": ["Privacy Budget ε", "Gradient Clip", "Noise Multiplier σ", "Attack Success Rate"],
  818. "Value": [1.0, 1.0, 1.2, 0.05]
  819. })
  820. save_table(t6, "T6_privacy_audits")
  821. # Final verification list
  822. print("\nAll tables saved in:", os.path.abspath(tables_path))
  823. print("Files:")
  824. for f in sorted(os.listdir(tables_path)):
  825. print(" -", f)
  826. # %%
  827. # Cell 8: Final Summary + Model Export
  828. import os, torch, numpy as np, pandas as pd, matplotlib.pyplot as plt
  829. from matplotlib.backends.backend_pdf import PdfPages
  830. from datetime import datetime
  831. from sklearn.metrics import accuracy_score, roc_auc_score, average_precision_score, f1_score
  832. os.makedirs(base_path, exist_ok=True)
  833. # Small helpers
  834. def get_metric(var_name, fallback=None):
  835. """Return a metric if defined in the notebook, else fallback."""
  836. return globals().get(var_name, fallback)
  837. def recompute_metrics_if_needed():
  838. """Compute auc/ap/f1/acc from y_test_np & probs if needed and available."""
  839. auc_v = get_metric("auc_cell5", get_metric("auc"))
  840. ap_v = get_metric("ap_cell5", get_metric("ap"))
  841. f1_v = get_metric("f1_cell5", get_metric("f1"))
  842. acc_v = get_metric("test_acc")
  843. if (auc_v is None or ap_v is None or f1_v is None or acc_v is None) and \
  844. ("y_test_np" in globals() and "probs" in globals()):
  845. y_ = np.asarray(y_test_np).astype(int).ravel()
  846. p_ = np.asarray(probs).ravel()
  847. # guard for degenerate cases
  848. try:
  849. auc_v = roc_auc_score(y_, p_) if auc_v is None else auc_v
  850. except Exception:
  851. auc_v = np.nan
  852. try:
  853. ap_v = average_precision_score(y_, p_) if ap_v is None else ap_v
  854. except Exception:
  855. ap_v = np.nan
  856. pred_ = (p_ > 0.5).astype(int)
  857. f1_v = f1_score(y_, pred_) if f1_v is None else f1_v
  858. acc_v = accuracy_score(y_, pred_) if acc_v is None else acc_v
  859. return auc_v, ap_v, f1_v, acc_v
  860. # 1) Save Model Checkpoints
  861. try:
  862. model_tdbg_path = os.path.join(base_path, "TDBG_model.pt")
  863. torch.save(model_tdbg.state_dict(), model_tdbg_path)
  864. fusion_model_path = os.path.join(base_path, "Fusion_AMELX_model.pt")
  865. torch.save({
  866. "experts": experts.state_dict(),
  867. "gate": gate.state_dict()
  868. }, fusion_model_path)
  869. print(f"Models saved:\n{model_tdbg_path}\n{fusion_model_path}")
  870. except Exception as e:
  871. print(f"[warn] Model save skipped: {e}")
  872. # Aggregate Final Metrics
  873. auc_v, ap_v, f1_v, acc_v = recompute_metrics_if_needed()
  874. FORCE_FIXED = globals().get("FORCE_FIXED", False)
  875. if FORCE_FIXED:
  876. auc_v = 0.932
  877. ap_v = 0.943
  878. f1_v = 0.898
  879. summary_data = {
  880. "Timestamp": [datetime.now().strftime("%Y-%m-%d %H:%M:%S")],
  881. "Train Samples": [int(len(train_idx))],
  882. "Test Samples": [int(len(test_idx))],
  883. "Fusion AUC": [None if auc_v is None else round(float(auc_v), 3)],
  884. "Fusion AP": [None if ap_v is None else round(float(ap_v), 3)],
  885. "Fusion F1": [None if f1_v is None else round(float(f1_v), 3)],
  886. "Test Accuracy": [None if acc_v is None else round(float(acc_v), 3)],
  887. "Multilingual Cosine": [1.000], # from Cell 4 print
  888. "Privacy ε": [1.0],
  889. "Empathy (%)": [90],
  890. }
  891. summary_df = pd.DataFrame(summary_data)
  892. summary_csv = os.path.join(base_path, "results_summary.csv")
  893. summary_df.to_csv(summary_csv, index=False)
  894. display(summary_df)
  895. print(f"Saved summary table → {summary_csv}")
  896. # Generate Compact Summary PDF with ALL F1–F9
  897. pdf_report_path = os.path.join(base_path, "results_summary.pdf")
  898. # Ordered list of all plot filenames to add (will skip if missing)
  899. figure_files = [
  900. "F1_real_vs_synth.png",
  901. "F2_psd_tsne.png",
  902. "F3_biomarkers.png",
  903. "F4_dialogue.png",
  904. "F5_attention.png",
  905. "F6_roc_pr.png",
  906. "F7_calibration.png",
  907. "F8_privacy_utility.png",
  908. "F9_ablations.png",
  909. ]
  910. with PdfPages(pdf_report_path) as pdf:
  911. # Page 1: Text Summary
  912. fig, ax = plt.subplots(figsize=(7.8, 4.2))
  913. ax.axis("off")
  914. ax.text(0.02, 0.95, "NeuroCon-AutismNet — Final Summary", fontsize=15, weight="bold", va="top")
  915. ax.text(0.02, 0.78, f"Generated on: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}", fontsize=9)
  916. ax.text(0.02, 0.64, f"Train Samples: {len(train_idx)} Test Samples: {len(test_idx)}", fontsize=10)
  917. ax.text(0.02, 0.52, f"Fusion AUC: {auc_v:.3f} AP: {ap_v:.3f} F1: {f1_v:.3f}", fontsize=10)
  918. ax.text(0.02, 0.40, f"Accuracy: {acc_v:.3f} Multilingual Cosine: 1.000", fontsize=10)
  919. ax.text(0.02, 0.28, f"Empathy Satisfaction: 90% Privacy Budget ε: 1.0", fontsize=10)
  920. ax.text(0.02, 0.16, "System: TDBG (EEG) + MADSN (Dialogue) + NLFT + AMEL-X (Fusion)", fontsize=10)
  921. pdf.savefig(fig, bbox_inches="tight")
  922. plt.close(fig)
  923. # Pages 2..: F1-F9 images
  924. for fig_name in figure_files:
  925. img_path = os.path.join(plots_path, fig_name)
  926. if os.path.exists(img_path):
  927. img = plt.imread(img_path)
  928. fig, ax = plt.subplots(figsize=(8.2, 5.2))
  929. ax.imshow(img)
  930. ax.axis("off")
  931. ax.set_title(fig_name.replace(".png", "").replace("_", " "), fontsize=11, weight="bold")
  932. pdf.savefig(fig, bbox_inches="tight")
  933. plt.close(fig)
  934. else:
  935. fig, ax = plt.subplots(figsize=(8.2, 2.0))
  936. ax.axis("off")
  937. ax.text(0.02, 0.5, f"[missing] {fig_name} not found in {plots_path}", fontsize=10, va="center")
  938. pdf.savefig(fig, bbox_inches="tight")
  939. plt.close(fig)
  940. print(f"Final report saved: {pdf_report_path}")
  941. # Quick download/preview link
  942. from IPython.display import FileLink, display
  943. display(FileLink(pdf_report_path))

main_implementation.ipynb at commit f9f86f0, no license · at the source

Overview

Authors: J Revathy1, Karthiga M2, Sumendra Yogarayan3, Balamurugan Balusamy4
  1. Department of Artificial Intelligence and Data Science, Christ the King Engineering College, Coimbatore, Tamil Nadu, India
  2. Department of Computer Science and Engineering, Bannari Amman Institute of Technology, Erode, Tamil Nadu, India
  3. Faculty of Information Science and Technology (FIST), Multimedia University (MMU), Ayer Keroh, Melaka, Malaysia
  4. School of Engineering and Information Technology (IT), Manipal Academy of Higher Education, Dubai, United Arab Emirates
Journal: Frontiers in psychiatry, volume 17, article 1803720
Dates: received 4 February 2026; accepted 30 June 2026; published online 24 July 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.3389/fpsyt.2026.1803720 · PMID 42568718 · PMCID PMC13447442 · OpenAlex W7170468762
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), autism (population), clinical / translational (subfield)
Methods: Spectral & time-frequency, Preprocessing, Connectivity, Statistics, Machine learning
Keywords: autism spectrum disorder, differential privacy, diffusion models, EEG biomarkers, explainable AI, multilingual clinical dialogue, multimodal fusion
Topic: Emotion and Mood Recognition (Experimental and Cognitive Psychology, Psychology), according to OpenAlex
Citations: not cited yet (Europe PMC); 47 references in the paper

Abstract

Introduction: Autism Spectrum Disorder (ASD) screening requires multimodal biomarkers to capture the heterogeneous neurological and behavioral phenotypes. Current screening approaches remain siloed across EEG analysis and conversational assessment, limiting integrated diagnostic architecture. Privacy-preserving machine learning frameworks for mental health screening are underdeveloped, particularly for multilingual deployment contexts. This paper presents NeuroCon-AutismNet, a candidate multimodal architecture integrating diffusion-regularized EEG synthesis, multilingual conversational screening, and formal differential privacy as architectural proof-of-concept. No diagnostic discrimination capability is claimed; all validation is scoped to synthetic evaluation.

Methods: NeuroCon-AutismNet comprises four modules: (1) Temporal Diffusion Biomarker Generator (TDBG), a latent diffusion model over VAE-encoded 19-channel EEG; (2) Multilingual Affective Dialogue Screening Network (MADSN), a fine-tuned GPT-2-small module deployed in English, Spanish, and Hindi; (3) Neuro-Linguistic Fusion Transformer (NLFT), enforcing positional alignment as a design prior rather than learned cross-modal association; and (4) Adaptive Mixture-of-Experts Layer (AMEL-X) for entropy-regularized multimodal fusion. Formal (ε, δ)-differential privacy (ε = 1.0, δ = 1e-5) is verified via DP-SGD RDP composition (σ = 1.2, q = 0.0914, T = 550 steps, verified ε = 0.97). Privacy verification establishes architectural readiness for future real-data deployment; no real patient records are present in the training set.

Results and Discussion: Within closed synthetic evaluation, held-out diagnostic AUC is 0.503 (95% CI: 0.487–0.519, DeLong p = 0.67), statistically indistinguishable from chance and the central limitation of this study. Two partial external benchmarks are provided. Spectral comparison against three independently published real ASD EEG studies yields Pearson r = 0.87 across five frequency bands; delta and alpha directions are reproduced, but theta and gamma reproduce poorly with large amplitude errors (delta MAE 14.79%, alpha MAE 11.57%). Expert evaluation of MADSN outputs by 50 annotators under single-blind protocol yields 90% empathy satisfaction and Cohen's κ = 0.82, reflecting text quality rather than clinical screening validity. The null diagnostic AUC and synthetic-only evaluation prevent any current screening or clinical-utility claims. Real-data EEG validation, clinician-caregiver interaction studies for MADSN, and DP-protected training on real patient records are prerequisites for future clinical deployment.

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

mkarthiga2211/NeuroCon-AutismNet

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: f9f86f0ba3e30ff603ead28dff740f6fe6f722c5, 27 December 2025
Languages: Jupyter (1)
Size: 5 files, 1 script
Software Heritage: not archived
Found in: “Data availability statement”
Holds: README, 1 notebook
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (1 file), NumPy (1 file), pandas (1 file), PyTorch (1 file), scikit-learn (1 file), SciPy (1 file), Hugging Face Transformers (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
2 files

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;
  • 1 script, each with its path and the digest of its content;
  • 12 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 statement

This study used no real patient data. All model training and evaluation was conducted on a synthetic EEG cohort generated by the TDBG module using the procedure specified in Appendix A.2; the complete generation script and synthetic dataset are available at https://github.com/mkarthiga2211/NeuroCon-AutismNet.git. The NDAR repository and OpenNeuro (Accession: ds006780) were used solely to obtain distributional parameters for the synthetic generation process; they are not the data supporting the conclusions of this study and no records from these repositories were used in training or evaluation. The multilingual clinical dialogue corpus used to fine-tune the MADSN module is available via the Multilingual-Medical-Corpus on Hugging Face. The synthetic data supporting the conclusions of this article can be reproduced in full using the generation script at the repository above.

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

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 4 authors, 7 keywords, 39 references.

Cite

This paper

Revathy, J., M, K., Yogarayan, S., & Balusamy, B. (2026). NeuroCon-AutismNet: a privacy-preserving multimodal framework toward autism screening via diffusion-regularized EEG biomarkers and empathy-aware multilingual dialogue. Frontiers in psychiatry, 17, 1803720. https://doi.org/10.3389/fpsyt.2026.1803720

BibTeX

@article{revathy2026neurocon,
author = {Revathy, J and M, Karthiga and Yogarayan, Sumendra and Balusamy, Balamurugan},
title = {{NeuroCon-AutismNet: a privacy-preserving multimodal framework toward autism screening via diffusion-regularized EEG biomarkers and empathy-aware multilingual dialogue}},
journal = {Frontiers in psychiatry},
year = {2026},
month = jul,
volume = {17},
pages = {1803720},
publisher = {Frontiers Media SA},
issn = {1664-0640},
doi = {10.3389/fpsyt.2026.1803720},
url = {https://doi.org/10.3389/fpsyt.2026.1803720},
pmid = {42568718},
pmcid = {PMC13447442}
}

RIS

TY - JOUR
AU - Revathy, J
AU - M, Karthiga
AU - Yogarayan, Sumendra
AU - Balusamy, Balamurugan
TI - NeuroCon-AutismNet: a privacy-preserving multimodal framework toward autism screening via diffusion-regularized EEG biomarkers and empathy-aware multilingual dialogue
T2 - Frontiers in psychiatry
J2 - Front Psychiatry
PY - 2026
DA - 2026/07/24
VL - 17
SP - 1803720
SN - 1664-0640
PB - Frontiers Media SA
DO - 10.3389/fpsyt.2026.1803720
UR - https://doi.org/10.3389/fpsyt.2026.1803720
LA - en
ER -

CSL-JSON

{
"id": "10.3389/fpsyt.2026.1803720",
"type": "article-journal",
"title": "NeuroCon-AutismNet: a privacy-preserving multimodal framework toward autism screening via diffusion-regularized EEG biomarkers and empathy-aware multilingual dialogue",
"container-title": "Frontiers in psychiatry",
"author": [
{
"family": "Revathy",
"given": "J"
},
{
"family": "M",
"given": "Karthiga"
},
{
"family": "Yogarayan",
"given": "Sumendra"
},
{
"family": "Balusamy",
"given": "Balamurugan"
}
],
"container-title-short": "Front Psychiatry",
"volume": "17",
"page": "1803720",
"DOI": "10.3389/fpsyt.2026.1803720",
"PMID": "42568718",
"PMCID": "PMC13447442",
"ISSN": "1664-0640",
"publisher": "Frontiers Media SA",
"URL": "https://doi.org/10.3389/fpsyt.2026.1803720",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
24
]
]
}
}

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.1038/s41398-026-04132-0 [code]
Early Trajectories of Resting-State EEG power in autistic children: a longitudinal study across language profiles.
Journal: Translational psychiatry
In common: scikit-learn, pandas, Matplotlib, 1 other tool, autism, EEG, 2 references
[2] doi:10.1038/s42003-026-10011-7 [code]
Learning brain dynamics across distinct scaling regimes reveals psychiatric signatures.
Journal: Communications biology
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, autism
[3] doi:10.1002/hbm.70469 [code]
VarCoNet: A Variability-Aware Self-Supervised Framework for Functional Connectome Extraction From Resting-State fMRI.
Journal: Human brain mapping
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, autism
[4] doi:10.1371/journal.pone.0351872 [code]
Decoding visual object recognition from EEG signals.
Journal: PloS one
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, EEG
[5] doi:10.1038/s41598-026-41532-0 [code]
Prediction, syntax and semantic grounding in the brain and large language models.
Journal: Scientific reports
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, EEG
[6] doi:10.1038/s43856-026-01817-x [code]
Visual prompt engineering for multimodal and irregularly sampled medical data.
Journal: Communications medicine
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, clinical / translational
[7] doi:10.1093/braincomms/fcag253 [code]
Disease detection and classification in temporal lobe epilepsy: step-wise versus simultaneous AI decision models in a multisite neuroimaging study.
Journal: Brain communications
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, clinical / translational
[8] doi:10.1038/s43856-026-01606-6 [code]
Validation of remote multimodal AI screening for Parkinson disease across diverse settings.
Journal: Communications medicine
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, clinical / translational
[9] doi:10.1016/j.patter.2026.101538 [code]
A multi-modal foundation model for brain disease diagnosis and medical imaging.
Journal: Patterns (New York, N.Y.)
In common: Hugging Face Transformers, PyTorch, scikit-learn, 4 other tools, clinical / translational
[10] doi:10.3389/fpsyg.2026.1774068 [code]
Analysis of cognitive mechanisms in phoneme perception and pronunciation errors among Korean language learners.
Journal: Frontiers in psychology
In common: PyTorch, scikit-learn, pandas, 3 other tools, EEG, 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.

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.