OSCR

Arousal modulates functional connectivity through structured and hemispherically asymmetric community architecture during wakefulness.

Code ↔ Paper

7 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 7 matches
  1. [1] § Appendix 4 › Control analysis on the optimal cluster number ↔ validation.py, lines 493–583 · score 0.83 · linear regression fit, elbow point, squared error, RMSE, minimized, optimal
  2. [2] § Appendix 4 › Control analysis on the optimal cluster number ↔ validation.py, lines 65–90 · score 0.71 · Calinski Harabasz, Davies Bouldin, CH, Silhouette, Score, clustering
  3. [3] § Appendix 4 › Control analysis on the optimal cluster number ↔ validation.py, lines 493–583 · score 0.70 · linear regression fit, minimum RMSE, squared error, optimal, clustering, metrics
  4. [4] § Appendix 1 › Robustness analysis of split-half reliability and participant-level resampling ↔ validation.py, lines 744–809 · score 0.57 · Hungarian algorithm, Dice coefficients, reliability, alignment, validations, iteration
  5. [5] § Appendix 1 › Robustness analysis of split-half reliability and participant-level resampling ↔ validation.py, lines 1899–1985 · score 0.57 · absolute deviation, relative error, resampled, metric, network, coupling
  6. [6] § Materials and methods › Quantifying hemispheric asymmetry: integration and segregation indices ↔ edge_analysis.py, lines 1108–1172 · score 0.51 · intra hemispheric, inter hemispheric, asymmetry, connections
  7. [7] § Results ↔ edge_analysis.py, lines 1108–1172 · score 0.51 · intra hemispheric connections, inter hemispheric connections, asymmetry, LR, RL, LL

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 · 2,310 lines · 92 KB · MIT · 5 matches

  1. #!/usr/bin/env python3
  2. # -*- coding: utf-8 -*-
  3. """
  4. Created on Mon Mar 16 10:27:44 2026
  5. @author: kongxiangyu
  6. """
  7. import sys
  8. import numpy as np
  9. import pandas as pd
  10. import matplotlib.pyplot as plt
  11. import seaborn as sns
  12. import matplotlib.colors as mcolors
  13. from src.utils import load_as_pickle,save_as_pickle
  14. from analysis.edge_wrapper import (run_community_edge_part_pipeline,
  15. run_community_node_part_pipeline,
  16. run_community_coupling_part_pipeline)
  17. from my_code.data_preprocessing import (optimized_assemble_to_flat,
  18. flat_to_assemble_matrix,
  19. unflatten_modes)
  20. from visualization.plot_matrix import plot_half_heatmap
  21. def aggregate_ftype_session_to_dict(sessions,
  22. data_path,
  23. prefix='REST_size60_step1_lag-4_arousal_tvFC_coupling',
  24. prefix1=None,
  25. N=400):
  26. from my_code.data_preprocessing import optimized_assemble_to_flat
  27. all_sessions_data = []
  28. import gc
  29. for ses in sessions:
  30. ll = np.stack(load_as_pickle(f'{data_path}/{prefix}_LL_{ses}{prefix1}.pkl'))
  31. rr = np.stack(load_as_pickle(f'{data_path}/{prefix}_RR_{ses}{prefix1}.pkl'))
  32. lr = np.stack(load_as_pickle(f'{data_path}/{prefix}_LR_{ses}{prefix1}.pkl'))
  33. rl = np.stack(load_as_pickle(f'{data_path}/{prefix}_RL_{ses}{prefix1}.pkl'))
  34. ll = ll.reshape(ll.shape[0],N//2,N//2)
  35. rr = rr.reshape(rr.shape[0],N//2,N//2)
  36. lr = lr.reshape(lr.shape[0],N//2,N//2)
  37. rl = rl.reshape(rl.shape[0],N//2,N//2)
  38. ses_flat = optimized_assemble_to_flat(ll, rr, lr, rl)
  39. del ll, rr, lr, rl
  40. gc.collect()
  41. all_sessions_data.append(ses_flat)
  42. arousal_tvFC_edge_coupling = np.vstack(all_sessions_data)
  43. arousal_tvFC_coupling = all_sessions_data
  44. save_as_pickle(arousal_tvFC_coupling,f'{data_path}/{prefix}{prefix1}.pkl')
  45. return arousal_tvFC_coupling
  46. from joblib import Parallel, delayed
  47. from tqdm import tqdm
  48. from sklearn.cluster import KMeans
  49. from sklearn.metrics import silhouette_score, calinski_harabasz_score, davies_bouldin_score
  50. from sklearn.preprocessing import StandardScaler
  51. def _compute_single_k(X_scaled, k, n_init=10, random_state=0):
  52. """
  53. Internal helper function to run KMeans and metrics for a single K.
  54. """
  55. km = KMeans(n_clusters=k, n_init=n_init, random_state=random_state)
  56. labels = km.fit_predict(X_scaled)
  57. # Calculate metrics
  58. inertia = km.inertia_
  59. # Note: silhouette_score is O(N^2), very slow for 79800 samples.
  60. # We use a sample (e.g., 10000) to speed up if necessary,
  61. # but here we keep it full unless it's too slow.
  62. sil_avg = silhouette_score(X_scaled, labels, sample_size=10000) if X_scaled.shape[0] > 10000 else silhouette_score(X_scaled, labels)
  63. ch_score = calinski_harabasz_score(X_scaled, labels)
  64. db_index = davies_bouldin_score(X_scaled, labels)
  65. return {
  66. 'k': k,
  67. 'inertia': inertia,
  68. 'silhouette': sil_avg,
  69. 'ch': ch_score,
  70. 'db': db_index,
  71. 'labels': labels
  72. }
  73. def run_kmeans_across_k(arousal_tvFC_coupling, K_range=range(2, 15), n_jobs=8):
  74. """
  75. Parallelized version of KMeans evaluation across K.
  76. Parameters
  77. ----------
  78. arousal_tvFC_coupling : array (n_runs, n_edges)
  79. K_range : iterable
  80. n_jobs : int, default=-1 (use all CPUs)
  81. """
  82. # 1. Preprocessing (Consistent with our previous discussion)
  83. X = arousal_tvFC_coupling.T # Shape (79800, 485)
  84. scaler = StandardScaler()
  85. X_scaled = scaler.fit_transform(X)
  86. print(f"Starting parallel processing for K in {list(K_range)}...")
  87. # 2. Parallel Execution using joblib
  88. # n_jobs=-1 uses all available cores
  89. results = Parallel(n_jobs=n_jobs)(
  90. delayed(_compute_single_k)(X_scaled, k) for k in K_range#, desc="Dispatching jobs")
  91. )
  92. # 3. Reassemble results
  93. summary_scores = {
  94. 'inertias': [],
  95. 'silhouette_scores': [],
  96. 'ch_scores': [],
  97. 'db_scores': [],
  98. 'labels': []
  99. }
  100. # Sort by K to ensure order
  101. results.sort(key=lambda x: x['k'])
  102. for res in results:
  103. summary_scores['inertias'].append(res['inertia'])
  104. summary_scores['silhouette_scores'].append(res['silhouette'])
  105. summary_scores['ch_scores'].append(res['ch'])
  106. summary_scores['db_scores'].append(res['db'])
  107. summary_scores['labels'].append(res['labels'])
  108. return summary_scores
  109. from sklearn.metrics import jaccard_score, adjusted_rand_score, normalized_mutual_info_score, confusion_matrix
  110. from scipy.optimize import linear_sum_assignment
  111. def align_and_calculate_stability(base_labels, target_labels, n_clusters):
  112. """
  113. Align target_labels to base_labels using the Hungarian Algorithm and
  114. calculate Jaccard, Dice, ARI, and NMI stability for each community.
  115. Parameters
  116. ----------
  117. base_labels : array, shape (n_samples,)
  118. Reference labels (e.g., from original parameters).
  119. target_labels : array, shape (n_samples,)
  120. Labels to be aligned (e.g., from shifted parameters).
  121. n_clusters : int
  122. Number of clusters (K).
  123. Returns
  124. -------
  125. jaccard_indices : dict
  126. dice_indices : dict
  127. ari_indices : dict
  128. nmi_indices : dict
  129. aligned_labels : array
  130. The target_labels remapped to match the base_labels' indexing.
  131. """
  132. # 1. Compute Confusion Matrix (Intersections)
  133. contingency_matrix = np.zeros((n_clusters, n_clusters))
  134. for i in range(n_clusters):
  135. for j in range(n_clusters):
  136. # i: index for base, j: index for target
  137. intersect = np.sum((base_labels == i) & (target_labels == j))
  138. contingency_matrix[i, j] = intersect
  139. # 2. Hungarian Algorithm (Minimize cost = Maximize intersection)
  140. # We use -contingency_matrix because linear_sum_assignment finds minimum
  141. row_ind, col_ind = linear_sum_assignment(-contingency_matrix)
  142. # 3. Create Mapping (target_label -> base_label)
  143. # col_ind[i] is the target label that best matches base label i
  144. mapping = {target_label: base_label for base_label, target_label in zip(row_ind, col_ind)}
  145. aligned_labels = np.array([mapping[label] for label in target_labels])
  146. # 4. Calculate Metrics for each Community
  147. jaccard_indices = {}
  148. dice_indices = {}
  149. ari_indices = {}
  150. nmi_indices = {}
  151. for k in range(n_clusters):
  152. # Mask for current community
  153. mask_base = (base_labels == k)
  154. mask_aligned = (aligned_labels == k)
  155. # --- Jaccard ---
  156. j_score = jaccard_score(mask_base, mask_aligned)
  157. jaccard_indices[f'community_{k+1}'] = j_score
  158. # --- Dice ---
  159. intersection = np.sum(mask_base & mask_aligned)
  160. size_base = np.sum(mask_base)
  161. size_aligned = np.sum(mask_aligned)
  162. d_score = (2.0 * intersection) / (size_base + size_aligned) if (size_base + size_aligned) > 0 else 0
  163. dice_indices[f'community_{k+1}'] = d_score
  164. # --- ARI ---
  165. a_score = adjusted_rand_score(mask_base, mask_aligned)
  166. ari_indices[f'community_{k+1}'] = a_score
  167. # --- NMI ---
  168. n_score = normalized_mutual_info_score(mask_base, mask_aligned)
  169. nmi_indices[f'community_{k+1}'] = n_score
  170. # Calculate means
  171. jaccard_indices['mean'] = np.mean(list(jaccard_indices.values()))
  172. dice_indices['mean'] = np.mean(list(dice_indices.values()))
  173. ari_indices['mean'] = np.mean(list(ari_indices.values()))
  174. nmi_indices['mean'] = np.mean(list(nmi_indices.values()))
  175. return jaccard_indices, dice_indices, ari_indices, nmi_indices
  176. def get_subject_split_data(data_matrix, session_dict):
  177. """
  178. Groups indices by subject ID to ensure subject-level independent splitting.
  179. Returns:
  180. - subject_to_indices: Dictionary mapping sub_id to list of row indices
  181. - unique_subjects: List of unique subject IDs
  182. """
  183. import random
  184. subject_to_indices = {}
  185. current_idx = 0
  186. session_order = ['rfMRI_REST1_7T_PA', 'rfMRI_REST2_7T_AP', 'rfMRI_REST3_7T_PA', 'rfMRI_REST4_7T_AP']
  187. for session in session_order:
  188. if session in session_dict:
  189. for run_name in session_dict[session]:
  190. sub_id = run_name.split('_')[0]
  191. if sub_id not in subject_to_indices:
  192. subject_to_indices[sub_id] = []
  193. subject_to_indices[sub_id].append(current_idx)
  194. current_idx += 1
  195. unique_subjects = list(subject_to_indices.keys())
  196. random.shuffle(unique_subjects)
  197. half_point = len(unique_subjects) // 2
  198. # Map subject lists back to matrix row indices
  199. idx_half1 = [idx for sub in unique_subjects[:half_point] for idx in subject_to_indices[sub]]
  200. idx_half2 = [idx for sub in unique_subjects[half_point:] for idx in subject_to_indices[sub]]
  201. # Slicing the data matrix
  202. data_half1 = data_matrix[idx_half1, :]
  203. data_half2 = data_matrix[idx_half2, :]
  204. return data_half1, data_half2
  205. def get_subject_split_indices(session_dict):
  206. """
  207. Step 1: Parse the session dictionary and group row indices by Subject ID.
  208. """
  209. subject_to_indices = {}
  210. current_idx = 0
  211. session_order = ['rfMRI_REST1_7T_PA', 'rfMRI_REST2_7T_AP', 'rfMRI_REST3_7T_PA', 'rfMRI_REST4_7T_AP']
  212. for session in session_order:
  213. if session in session_dict:
  214. for run_name in session_dict[session]:
  215. sub_id = run_name.split('_')[0]
  216. if sub_id not in subject_to_indices:
  217. subject_to_indices[sub_id] = []
  218. subject_to_indices[sub_id].append(current_idx)
  219. current_idx += 1
  220. return subject_to_indices, list(subject_to_indices.keys())
  221. def align_labels_to_reference(ref_labels, target_labels, k):
  222. """
  223. New Utility: Aligns target_labels to match ref_labels numerically.
  224. Uses Hungarian Algorithm to maximize the overlap between clusters.
  225. Parameters:
  226. - ref_labels: Reference cluster labels (1D array)
  227. - target_labels: Labels to be reordered (1D array)
  228. - k: Number of clusters
  229. Returns:
  230. - aligned_labels: target_labels mapped to ref_labels' indexing
  231. """
  232. # 1. Build contingency matrix (Jaccard similarity between all pairs)
  233. contingency = np.zeros((k, k))
  234. for i in range(k):
  235. mask_ref = (ref_labels == i)
  236. for j in range(k):
  237. mask_target = (target_labels == j)
  238. intersection = np.logical_and(mask_ref, mask_target).sum()
  239. union = np.logical_or(mask_ref, mask_target).sum()
  240. contingency[i, j] = intersection / union if union > 0 else 0
  241. # 2. Find optimal matching (maximize Jaccard overlap)
  242. # row_ind corresponds to ref_labels index, col_ind to target_labels index
  243. row_ind, col_ind = linear_sum_assignment(-contingency)
  244. # 3. Create mapping and reassign labels
  245. # mapping[original_target_label] = matched_ref_label
  246. mapping = {target_l: ref_l for ref_l, target_l in zip(row_ind, col_ind)}
  247. aligned_labels = np.array([mapping[l] for l in target_labels])
  248. return aligned_labels
  249. def calculate_robust_stability(labels_ref, labels_target, k, return_matrix=False):
  250. """
  251. Step 3: Evaluate stability and alignment.
  252. """
  253. # 1. Hungarian Alignment for Jaccard
  254. jaccard_matrix = np.zeros((k, k))
  255. for i in range(k):
  256. mask_ref = (labels_ref == i)
  257. for j in range(k):
  258. mask_target = (labels_target == j)
  259. intersection = np.logical_and(mask_ref, mask_target).sum()
  260. union = np.logical_or(mask_ref, mask_target).sum()
  261. jaccard_matrix[i, j] = intersection / union if union > 0 else 0
  262. row_ind, col_ind = linear_sum_assignment(-jaccard_matrix)
  263. matched_jaccard = jaccard_matrix[row_ind, col_ind].mean()
  264. # 2. Global Metrics
  265. ari = adjusted_rand_score(labels_ref, labels_target)
  266. nmi = normalized_mutual_info_score(labels_ref, labels_target)
  267. if return_matrix:
  268. # Create a re-mapped version of target labels for confusion matrix visualization
  269. mapping = {old: new for new, old in zip(row_ind, col_ind)}
  270. aligned_target = np.array([mapping[l] for l in labels_target])
  271. conf_mat = confusion_matrix(labels_ref, aligned_target, normalize='true')
  272. return matched_jaccard, ari, nmi, conf_mat
  273. return matched_jaccard, ari, nmi
  274. def split_data_by_subject(data_matrix, subject_to_indices, unique_subjects):
  275. """
  276. Step 2: Physically split the data matrix into two halves based on subject identity.
  277. """
  278. import random
  279. random.shuffle(unique_subjects)
  280. half_point = len(unique_subjects) // 2
  281. idx_half1 = [idx for sub in unique_subjects[:half_point] for idx in subject_to_indices[sub]]
  282. idx_half2 = [idx for sub in unique_subjects[half_point:] for idx in subject_to_indices[sub]]
  283. return data_matrix[idx_half1, :], data_matrix[idx_half2, :],idx_half1,idx_half2
  284. def run_comprehensive_similarity_analysis(summary_scores_dict, target_k, k_range, n_jobs=20):
  285. """
  286. Optimized calculation of all-to-all similarity matrices using parallel processing.
  287. Args:
  288. summary_scores_dict: Dictionary containing 'labels' for different parameters.
  289. target_k: The specific number of clusters to analyze.
  290. k_range: The range of k values used in the original scoring.
  291. n_jobs: Number of CPU cores to use. -1 uses all available processors.
  292. """
  293. from itertools import combinations_with_replacement
  294. # Locating the index for the specific target_k in the k_range list
  295. k_index = list(k_range).index(target_k)
  296. keys = list(summary_scores_dict.keys())
  297. n_params = len(keys)
  298. # Pre-extract labels to avoid passing the massive dictionary to worker processes
  299. # This reduces IPC (Inter-Process Communication) overhead significantly
  300. all_labels = [summary_scores_dict[k]['labels'][k_index] for k in keys]
  301. # Define tracking keys for communities and the overall mean
  302. com_keys = [f'community_{i+1}' for i in range(target_k)] + ['mean']
  303. metrics_names = ['jaccard', 'dice', 'ari', 'nmi']
  304. # Generate unique pairs (i, j) for the upper triangle including the diagonal
  305. # This reduces calculations from N^2 to (N*(N+1))/2
  306. pairs = list(combinations_with_replacement(range(n_params), 2))
  307. print(f"Calculating parallel all-to-all similarity for K={target_k} using {n_jobs} cores...")
  308. # Helper function to process a single pair of label sets
  309. def compute_pair(i, j):
  310. # res: (jaccard_dict, dice_dict, ari_dict, nmi_dict, aligned_labels)
  311. res = align_and_calculate_stability(all_labels[i], all_labels[j], n_clusters=target_k)
  312. # Return indices and the first 4 dictionary results (metrics)
  313. return i, j, res[:4]
  314. # Execute parallel processing
  315. results_list = Parallel(n_jobs=n_jobs)(
  316. delayed(compute_pair)(i, j) for i, j in tqdm(pairs, desc='Computing Pairs')
  317. )
  318. # Initialize empty matrices for each metric and community
  319. all_matrices = {
  320. m: {ck: np.zeros((n_params, n_params)) for ck in com_keys}
  321. for m in metrics_names
  322. }
  323. # Populate matrices using the results from parallel workers
  324. for i, j, metrics_dicts in results_list:
  325. for m_idx, m_name in enumerate(metrics_names):
  326. current_metric_dict = metrics_dicts[m_idx]
  327. for ck in com_keys:
  328. val = current_metric_dict[ck]
  329. # Fill the upper triangle
  330. all_matrices[m_name][ck][i, j] = val
  331. # Mirror to the lower triangle if not on the diagonal
  332. if i != j:
  333. all_matrices[m_name][ck][j, i] = val
  334. # Pack the result matrices into a nested dictionary of Pandas DataFrames
  335. results = {m: {} for m in metrics_names}
  336. for m in metrics_names:
  337. for ck in com_keys:
  338. results[m][ck] = pd.DataFrame(
  339. all_matrices[m][ck],
  340. index=keys,
  341. columns=keys
  342. )
  343. return results
  344. def analyze_similarity_by_parameter(similarity_df, mode='lag'):
  345. def extract_param(name):
  346. parts = name.split('_')
  347. for p in parts:
  348. if mode in p:
  349. return p
  350. return 'unknown'
  351. temp_df = similarity_df.copy()
  352. temp_df['group'] = [extract_param(idx) for idx in temp_df.index]
  353. grouped_rows = temp_df.groupby('group').mean()
  354. grouped_rows = grouped_rows.T
  355. grouped_rows['group'] = [extract_param(idx) for idx in grouped_rows.index]
  356. param_sim_matrix = grouped_rows.groupby('group').mean()
  357. def sort_key(s):
  358. import re
  359. match = re.search(r"(-?\d+)", s)
  360. return int(match.group(1)) if match else 0
  361. sorted_index = sorted(param_sim_matrix.index, key=sort_key)
  362. param_sim_matrix = param_sim_matrix.reindex(index=sorted_index, columns=sorted_index)
  363. return param_sim_matrix
  364. def summary_noisy_data(info, working_path, version):
  365. sessions=list(info['IDRuns_group'].keys())
  366. IDRuns = info['IDRuns_group']
  367. ET_group = info['ET_group_1hz']
  368. summary_noisy = {'motion':{session:np.zeros((24,900,IDRuns[session].shape[0]))
  369. for session in sessions},
  370. 'blinks':{session:np.zeros((900,IDRuns[session].shape[0]))
  371. for session in sessions},
  372. 'GS':{session:np.zeros((900,IDRuns[session].shape[0]))
  373. for session in sessions},
  374. }
  375. image_path = '/data/disk0/kongxiangyu/image/32k/'
  376. for session in sessions:
  377. for i, IDRun in enumerate(IDRuns[session]):
  378. ID = IDRun[0:6]
  379. # movement
  380. movement = np.loadtxt(f'{image_path}/{ID}/{session}/Movement_Regressors.txt')
  381. movement_sq = movement**2
  382. movement_friston24 = np.hstack([movement, movement_sq])
  383. summary_noisy['motion'][session][:,:,i] = movement_friston24.T
  384. # blinks
  385. summary_noisy['blinks'][session] = np.isnan(ET_group[session])
  386. # GS
  387. cortex_node_ts_group = load_as_pickle(os.path.join(
  388. working_path, f'result/{version}/cortex_node_ts_group_{session}.joblib'))
  389. for i, IDRun in enumerate(IDRuns[session]):
  390. ID = IDRun[0:6]
  391. summary_noisy['GS'][session][:,i] = np.mean(cortex_node_ts_group[IDRun],axis=1)
  392. return summary_noisy
  393. from sklearn.linear_model import LinearRegression
  394. from sklearn.metrics import mean_squared_error
  395. import os
  396. def plot_L_method(metrics, K_range, metrics_name='SSE', output_path='.'):
  397. """
  398. Automatically finds optimal K and plots:
  399. 1. L-method regression fits.
  400. 2. Total RMSE for each candidate K to prove why the best K was chosen.
  401. """
  402. k_values = np.array(list(K_range))
  403. metrics = np.array(metrics)
  404. k_values_reshaped = k_values.reshape(-1, 1)
  405. # --- Step 1: Search for Optimal K by minimizing Total RMSE ---
  406. # We need at least 2 points for each line (left and right)
  407. # search_space stores the indices of candidate 'elbow' points
  408. search_indices = range(1, len(k_values) - 1)
  409. rmse_results = []
  410. candidate_ks = []
  411. for i in search_indices:
  412. # Define current candidate K
  413. curr_k = k_values[i]
  414. candidate_ks.append(curr_k)
  415. # Split data at index i (including curr_k in both for the elbow)
  416. x_left, y_left = k_values_reshaped[:i+1], metrics[:i+1]
  417. x_right, y_right = k_values_reshaped[i:], metrics[i:]
  418. # Fit linear models
  419. reg_left = LinearRegression().fit(x_left, y_left)
  420. reg_right = LinearRegression().fit(x_right, y_right)
  421. # Calculate RMSE for both sides
  422. mse_left = mean_squared_error(y_left, reg_left.predict(x_left))
  423. mse_right = mean_squared_error(y_right, reg_right.predict(x_right))
  424. # Calculate Total Weighted RMSE
  425. n_total = len(metrics)
  426. total_rmse = (len(y_left) * np.sqrt(mse_left) + len(y_right) * np.sqrt(mse_right)) / n_total
  427. rmse_results.append(total_rmse)
  428. # Identify the best K
  429. best_idx = np.argmin(rmse_results)
  430. optimal_k = candidate_ks[best_idx]
  431. # --- Step 2: Visualization ---
  432. png_path = os.path.join(output_path, 'png')
  433. os.makedirs(png_path, exist_ok=True)
  434. # FIGURE 1: Regression Fitting
  435. plt.figure(figsize=(10, 6))
  436. split_i = search_indices[best_idx]
  437. # Re-run best fit for plotting
  438. rl = LinearRegression().fit(k_values_reshaped[:split_i+1], metrics[:split_i+1])
  439. rr = LinearRegression().fit(k_values_reshaped[split_i:], metrics[split_i:])
  440. plt.plot(k_values, metrics, 'ko-', label=f'Original {metrics_name}', alpha=0.4)
  441. plt.plot(k_values[:split_i+1], rl.predict(k_values_reshaped[:split_i+1]), 'r--', linewidth=2, label='Left Fit')
  442. plt.plot(k_values[split_i:], rr.predict(k_values_reshaped[split_i:]), 'b--', linewidth=2, label='Right Fit')
  443. plt.axvline(x=optimal_k, color='green', linestyle=':', label=f'Optimal K={optimal_k}')
  444. plt.title(f'L-Method: Optimal K Search ({metrics_name})', fontsize=15)
  445. plt.xlabel('Number of Clusters (K)', fontsize=12)
  446. plt.ylabel('Metric Value', fontsize=12)
  447. plt.xticks(K_range)
  448. plt.legend()
  449. plt.grid(True, alpha=0.2)
  450. plt.savefig(os.path.join(png_path, f'L_method_fit_{metrics_name}.png'), dpi=300)
  451. plt.show()
  452. plt.figure(figsize=(12, 6))
  453. k_labels = [str(x) for x in candidate_ks]
  454. bars = plt.bar(k_labels, rmse_results, color='lightgray', edgecolor='black', alpha=0.6)
  455. # Highlight the minimum RMSE bar
  456. bars[best_idx].set_color('mediumseagreen')
  457. bars[best_idx].set_edgecolor('darkgreen')
  458. bars[best_idx].set_alpha(0.9)
  459. # Add labels on top of bars
  460. for i, val in enumerate(rmse_results):
  461. plt.text(i, val, f'{val:.4f}', ha='center', va='bottom', fontsize=9)
  462. plt.title(f'L-Method Evaluation: Total RMSE per Candidate K', fontsize=15)
  463. plt.xlabel('Candidate Elbow Point (K)', fontsize=12)
  464. plt.ylabel('Total RMSE (Lower is Better)', fontsize=12)
  465. plt.axhline(min(rmse_results), color='red', linestyle='--', alpha=0.3)
  466. plt.savefig(os.path.join(png_path, f'L_method_rmse_eval_{metrics_name}.png'), dpi=300)
  467. plt.show()
  468. return optimal_k
  469. def parse_param_string(param_name):
  470. """
  471. Parse parameter string.
  472. Naming convention: 'size30_step5_lag0' or 'size30_step5_lag0_con1'
  473. """
  474. if 'con' in param_name:
  475. parts = param_name.split('_')
  476. # Format: sizeXX_stepXX_lagXX_conX
  477. size_param, step_param, lag_param, confound = parts[0], parts[1], parts[2], parts[3]
  478. # Mapping according to definitions: con1=Full, con2=Medium, con3=Basic
  479. if confound.endswith('1'):
  480. con_param = 'Motion+Blinks+GS'
  481. elif confound.endswith('2'):
  482. con_param = 'Motion+GS'
  483. elif confound.endswith('3'):
  484. con_param = 'Motion'
  485. else:
  486. con_param = confound
  487. else:
  488. parts = param_name.split('_')
  489. step_param, lag_param = parts[1], parts[2]
  490. con_param = 'Raw'
  491. return step_param, lag_param, con_param
  492. def prepare_plotting_data(sim_results, target_metric='jaccard', target_com='mean'):
  493. """
  494. Convert results to long-form DataFrame, keeping only Confound vs. Raw similarity.
  495. """
  496. df_matrix = sim_results[target_metric][target_com]
  497. plot_rows = []
  498. all_keys = df_matrix.index.tolist()
  499. for param_name in all_keys:
  500. step, lag, con = parse_param_string(param_name)
  501. # Only calculate score for Confound versions relative to Raw, exclude Raw self-comparison
  502. if con != "Raw":
  503. base_key = f"size30_{step}_{lag}"
  504. if base_key in df_matrix.columns:
  505. score = df_matrix.loc[param_name, base_key]
  506. plot_rows.append({
  507. "Step": step,
  508. "Lag": lag,
  509. "Confound_Version": con,
  510. "Similarity": score
  511. })
  512. return pd.DataFrame(plot_rows)
  513. def plot_regression_comparison_line(plot_df, output_path, target_step='step5',prefix=''):
  514. """
  515. Plot line chart: X-axis as Lag, colors representing different regression versions.
  516. Excludes Raw self-line to focus on robustness contrast.
  517. """
  518. import re
  519. # Filter by specific Step
  520. sub_df = plot_df[plot_df['Step'] == target_step].copy()
  521. # Ensure Lag is sorted numerically (handles strings like 'lag0', 'lag1')
  522. if sub_df['Lag'].dtype == object:
  523. sub_df['lag_val'] = sub_df['Lag'].apply(lambda x: int(re.search(r'(-?\d+)', x).group(1)))
  524. sub_df = sub_df.sort_values('lag_val')
  525. plt.figure(figsize=(10, 7))
  526. sns.set_theme(style="ticks", context="talk")
  527. # Plotting
  528. ax = sns.lineplot(
  529. data=sub_df,
  530. x="Lag",
  531. y="Similarity",
  532. hue="Confound_Version",
  533. style="Confound_Version",
  534. markers=True,
  535. dashes=False,
  536. markersize=10,
  537. linewidth=2.5,
  538. palette="flare"
  539. )
  540. plt.title(f"Stability: Confound Regressed vs. Raw ({target_step})", pad=60)
  541. plt.xlabel("Temporal Lag (TR)")
  542. plt.ylabel("Similarity to Raw Version")
  543. # plt.ylim(0.2, 1)
  544. plt.grid(True, axis='y', alpha=0.3)
  545. n_versions = sub_df['Confound_Version'].nunique()
  546. plt.legend(
  547. loc='lower center',
  548. bbox_to_anchor=(0.5, 1.01),
  549. ncol=n_versions,
  550. frameon=False,
  551. fontsize='small',
  552. handletextpad=0.5,
  553. columnspacing=1.5
  554. )
  555. plt.tight_layout()
  556. if output_path:
  557. plt.savefig(f'{output_path}/png/Confound_Regressed_Stability{prefix}_{target_step}.png', dpi=600, bbox_inches='tight')
  558. plt.show()
  559. def plot_comparison_matrices(matrix1, matrix2, title1="Original", title2="Aligned", cmap='viridis'):
  560. """
  561. Helper function to plot two matrices side-by-side for comparison.
  562. """
  563. fig, axes = plt.subplots(1, 2, figsize=(15, 6))
  564. im1 = axes[0].imshow(matrix1, cmap=cmap)
  565. axes[0].set_title(title1)
  566. axes[0].set_xticks([])
  567. axes[0].set_yticks([])
  568. fig.colorbar(im1, ax=axes[0])
  569. im2 = axes[1].imshow(matrix2, cmap=cmap)
  570. axes[1].set_title(title2)
  571. axes[1].set_xticks([])
  572. axes[1].set_yticks([])
  573. fig.colorbar(im2, ax=axes[1])
  574. plt.tight_layout()
  575. plt.show()
  576. def add_mean_in_result(res_level, selected_communities=['community_2', 'community_3', 'community_4']):
  577. updated_results = {}
  578. for metric_name, communities_dict in res_level.items():
  579. updated_results[metric_name] = communities_dict.copy()
  580. target_dfs = []
  581. for ck in selected_communities:
  582. if ck in communities_dict:
  583. target_dfs.append(communities_dict[ck])
  584. else:
  585. print(f"Warning: {ck} not found in metric {metric_name}")
  586. if target_dfs:
  587. combined_arrays = np.stack([df.values for df in target_dfs], axis=0)
  588. mean_array = np.mean(combined_arrays, axis=0)
  589. mean_df = pd.DataFrame(
  590. mean_array,
  591. index=target_dfs[0].index,
  592. columns=target_dfs[0].columns
  593. )
  594. updated_results[metric_name]['selected-mean'] = mean_df
  595. return updated_results
  596. def process_single_k(k, data1, data2, iter_idx):
  597. """
  598. Processes a single K value: clustering, explicit label alignment,
  599. reliability metrics calculation, and returning labels.
  600. """
  601. samples1 = data1.T
  602. samples2 = data2.T
  603. # 1. Independent Clustering
  604. # Using iter_idx to ensure different seeds for each iteration
  605. km1 = KMeans(n_clusters=k, n_init=10, random_state=iter_idx * 100 + 42).fit(samples1)
  606. labels1 = km1.labels_ # Reference labels
  607. km2 = KMeans(n_clusters=k, n_init=10, random_state=iter_idx * 100 + 43).fit(samples2)
  608. labels2 = km2.labels_ # Labels to be aligned
  609. # 2. Label Alignment
  610. # Based on the overlap of samples between the two clusterings
  611. contingency_matrix = np.zeros((k, k))
  612. for i in range(k):
  613. for j in range(k):
  614. # Calculate the number of overlapping samples between cluster i (labels1) and cluster j (labels2)
  615. contingency_matrix[i, j] = np.sum((labels1 == i) & (labels2 == j))
  616. # Use Hungarian algorithm to find optimal matching (minimize negative overlap = maximize overlap)
  617. row_idx, col_idx = linear_sum_assignment(-contingency_matrix)
  618. # Map labels2 categories to labels1's coordinate system
  619. mapping = {old_label: new_label for new_label, old_label in zip(row_idx, col_idx)}
  620. mapped_labels2 = np.array([mapping[l] for l in labels2])
  621. # 3. Calculate Similarity Metrics
  622. ari = adjusted_rand_score(labels1, mapped_labels2)
  623. nmi = normalized_mutual_info_score(labels1, mapped_labels2)
  624. # Calculate Dice coefficient (overlap after alignment)
  625. def calculate_dice(l1, l2, num_k):
  626. dices = []
  627. for i in range(num_k):
  628. mask1 = (l1 == i)
  629. mask2 = (l2 == i)
  630. intersection = np.sum(mask1 & mask2)
  631. dice = (2. * intersection) / (np.sum(mask1) + np.sum(mask2) + 1e-10)
  632. dices.append(dice)
  633. return np.mean(dices), dices
  634. avg_dice, per_cluster_dice = calculate_dice(labels1, mapped_labels2, k)
  635. # 4. Normalized Confusion Matrix
  636. conf_mat = confusion_matrix(
  637. labels1,
  638. mapped_labels2,
  639. labels=range(k),
  640. normalize='true'
  641. )
  642. return {
  643. 'k': k,
  644. 'ari': ari,
  645. 'nmi': nmi,
  646. 'dice': avg_dice,
  647. 'per_cluster_dice': per_cluster_dice,
  648. 'conf_mat': conf_mat,
  649. 'labels_ref': labels1,
  650. 'labels_aligned': mapped_labels2
  651. }
  652. def process_iteration(iter_idx, data_matrix, subject_to_indices, unique_subjects, k_range):
  653. """
  654. Task for a single iteration: split data and compute all K values.
  655. This is the primary unit for parallelization.
  656. """
  657. # 1. Split data by subject (performed inside each process for unique splits)
  658. data1, data2 = split_data_by_subject(data_matrix, subject_to_indices, unique_subjects)
  659. # 2. Sequential loop over k_range within this iteration
  660. iter_results = []
  661. for k in k_range:
  662. res = process_single_k(k, data1, data2, iter_idx)
  663. iter_results.append(res)
  664. return iter_results
  665. def run_split_half_analysis(data_matrix, session_dict, k_range=range(2, 15), iterations=20, n_jobs=15):
  666. """
  667. Main orchestration function: performs split-half reliability analysis in parallel.
  668. Parallelization is executed over the 'iterations' level.
  669. """
  670. # Get subject-level split indices
  671. subject_to_indices, unique_subjects = get_subject_split_indices(session_dict)
  672. print(f"Starting parallel analysis: Total iterations={iterations}, Jobs={n_jobs}...")
  673. # CORE PARALLEL STEP: Dispatch iterations to multiple workers
  674. all_iterations_results = Parallel(n_jobs=n_jobs)(
  675. delayed(process_iteration)(i, data_matrix, subject_to_indices, unique_subjects, k_range)
  676. for i in tqdm(range(iterations),desc='running iterations ...')
  677. )
  678. # Initialize result storage
  679. stats = {
  680. k: {m: [] for m in ['dice', 'ari', 'nmi', 'per_cluster_dice', 'idx_half1', 'idx_half2']}
  681. for k in k_range
  682. }
  683. conf_mat = {}
  684. half_labels = {}
  685. # Aggregating results from parallel workers
  686. for i, iter_res in enumerate(all_iterations_results):
  687. # We designate the result from the "last" index in the returned list as the sample for visualization
  688. is_last_item = (i == len(all_iterations_results) - 1)
  689. for res in iter_res:
  690. k = res['k']
  691. # Basic statistics
  692. for metric in ['dice', 'ari', 'nmi', 'per_cluster_dice', 'idx_half1', 'idx_half2']:
  693. stats[k][metric].append(res[metric])
  694. # Capture the data from the last iteration for downstream visualization/debug
  695. if is_last_item:
  696. conf_mat[k] = res['conf_mat']
  697. half_labels[k] = {
  698. 'labels1': res['labels_ref'],
  699. 'labels2': res['labels_aligned']
  700. }
  701. print(f"Analysis finished. Captured labels for K range: {list(half_labels.keys())}")
  702. return stats, conf_mat, half_labels
  703. def align_to_reference(labels, ref_labels, k):
  704. """
  705. Align current clustering labels to a fixed reference using the Hungarian algorithm.
  706. """
  707. offset = 1
  708. # Create a contingency matrix (Reference x Current)
  709. contingency_matrix = np.zeros((k, k))
  710. for i in range(k):
  711. for j in range(k):
  712. # i = reference cluster index, j = current cluster index
  713. contingency_matrix[i, j] = np.sum((ref_labels == (i+offset)) & (labels == (j+offset)))
  714. # Use Hungarian algorithm to maximize overlap (minimize negative overlap)
  715. row_idx, col_idx = linear_sum_assignment(-contingency_matrix)
  716. # Create mapping: current_label -> reference_label
  717. # col_idx contains the original labels, row_idx contains their new aligned identities
  718. mapping = {old_idx + offset: new_idx + offset for new_idx, old_idx in zip(row_idx, col_idx)}
  719. # Transform labels
  720. aligned_labels = np.array([mapping.get(l, l) for l in labels])
  721. return aligned_labels
  722. def process_single_k(k, data1, data2, ref_labels, iter_idx):
  723. """
  724. Processes a single K: Align both halves to a common reference.
  725. """
  726. # 1. Independent Clustering
  727. km1 = KMeans(n_clusters=k, n_init=10, random_state=iter_idx * 100 + 42).fit(data1.T)
  728. km2 = KMeans(n_clusters=k, n_init=10, random_state=iter_idx * 100 + 43).fit(data2.T)
  729. # 2. Dual Alignment to Reference
  730. # Both sets of labels are now in the "Reference Label Space"
  731. mapped_labels1 = align_to_reference(km1.labels_, ref_labels, k)
  732. mapped_labels2 = align_to_reference(km2.labels_, ref_labels, k)
  733. # 3. Calculate Similarity Metrics between the two aligned halves
  734. ari = adjusted_rand_score(mapped_labels1, mapped_labels2)
  735. nmi = normalized_mutual_info_score(mapped_labels1, mapped_labels2)
  736. # Calculate Dice (Consistency between aligned halves)
  737. def calculate_dice(l1, l2, num_k):
  738. dices = []
  739. for i in range(num_k):
  740. mask1 = (l1 == i)
  741. mask2 = (l2 == i)
  742. intersection = np.sum(mask1 & mask2)
  743. denom = np.sum(mask1) + np.sum(mask2)
  744. dice = (2. * intersection) / (denom + 1e-10)
  745. dices.append(dice)
  746. return np.mean(dices), dices
  747. avg_dice, per_cluster_dice = calculate_dice(mapped_labels1, mapped_labels2, k)
  748. # 4. Confusion Matrix (How well Half 2 matches Half 1 after global alignment)
  749. conf_mat = confusion_matrix(
  750. mapped_labels1,
  751. mapped_labels2,
  752. labels=range(k),
  753. normalize='true'
  754. )
  755. return {
  756. 'k': k,
  757. 'ari': ari,
  758. 'nmi': nmi,
  759. 'dice': avg_dice,
  760. 'per_cluster_dice': per_cluster_dice,
  761. 'conf_mat': conf_mat,
  762. 'labels_h1': mapped_labels1,
  763. 'labels_h2': mapped_labels2
  764. }
  765. def process_iteration(iter_idx, data_matrix, subject_to_indices, unique_subjects, best_k, reference_labels):
  766. """
  767. Task for a single iteration: Split data and align both halves to a global reference.
  768. """
  769. # 1. Split data by subject (Half-split)
  770. # Ensure this function is defined in your environment
  771. data1, data2, idx_half1, idx_half2 = split_data_by_subject(data_matrix,
  772. subject_to_indices,
  773. unique_subjects)
  774. # 2. Independent Clustering for each half
  775. # Use distinct seeds for each half and each iteration to ensure variability
  776. km1 = KMeans(n_clusters=best_k, n_init=10, random_state=iter_idx * 100 + 1).fit(data1.T)
  777. km2 = KMeans(n_clusters=best_k, n_init=10, random_state=iter_idx * 100 + 2).fit(data2.T)
  778. label1 = km1.labels_+1
  779. label2 = km2.labels_+1
  780. # 3. Align BOTH halves to the GLOBAL REFERENCE
  781. # Now both label sets share the same "physical meaning" defined by reference_labels
  782. aligned_labels1 = align_to_reference(label1, reference_labels, best_k)
  783. aligned_labels2 = align_to_reference(label2, reference_labels, best_k)
  784. # aligned_labels1=align_labels_to_reference(label1, labels_edge)
  785. # 4. Calculate Reliability Metrics between the two aligned halves
  786. ari = adjusted_rand_score(aligned_labels1, aligned_labels2)
  787. nmi = normalized_mutual_info_score(aligned_labels1, aligned_labels2)
  788. # Function to calculate Dice coefficient for each cluster
  789. def calculate_dice(l1, l2, k):
  790. dices = []
  791. for i in range(k):
  792. mask1 = (l1 == (i+1))
  793. mask2 = (l2 == (i+1))
  794. intersection = np.sum(mask1 & mask2)
  795. sum_val = np.sum(mask1) + np.sum(mask2)
  796. dice = (2. * intersection) / (sum_val + 1e-10)
  797. dices.append(dice)
  798. return np.mean(dices), dices
  799. avg_dice, per_cluster_dice = calculate_dice(aligned_labels1, aligned_labels2, best_k)
  800. # ci_matrix = flat_to_assemble_matrix(aligned_labels1)
  801. # 5. Normalized Confusion Matrix (Half 1 as row, Half 2 as column)
  802. conf_mat_h1h2 = confusion_matrix(
  803. aligned_labels1,
  804. aligned_labels2,
  805. labels=range(1,best_k+1),
  806. normalize='true'
  807. )
  808. conf_mat_origh1 = confusion_matrix(
  809. reference_labels,
  810. aligned_labels1,
  811. labels=range(1,best_k+1),
  812. normalize='true'
  813. )
  814. conf_mat_origh2 = confusion_matrix(
  815. reference_labels,
  816. aligned_labels2,
  817. labels=range(1,best_k+1),
  818. normalize='true'
  819. )
  820. return {
  821. 'ari': ari,
  822. 'nmi': nmi,
  823. 'dice': avg_dice,
  824. 'per_cluster_dice': per_cluster_dice,
  825. 'conf_mat_h1h2': conf_mat_h1h2,
  826. 'conf_mat_origh1': conf_mat_origh1,
  827. 'conf_mat_origh2': conf_mat_origh2,
  828. 'labels_h1': aligned_labels1,
  829. 'labels_h2': aligned_labels2,
  830. 'idx_half1': idx_half1,
  831. 'idx_half2': idx_half2
  832. }
  833. def run_split_half_analysis(data_matrix, session_dict, reference_labels, best_k, iterations=100, n_jobs=15):
  834. """
  835. Main orchestration function for fixed-K split-half reliability analysis.
  836. Args:
  837. data_matrix: Voxel/Vertex by Time/Subject matrix.
  838. session_dict: Metadata for splitting.
  839. reference_labels: The "Gold Standard" labels (e.g., from full-sample clustering).
  840. best_k: The fixed number of clusters to analyze.
  841. iterations: Number of random splits.
  842. n_jobs: Parallel workers.
  843. """
  844. # Get subject-level split indices
  845. subject_to_indices, unique_subjects = get_subject_split_indices(session_dict)
  846. print("Starting Fixed-K Reliability Analysis...")
  847. print(f"K = {best_k}, Iterations = {iterations}, Reference-based Alignment.")
  848. # Parallel execution over iterations
  849. results = Parallel(n_jobs=n_jobs)(
  850. delayed(process_iteration)(
  851. i, data_matrix, subject_to_indices, unique_subjects, best_k, reference_labels
  852. )
  853. for i in tqdm(range(iterations), desc='Processing iterations')
  854. )
  855. # Initialize storage for aggregated statistics
  856. stats = {
  857. 'dice': [],
  858. 'ari': [],
  859. 'nmi': [],
  860. 'per_cluster_dice': [],
  861. 'idx_half1':[],
  862. 'idx_half2':[],
  863. 'conf_mat_h1h2':[],
  864. 'conf_mat_origh1':[],
  865. 'conf_mat_origh2':[],
  866. }
  867. # Extract results
  868. all_conf_mats = []
  869. for res in results:
  870. stats['dice'].append(res['dice'])
  871. stats['ari'].append(res['ari'])
  872. stats['nmi'].append(res['nmi'])
  873. stats['per_cluster_dice'].append(res['per_cluster_dice'])
  874. stats['idx_half1'].append(res['idx_half1'])
  875. stats['idx_half2'].append(res['idx_half2'])
  876. stats['conf_mat_h1h2'].append(res['conf_mat_h1h2'])
  877. stats['conf_mat_origh1'].append(res['conf_mat_origh1'])
  878. stats['conf_mat_origh2'].append(res['conf_mat_origh2'])
  879. # Average confusion matrix across all iterations for visualization
  880. avg_conf_mat = np.mean(all_conf_mats, axis=0)
  881. # Store the last iteration labels for potential debugging
  882. half_labels = {
  883. 'labels1': [r['labels_h1'] for r in results],
  884. 'labels2': [r['labels_h2'] for r in results]
  885. }
  886. print(f"Analysis complete for K={best_k}.")
  887. return stats, half_labels
  888. def plot_split_half_reliability(stats, best_k, cmap, output_path,prefix):
  889. """
  890. Visualize stability analysis results for a specific K.
  891. """
  892. # 1. Data preparation for Raincloud plot
  893. dice_data = np.array(stats['per_cluster_dice'])
  894. cluster_names = [f'C{i+1}' for i in range(best_k)]
  895. df_dice = pd.DataFrame(dice_data, columns=cluster_names)
  896. df_melted = df_dice.melt(var_name='Cluster', value_name='Dice Score')
  897. colors = [cmap(i) for i in range(best_k)]
  898. # 2. Initialize figure with 2x2 subplots
  899. sns.set_theme(style="whitegrid")
  900. fig, axes = plt.subplots(2, 2, figsize=(15, 12))
  901. axes = axes.flatten()
  902. # --- Panel A: Heatmap 1 ---
  903. cm_h1h2 = np.mean(np.array(stats['conf_mat_h1h2']), axis=0)
  904. cm_h1h2_norm = cm_h1h2.astype('float') / (cm_h1h2.sum(axis=1)[:, np.newaxis] + 1e-10)
  905. sns.heatmap(cm_h1h2_norm, annot=True, fmt=".2f", cmap="Blues", ax=axes[0],
  906. xticklabels=cluster_names, yticklabels=cluster_names,
  907. square=True, cbar_kws={'label': 'Proportion of Samples'})
  908. axes[0].set_title(f"Panel A: Alignment Accuracy (K={best_k})\n(Half 1 vs. Aligned Half 2)",
  909. fontsize=14, fontweight='bold', pad=15)
  910. axes[0].set_xlabel("Aligned Half 2 Labels", fontsize=12)
  911. axes[0].set_ylabel("Half 1 Reference Labels", fontsize=12)
  912. # --- Panel B: Raincloud Plot 1 ---
  913. ax_b = axes[1]
  914. sns.violinplot(x='Cluster', y='Dice Score', data=df_melted, palette=colors,
  915. alpha=0.3, inner=None, density_norm='width', ax=ax_b, hue='Cluster', legend=False)
  916. for violin in [c for c in ax_b.collections if isinstance(c, plt.matplotlib.collections.PolyCollection)]:
  917. for path in violin.get_paths():
  918. m = path.vertices[:, 0].mean()
  919. path.vertices[:, 0] = np.clip(path.vertices[:, 0], m, np.inf)
  920. for i, cluster in enumerate(cluster_names):
  921. cluster_data = df_melted[df_melted['Cluster'] == cluster]
  922. sns.boxplot(x='Cluster', y='Dice Score', data=cluster_data, width=0.12,
  923. boxprops={'facecolor': colors[i], 'edgecolor': '0.2', 'alpha': 0.8, 'zorder': 10},
  924. whiskerprops={'color': '0.3', 'linewidth': 1.5},
  925. capprops={'color': '0.3', 'linewidth': 1.5},
  926. medianprops={'color': 'black', 'linewidth': 2.5, 'zorder': 11},
  927. showfliers=False, ax=ax_b)
  928. sns.stripplot(x='Cluster', y='Dice Score', data=df_melted, palette=colors,
  929. size=3, jitter=0.15, alpha=0.5, dodge=False, ax=ax_b, hue='Cluster', legend=False)
  930. ax_b.set_title(f'Panel B: {best_k} Communities Stability\n(Dice Distribution)', fontsize=14, fontweight='bold', pad=15)
  931. ax_b.set_ylim(0, 1.05)
  932. sns.despine(ax=ax_b, left=True, bottom=True)
  933. # --- Panel C: Heatmap 2 ---
  934. cm_orighalf_data = np.concatenate((stats['conf_mat_origh1'], stats['conf_mat_origh2']))
  935. cm_orighalf = np.mean(cm_orighalf_data, axis=0)
  936. cm_orighalf_norm = cm_orighalf.astype('float') / (cm_orighalf.sum(axis=1)[:, np.newaxis] + 1e-10)
  937. sns.heatmap(cm_orighalf_norm, annot=True, fmt=".2f", cmap="Blues", ax=axes[2],
  938. xticklabels=cluster_names, yticklabels=cluster_names,
  939. square=True, cbar_kws={'label': 'Proportion of Samples'})
  940. axes[2].set_title(f"Panel C: Alignment Accuracy (K={best_k})\n(Origin vs. Aligned Half)",
  941. fontsize=14, fontweight='bold', pad=15)
  942. axes[2].set_xlabel("Aligned Half Labels", fontsize=12)
  943. axes[2].set_ylabel("Origin Labels", fontsize=12)
  944. # --- Panel D: Raincloud Plot 2 ---
  945. ax_d = axes[3]
  946. sns.violinplot(x='Cluster', y='Dice Score', data=df_melted, palette=colors,
  947. alpha=0.3, inner=None, density_norm='width', ax=ax_d, hue='Cluster', legend=False)
  948. for violin in [c for c in ax_d.collections if isinstance(c, plt.matplotlib.collections.PolyCollection)]:
  949. for path in violin.get_paths():
  950. m = path.vertices[:, 0].mean()
  951. path.vertices[:, 0] = np.clip(path.vertices[:, 0], m, np.inf)
  952. for i, cluster in enumerate(cluster_names):
  953. cluster_data = df_melted[df_melted['Cluster'] == cluster]
  954. sns.boxplot(x='Cluster', y='Dice Score', data=cluster_data, width=0.12,
  955. boxprops={'facecolor': colors[i], 'edgecolor': '0.2', 'alpha': 0.8, 'zorder': 10},
  956. showfliers=False,ax=ax_d)
  957. sns.stripplot(x='Cluster', y='Dice Score', data=df_melted, palette=colors, size=3, ax=ax_d, hue='Cluster', legend=False)
  958. ax_d.set_title(f'Panel D: {best_k} Communities Stability', fontsize=14, fontweight='bold', pad=15)
  959. ax_d.set_ylim(0, 1.05)
  960. sns.despine(ax=ax_d, left=True, bottom=True)
  961. # 7. Save and show
  962. plt.tight_layout()
  963. if output_path:
  964. os.makedirs(output_path, exist_ok=True)
  965. save_path = os.path.join(output_path, 'validation', f'split-half_reliability_k{best_k}_{prefix}.png')
  966. plt.savefig(save_path, dpi=300, bbox_inches='tight')
  967. plt.show()
  968. from matplotlib.collections import PolyCollection
  969. def draw_raincloud_plot(ax,
  970. data,
  971. labels,
  972. colors,
  973. yrange,
  974. title="",
  975. ylabel="Value",
  976. xlabel="Group"):
  977. """
  978. Internal utility to draw a horizontal or vertical raincloud plot on a given axis.
  979. Parameters:
  980. - ax (matplotlib.axes.Axes): The target axis to draw on.
  981. - data (pd.DataFrame): Long-format dataframe with columns matched to seaborn mapping.
  982. - labels (list): List of group names (e.g., ['C1', 'C2', ...]).
  983. - colors (list): List of RGB/RGBA tuples for each group.
  984. - title (str): Title of the subplot.
  985. - ylabel (str): Label for the y-axis (usually the metric name).
  986. - xlabel (str): Label for the x-axis (usually the group ID).
  987. """
  988. from matplotlib.ticker import MaxNLocator
  989. # 1. Draw half-violin plot (The Cloud)
  990. # 'density_norm' is the parameter in newer seaborn versions, replaces 'scale'
  991. sns.violinplot(
  992. x='Group', y='Score', data=data,
  993. palette=colors, alpha=0.3, inner=None, density_norm='width',
  994. ax=ax, hue='Group', legend=False
  995. )
  996. # Clip the violin plots to show only the right half
  997. for violin in [c for c in ax.collections if isinstance(c, PolyCollection)]:
  998. for path in violin.get_paths():
  999. m = path.vertices[:, 0].mean()
  1000. path.vertices[:, 0] = np.clip(path.vertices[:, 0], m, np.inf)
  1001. # 2. Draw boxplots (The Core)
  1002. # We plot each box individually to ensure colors and z-order are correct
  1003. for i, label in enumerate(labels):
  1004. group_subset = data[data['Group'] == label]
  1005. sns.boxplot(
  1006. x='Group', y='Score', data=group_subset,
  1007. width=0.12,
  1008. boxprops={'facecolor': colors[i], 'edgecolor': '0.3', 'alpha': 0.8, 'zorder': 10},
  1009. whiskerprops={'color': '0.3', 'linewidth': 1.5},
  1010. capprops={'color': '0.3', 'linewidth': 1.5},
  1011. medianprops={'color': 'black', 'linewidth': 2.0, 'zorder': 11},
  1012. showfliers=False, ax=ax
  1013. )
  1014. # 3. Draw strip plot (The Raindrops)
  1015. sns.stripplot(
  1016. x='Group', y='Score', data=data,
  1017. palette=colors, size=3, jitter=0.15, alpha=0.4,
  1018. dodge=False, ax=ax, hue='Group', legend=False, zorder=1
  1019. )
  1020. # 4. Aesthetic refinements
  1021. ax.set_title(title, fontsize=14, fontweight='bold', pad=15)
  1022. ax.set_ylabel(ylabel, fontsize=14)
  1023. ax.set_xlabel(xlabel, fontsize=14)
  1024. ax.tick_params(axis='x', labelsize=14)
  1025. ax.tick_params(axis='y', labelsize=14)
  1026. ax.set_ylim(yrange)
  1027. # ax.set_yticks(np.linspace(0, 1, 5))
  1028. ax.grid(axis='y', linestyle='--', alpha=0.3)
  1029. sns.despine(ax=ax, left=True, bottom=False)
  1030. ax.spines['bottom'].set_linewidth(0.5)
  1031. ax.spines['bottom'].set_color('#444444')
  1032. return ax
  1033. def process_iteration_half_split(n_iter, half_type, labels, idx_half, base_coupling,
  1034. ci_matrix_null_dicts, ci_matrix_null, labels_dict,
  1035. output_path, N):
  1036. network_file = f'{output_path}/half_split/network_result_{half_type}_n{n_iter}.pkl'
  1037. coupling_file = f'{output_path}/half_split/LI_coupling_{half_type}_n{n_iter}.pkl'
  1038. edge_file = f'{output_path}/half_split/LI_edge_{half_type}_n{n_iter}.pkl'
  1039. if os.path.exists(network_file) and os.path.exists(coupling_file) and os.path.exists(edge_file):
  1040. network_result = load_as_pickle(network_file)
  1041. LI_coupling = load_as_pickle(coupling_file)
  1042. LI_edge_combined = load_as_pickle(edge_file)
  1043. LI_obs_iter, LI_sig_iter = LI_edge_combined['obs'], LI_edge_combined['sig']
  1044. else:
  1045. ci_matrix = flat_to_assemble_matrix(labels[n_iter])
  1046. arousal_tvFC_coupling = base_coupling[idx_half[n_iter], :]
  1047. arousal_coupling_matrix = np.stack([
  1048. flat_to_assemble_matrix(arousal_tvFC_coupling[i])
  1049. for i in range(arousal_tvFC_coupling.shape[0])
  1050. ])
  1051. ci_matrix_dicts = {
  1052. 'cortex': ci_matrix,
  1053. 'LL': ci_matrix[:N//2, :N//2],
  1054. 'RR': ci_matrix[N//2:, N//2:],
  1055. 'LR': ci_matrix[:N//2, N//2:],
  1056. 'RL': ci_matrix[N//2:, :N//2]
  1057. }
  1058. # 1. Edge pipeline
  1059. LI_obs_iter, LI_sig_iter = run_community_edge_part_pipeline(
  1060. ci_matrix_dicts, ci_matrix_null_dicts, labels_dict, output_path
  1061. )
  1062. save_as_pickle({'obs': LI_obs_iter, 'sig': LI_sig_iter}, edge_file)
  1063. # 2. Node pipeline
  1064. network_result = run_community_node_part_pipeline(
  1065. ci_matrix, ci_matrix_null, labels_dict, output_path, n_perm=10000
  1066. )
  1067. save_as_pickle(network_result, network_file)
  1068. # 3. Coupling pipeline
  1069. LI_coupling = run_community_coupling_part_pipeline(
  1070. ci_matrix_dicts, arousal_coupling_matrix, labels_dict
  1071. )
  1072. save_as_pickle(LI_coupling, coupling_file)
  1073. return {
  1074. 'n_iter': n_iter,
  1075. 'obs_edge': LI_obs_iter,
  1076. 'sig_edge': LI_sig_iter,
  1077. 'node': network_result,
  1078. 'coupling': LI_coupling
  1079. }
  1080. def concordance_correlation_coefficient(y_true, y_pred):
  1081. """
  1082. Computes the Concordance Correlation Coefficient (CCC).
  1083. CCC measures both precision (correlation) and accuracy (deviation from 45-degree line).
  1084. Range: [-1, 1], where 1 is perfect agreement.
  1085. """
  1086. # Remove NaNs
  1087. mask = ~np.isnan(y_true) & ~np.isnan(y_pred)
  1088. if not np.any(mask):
  1089. return np.nan
  1090. y_t = y_true[mask]
  1091. y_p = y_pred[mask]
  1092. if len(y_t) < 2:
  1093. return np.nan
  1094. # Calculate means and variances
  1095. mean_t = np.mean(y_t)
  1096. mean_p = np.mean(y_p)
  1097. var_t = np.var(y_t)
  1098. var_p = np.var(y_p)
  1099. covar = np.mean((y_t - mean_t) * (y_p - mean_p))
  1100. # CCC Formula: (2 * covariance) / (var_t + var_p + (mean_t - mean_p)^2)
  1101. numerator = 2 * covar
  1102. denominator = var_t + var_p + (mean_t - mean_p)**2
  1103. if denominator == 0:
  1104. return np.nan
  1105. return numerator / denominator
  1106. def analyze_stability_core_v1(LI_obs,
  1107. LI_obs_h1_iters,
  1108. LI_obs_h2_iters,
  1109. min_denominator=1e-4):
  1110. """
  1111. Analyzes the consistency and reliability of Lateralization Index (LI)
  1112. using split-half iterations and Concordance Correlation Coefficient (CCC).
  1113. Args:
  1114. LI_obs: Dict containing original observed LI ('inte', 'segre').
  1115. LI_obs_h1_iters: Dict containing LI from the first half iterations.
  1116. LI_obs_h2_iters: Dict containing LI from the second half iterations.
  1117. min_denominator: Small value to avoid division by zero.
  1118. """
  1119. n_comm = LI_obs['inte'].shape[0]
  1120. n_net = LI_obs['inte'].shape[1]
  1121. n_iter = len(LI_obs_h1_iters['inte'])
  1122. types = ['inte', 'segre']
  1123. # Initialize results structure
  1124. stats_results = {t: {
  1125. # 1. Bias metrics (Point-wise accuracy)
  1126. 'abs_bias': np.zeros((n_comm, n_net, n_net)),
  1127. 'rel_bias': np.full((n_comm, n_net, n_net), np.nan),
  1128. 'z_score': np.full((n_comm, n_net, n_net), np.nan),
  1129. # 2. Pattern consistency (CCC is better for continuous agreement)
  1130. 'ccc_h1_vs_h2': np.full((n_comm, n_iter), np.nan), # Reliability
  1131. 'ccc_h1_vs_orig': np.full((n_comm, n_iter), np.nan), # Representativeness
  1132. # 3. Summary statistics
  1133. 'avg_reliability_ccc': np.full(n_comm, np.nan),
  1134. 'avg_stability_ccc': np.full(n_comm, np.nan)
  1135. } for t in types}
  1136. for t in types:
  1137. # Convert list to array for vectorized access
  1138. h1_stack = np.array(LI_obs_h1_iters[t])
  1139. h2_stack = np.array(LI_obs_h2_iters[t])
  1140. for icom in range(n_comm):
  1141. orig_pattern = LI_obs[t][icom].flatten()
  1142. h1_h2_ccc_list = []
  1143. h_orig_ccc_list = []
  1144. for i in range(n_iter):
  1145. p1 = h1_stack[i, icom].flatten()
  1146. p2 = h2_stack[i, icom].flatten()
  1147. # --- Step A: Calculate Split-half Reliability (CCC) ---
  1148. # CCC is more strict than correlation as it checks for y=x
  1149. c_val = concordance_correlation_coefficient(p1, p2)
  1150. stats_results[t]['ccc_h1_vs_h2'][icom, i] = c_val
  1151. if not np.isnan(c_val):
  1152. h1_h2_ccc_list.append(c_val)
  1153. # --- Step B: Calculate Representativeness (CCC) ---
  1154. c_stab = concordance_correlation_coefficient(p1, orig_pattern)
  1155. stats_results[t]['ccc_h1_vs_orig'][icom, i] = c_stab
  1156. if not np.isnan(c_stab):
  1157. h_orig_ccc_list.append(c_stab)
  1158. # --- Step C: Aggregate CCC Scores ---
  1159. if h1_h2_ccc_list:
  1160. stats_results[t]['avg_reliability_ccc'][icom] = np.mean(h1_h2_ccc_list)
  1161. if h_orig_ccc_list:
  1162. stats_results[t]['avg_stability_ccc'][icom] = np.mean(h_orig_ccc_list)
  1163. # --- Step D: Point-wise Bias and Z-score ---
  1164. # Combine h1 and h2 to get a distribution for each element
  1165. combined_iters = np.concatenate([h1_stack[:, icom, :, :], h2_stack[:, icom, :, :]], axis=0)
  1166. m_iter = np.nanmean(combined_iters, axis=0)
  1167. s_iter = np.nanstd(combined_iters, axis=0)
  1168. orig_vals = LI_obs[t][icom]
  1169. # 1. Absolute Bias
  1170. abs_bias = m_iter - orig_vals
  1171. stats_results[t]['abs_bias'][icom] = abs_bias
  1172. # 2. Relative Bias
  1173. denom = np.abs(orig_vals)
  1174. valid_rel = denom > min_denominator
  1175. stats_results[t]['rel_bias'][icom][valid_rel] = abs_bias[valid_rel] / denom[valid_rel]
  1176. # 3. Z-Score (How many SDs original value is from iteration mean)
  1177. valid_z = s_iter > min_denominator
  1178. stats_results[t]['z_score'][icom][valid_z] = (orig_vals[valid_z] - m_iter[valid_z]) / s_iter[valid_z]
  1179. return stats_results
  1180. def analyze_stability_core(obs_data, h1_stack, h2_stack, n_iter, min_denominator=1e-4):
  1181. # 1. Pattern consistency (CCC)
  1182. ccc_h1_h2 = np.array([concordance_correlation_coefficient(h1_stack[i].flatten(), h2_stack[i].flatten()) for i in range(n_iter)])
  1183. ccc_h1_orig = np.array([concordance_correlation_coefficient(h1_stack[i].flatten(), obs_data.flatten()) for i in range(n_iter)])
  1184. # 2. Point-wise metrics
  1185. combined = np.concatenate([h1_stack, h2_stack], axis=0)
  1186. m_iter = np.nanmean(combined, axis=0)
  1187. s_iter = np.nanstd(combined, axis=0)
  1188. abs_bias = m_iter - obs_data
  1189. # Relative Bias
  1190. rel_bias = np.full(obs_data.shape, np.nan)
  1191. denom = np.abs(obs_data)
  1192. mask_rel = denom > min_denominator
  1193. rel_bias[mask_rel] = abs_bias[mask_rel] / denom[mask_rel]
  1194. # Z-score
  1195. z_score = np.full(obs_data.shape, np.nan)
  1196. mask_z = s_iter > min_denominator
  1197. z_score[mask_z] = (obs_data[mask_z] - m_iter[mask_z]) / s_iter[mask_z]
  1198. return {
  1199. 'ccc_h1_vs_h2': ccc_h1_h2,
  1200. 'ccc_h1_vs_orig': ccc_h1_orig,
  1201. 'avg_reliability_ccc': np.nanmean(ccc_h1_h2),
  1202. 'avg_stability_ccc': np.nanmean(ccc_h1_orig),
  1203. 'abs_bias': abs_bias,
  1204. 'rel_bias': rel_bias,
  1205. 'z_score': z_score
  1206. }
  1207. def analyze_LI_all_levels_stability(
  1208. LI_obs_edge, LI_h1_edge_iters, LI_h2_edge_iters,
  1209. LI_obs_node, LI_h1_node_iters, LI_h2_node_iters,
  1210. LI_obs_coupling, LI_h1_coupling_iters, LI_h2_coupling_iters,
  1211. min_denominator=1e-4
  1212. ):
  1213. n_iter = 500#len(LI_h1_edge_iters['inte'])
  1214. types = ['inte', 'segre']
  1215. results = {}
  1216. # --- 1. Edge Level ---
  1217. results['edge'] = {t: {} for t in types}
  1218. for t in types:
  1219. h1_stack = np.array(LI_h1_edge_iters[t]) # (n_iter, n_comm, n_net, n_net)
  1220. h2_stack = np.array(LI_h2_edge_iters[t])
  1221. n_comm = LI_obs_edge[t].shape[0]
  1222. comm_results = []
  1223. for icom in range(n_comm):
  1224. res = analyze_stability_core(LI_obs_edge[t][icom], h1_stack[:, icom], h2_stack[:, icom], n_iter, min_denominator)
  1225. comm_results.append(res)
  1226. for key in comm_results[0].keys():
  1227. results['edge'][t][key] = np.array([r[key] for r in comm_results])
  1228. # --- 2. Node Level ---
  1229. results['node'] = {t: {} for t in types}
  1230. for t in types:
  1231. obs_node_vals = np.array(LI_obs_node[f'mean_{t}']) # (n_comm, n_nodes)
  1232. # h1_node_vals = np.array([it['mean'] for it in LI_h1_node_iters[t]]) # (n_iter, n_comm, n_nodes)
  1233. h1_node_stack = np.swapaxes(np.array([iter_list['mean'] for iter_list in LI_h1_node_iters[t]]), 0, 1)
  1234. h2_node_stack = np.swapaxes(np.array([iter_list['mean'] for iter_list in LI_h2_node_iters[t]]), 0, 1)
  1235. comm_node_res = []
  1236. for icom in range(n_comm):
  1237. res = analyze_stability_core(obs_node_vals[icom], h1_node_stack[icom], h2_node_stack[icom], n_iter, min_denominator)
  1238. comm_node_res.append(res)
  1239. for key in comm_node_res[0].keys():
  1240. results['node'][t][key] = np.array([r[key] for r in comm_node_res])
  1241. results['node'][t]['half_stack'] = np.concat((h1_node_stack,h2_node_stack),axis=1)
  1242. results['node'][t]['obs_values'] = obs_node_vals
  1243. # --- 3. Coupling Level ---
  1244. feature_full = 'slope_pair_comm'
  1245. feat_name = feature_full.split('_')[0]
  1246. res_key = f'coupling_{feat_name}'
  1247. results[res_key] = {t: {} for t in types}
  1248. for t in types:
  1249. obs_vals = np.array([np.nanmean(LI_obs_coupling[t][feature_full][icom]) for icom in range(n_comm)])
  1250. h1_vals = np.swapaxes(np.array([np.mean(np.array(iter_list[feat_name]),axis=1)
  1251. for iter_list in LI_h1_coupling_iters[t]]),
  1252. 0, 1)
  1253. h2_vals = np.swapaxes(np.array([np.mean(np.array(iter_list[feat_name]),axis=1)
  1254. for iter_list in LI_h2_coupling_iters[t]]),
  1255. 0, 1)
  1256. comm_cp_res = []
  1257. for icom in range(n_comm):
  1258. res = analyze_stability_core(obs_vals[icom], h1_vals[icom], h2_vals[icom], n_iter, min_denominator)
  1259. comm_cp_res.append(res)
  1260. for key in comm_cp_res[0].keys():
  1261. results[res_key][t][key] = np.array([r[key] for r in comm_cp_res])
  1262. results[res_key][t]['half_stack'] = np.concat((h1_vals,h2_vals),axis=1)
  1263. results[res_key][t]['obs_values'] = obs_vals
  1264. feature_full = 'strength_pair_comm'
  1265. feat_name = feature_full.split('_')[0]
  1266. res_key = f'coupling_{feat_name}'
  1267. results[res_key] = {t: {} for t in types}
  1268. for t in types:
  1269. obs_vals = np.array([np.nanmean(LI_obs_coupling[t][feature_full][icom]) for icom in range(n_comm)])
  1270. h1_vals = np.swapaxes(np.array([np.nanmean(np.array(iter_list[feat_name]),axis=(1,2))
  1271. for iter_list in LI_h1_coupling_iters[t]]),
  1272. 0, 1)
  1273. h2_vals = np.swapaxes(np.array([np.nanmean(np.array(iter_list[feat_name]),axis=(1,2))
  1274. for iter_list in LI_h2_coupling_iters[t]]),
  1275. 0, 1)
  1276. comm_cp_res = []
  1277. for icom in range(n_comm):
  1278. res = analyze_stability_core(obs_vals[icom], h1_vals[icom], h2_vals[icom], n_iter, min_denominator)
  1279. comm_cp_res.append(res)
  1280. for key in comm_cp_res[0].keys():
  1281. results[res_key][t][key] = np.array([r[key] for r in comm_cp_res])
  1282. results[res_key][t]['half_stack'] = np.concat((h1_vals,h2_vals),axis=1)
  1283. results[res_key][t]['obs_values'] = obs_vals
  1284. return results
  1285. def run_single_iteration(data, sub_ids, k, original_labels):
  1286. unique_subs = np.unique(sub_ids)
  1287. n_subs = len(unique_subs)
  1288. n_edges = data.shape[1]
  1289. sub_sample_matrix = np.zeros((n_subs, n_edges))
  1290. selected_run_indices = []
  1291. for i, sub in enumerate(unique_subs):
  1292. sub_indices = np.where(sub_ids == sub)[0]
  1293. chosen_idx = np.random.choice(sub_indices)
  1294. selected_run_indices.append(chosen_idx)
  1295. sub_sample_matrix[i, :] = data[chosen_idx, :]
  1296. edge_features = sub_sample_matrix.T
  1297. kmeans_sub = KMeans(n_clusters=k, n_init=10, random_state=42).fit(edge_features)
  1298. labels_sub = kmeans_sub.labels_+1
  1299. aligned_sub = align_to_reference(labels_sub, original_labels,k)
  1300. # sns.heatmap(flat_to_assemble_matrix(labels_sub),vmin=1,cmap=cmap_k5)
  1301. ari = adjusted_rand_score(original_labels, aligned_sub)
  1302. nmi = normalized_mutual_info_score(original_labels, aligned_sub)
  1303. def calculate_dice(l1, l2, k):
  1304. dices = []
  1305. for i in range(k):
  1306. mask1 = (l1 == (i+1))
  1307. mask2 = (l2 == (i+1))
  1308. intersection = np.sum(mask1 & mask2)
  1309. sum_val = np.sum(mask1) + np.sum(mask2)
  1310. dice = (2. * intersection) / (sum_val + 1e-10)
  1311. dices.append(dice)
  1312. return np.mean(dices), dices
  1313. avg_dice, per_cluster_dice = calculate_dice(original_labels, aligned_sub, k)
  1314. return {
  1315. "ari": ari,
  1316. "nmi": nmi,
  1317. "mean_dice": avg_dice,
  1318. "per_cluster_dice": per_cluster_dice,
  1319. "aligned_labels": aligned_sub,
  1320. "selected_run_indices": selected_run_indices
  1321. }
  1322. def analyze_subject_stability(data,
  1323. sub_ids,
  1324. original_labels,
  1325. k,
  1326. n_iterations=50,
  1327. n_jobs=1):
  1328. print(f"Starting Stability Analysis: k={k}, iterations={n_iterations}")
  1329. unique_subs = np.unique(sub_ids)
  1330. n_edges = data.shape[1]
  1331. full_sub_matrix = np.zeros((len(unique_subs), n_edges))
  1332. for i, sub in enumerate(unique_subs):
  1333. full_sub_matrix[i, :] = np.mean(data[sub_ids == sub], axis=0)
  1334. kmeans_strat1 = KMeans(n_clusters=k, n_init=20, random_state=42).fit(full_sub_matrix.T)
  1335. strat1_aligned = align_to_reference(kmeans_strat1.labels_+1, original_labels,k)
  1336. ari_strat1 = adjusted_rand_score(original_labels, strat1_aligned)
  1337. nmi_strat1 = normalized_mutual_info_score(original_labels, strat1_aligned)
  1338. def calculate_dice(l1, l2, k):
  1339. dices = []
  1340. for i in range(k):
  1341. mask1 = (l1 == (i+1))
  1342. mask2 = (l2 == (i+1))
  1343. intersection = np.sum(mask1 & mask2)
  1344. sum_val = np.sum(mask1) + np.sum(mask2)
  1345. dice = (2. * intersection) / (sum_val + 1e-10)
  1346. dices.append(dice)
  1347. return np.mean(dices), dices
  1348. avg_dice, per_cluster_dice = calculate_dice(original_labels, strat1_aligned, k)
  1349. conf_mat = confusion_matrix(
  1350. original_labels,
  1351. strat1_aligned,
  1352. labels=range(1,k+1),
  1353. normalize='true'
  1354. )
  1355. print(f"Running Strategy 2 on {n_jobs} CPU cores...")
  1356. results_list = Parallel(n_jobs=n_jobs)(
  1357. delayed(run_single_iteration)(data, sub_ids, k, original_labels)
  1358. for _ in range(n_iterations)
  1359. )
  1360. all_edge_labels = [r["aligned_labels"] for r in results_list]
  1361. all_sampled_indices = [r["selected_run_indices"] for r in results_list]
  1362. ari_list = [r["ari"] for r in results_list]
  1363. nmi_list = [r["nmi"] for r in results_list]
  1364. mean_dice_list = [r["mean_dice"] for r in results_list]
  1365. all_per_cluster_dice = np.array([r["per_cluster_dice"] for r in results_list])
  1366. ari_array = np.array(ari_list)
  1367. print("Stability analysis workflow completed.")
  1368. return {
  1369. "original_labels": original_labels,
  1370. "strat1_results": {
  1371. "ari": ari_strat1,
  1372. "aligned_labels": strat1_aligned,
  1373. 'confusion_matrix':conf_mat
  1374. },
  1375. "strat2_metrics": {
  1376. "ari": ari_list,
  1377. "nmi": nmi_list,
  1378. "mean_dice": mean_dice_list,
  1379. "per_cluster_dice": all_per_cluster_dice,
  1380. "all_edge_labels": all_edge_labels,
  1381. "all_sampled_indices": all_sampled_indices
  1382. },
  1383. }
  1384. def plot_LI_stability(all_levels_obs_results,
  1385. all_levels_sig_results,
  1386. net_yeo7,
  1387. cmap_k,
  1388. output_path,
  1389. mode):
  1390. net_yeo7 = ['DMN', 'FPN', 'LIMB', 'VAN', 'DAN', 'SMN', 'VIS']
  1391. types = ['inte', 'segre']
  1392. n_comm = all_levels_obs_results['edge']['inte']['abs_bias'].shape[0]
  1393. comm_labels = [f'C{i+1}' for i in range(n_comm)]
  1394. colors_list = [cmap_k(i) for i in range(n_comm)]
  1395. sns.set_context("talk")
  1396. sns.set_style("white")
  1397. result_obs_edge = all_levels_obs_results['edge']
  1398. result_sig_edge = all_levels_sig_results['edge']
  1399. result_node = all_levels_obs_results['node']
  1400. result_slope = all_levels_obs_results['coupling_slope']
  1401. result_strength = all_levels_obs_results['coupling_strength']
  1402. os.makedirs(f"{output_path}/validation", exist_ok=True)
  1403. # --- 1 & 2. Edge & Node CCC Plots ---
  1404. for level_name, result_dict in zip(['edge', 'node'], [result_obs_edge, result_node]):
  1405. for t in types:
  1406. for metric in ['ccc_h1_vs_h2', 'ccc_h1_vs_orig']:
  1407. fig, ax = plt.subplots(figsize=(10, 5))
  1408. df_corr = pd.DataFrame(result_dict[t][metric].T, columns=comm_labels)
  1409. df_melted = df_corr.melt(var_name='Group', value_name='Score')
  1410. suffix = "reliability" if "h2" in metric else "stability"
  1411. ax = draw_raincloud_plot(ax, df_melted, comm_labels, colors_list, (0,1),
  1412. title=f"{level_name.capitalize()}_{t}_{suffix}",
  1413. ylabel="CCC", xlabel="Communities ID")
  1414. plt.savefig(f"{output_path}/validation/{mode}_LI-{level_name}_CCC-{suffix}_{t}.png", dpi=600, bbox_inches='tight')
  1415. plt.close()
  1416. # --- 3. Node Level - Stability Comparison (Box + Strip) ---
  1417. for t in types:
  1418. iter_data = result_node[t]['half_stack']
  1419. obs_vals = result_node[t].get('obs_values', None)
  1420. node_records, obs_records, relative_records = [], [], []
  1421. for icom in range(n_comm):
  1422. for inode, net_name in enumerate(net_yeo7):
  1423. node_iters = iter_data[icom, :, inode]
  1424. mean_resampled = np.mean(iter_data[icom, :, inode])
  1425. original_val = obs_vals[icom, inode]
  1426. if original_val != 0:
  1427. relative_val = ((mean_resampled - original_val) / np.abs(original_val)) * 100
  1428. else:
  1429. relative_val = 0
  1430. relative_records.append({
  1431. 'Network': net_name,
  1432. 'Community': f'C{icom+1}',
  1433. 'Relative Value (%)': relative_val
  1434. })
  1435. for val in node_iters:
  1436. node_records.append({'Network': net_name, 'Community': f'C{icom+1}', 'Lateralization index': val})
  1437. if obs_vals is not None:
  1438. obs_records.append({'Network': net_name, 'Community': f'C{icom+1}', 'LI_obs': obs_vals[icom, inode]})
  1439. df_iter, df_obs = pd.DataFrame(node_records), pd.DataFrame(obs_records).groupby(['Network', 'Community']).mean().reset_index()
  1440. plt.figure(figsize=(12, 4))
  1441. ax = sns.boxplot(data=df_iter, x='Network', y='Lateralization index', hue='Community', palette=colors_list, showfliers=False, whis=1.5, boxprops=dict(alpha=.3))
  1442. sns.stripplot(data=df_obs, x='Network', y='LI_obs', hue='Community', dodge=True, marker='D', size=8, edgecolor='black', linewidth=1, palette=colors_list, ax=ax)
  1443. ax.tick_params(axis='both', which='major', direction='out', length=6, width=1.5, colors='black', bottom=True, left=True)
  1444. if ax.get_legend() is not None: ax.get_legend().remove()
  1445. ax.axhline(0, color='black', linestyle='--', linewidth=1.2, alpha=0.6)
  1446. plt.title(f"Node-level LI Distribution across Networks ({t})")
  1447. plt.savefig(f"{output_path}/validation/{mode}_node_stability_comparison_{t}.png", dpi=600, bbox_inches='tight')
  1448. plt.close()
  1449. # --- 4. Slope & Strength Stability Comparison ---
  1450. for metric_name, result_dict, yrange in zip(['slope', 'strength'],
  1451. [result_slope, result_strength],
  1452. [(0.1,0.6),(-0.02,0.02)]):
  1453. for t in types:
  1454. iter_data = result_dict[t]['half_stack'] # (n_comm, n_iter)
  1455. obs_vals = result_dict[t].get('obs_values', None) # (n_comm,)
  1456. records = []
  1457. relative_records = []
  1458. for icom in range(n_comm):
  1459. iters = iter_data[icom, :]
  1460. for val in iters:
  1461. records.append({'Group': f'C{icom+1}', 'Score': val})
  1462. mean_resampled = np.mean(iters)
  1463. original_val = obs_vals[icom] if obs_vals is not None else 0
  1464. if original_val != 0:
  1465. relative_val = ((mean_resampled - original_val) / np.abs(original_val)) * 100
  1466. else:
  1467. relative_val = 0
  1468. relative_records.append({
  1469. 'Community': f'C{icom+1}',
  1470. 'Relative Value (%)': relative_val
  1471. })
  1472. df_iter = pd.DataFrame(records)
  1473. df_relative = pd.DataFrame(relative_records)
  1474. obs_recs = []
  1475. if obs_vals is not None:
  1476. for icom in range(n_comm):
  1477. obs_recs.append({'Group': f'C{icom+1}', 'Score': obs_vals[icom]})
  1478. df_obs = pd.DataFrame(obs_recs)
  1479. fig, ax = plt.subplots(figsize=(6, 5))
  1480. ax = draw_raincloud_plot(ax, df_iter, comm_labels, colors_list, yrange,
  1481. title=f"{metric_name} LI Stability Distribution ({t})",
  1482. ylabel="Lateralization Index", xlabel="Communities ID")
  1483. if not df_obs.empty:
  1484. sns.stripplot(data=df_obs, x='Group', y='Score',
  1485. order=comm_labels,
  1486. marker='D', size=10, color='white',
  1487. edgecolor='black', linewidth=1.5,
  1488. jitter=False,
  1489. ax=ax, zorder=10)
  1490. if ax.get_legend() is not None:
  1491. ax.get_legend().remove()
  1492. ax.axhline(0, color='black', linestyle='--', alpha=0.4, zorder=0)
  1493. plt.tight_layout()
  1494. plt.savefig(f"{output_path}/validation/{mode}_{metric_name}_LI_stability_comparison_{t}.png", dpi=600, bbox_inches='tight')
  1495. # plt.close()
  1496. plt.figure(figsize=(6, 4))
  1497. ax_rel = sns.barplot(
  1498. data=df_relative,
  1499. x='Community',
  1500. y='Relative Value (%)',
  1501. palette=colors_list,
  1502. edgecolor='black',
  1503. linewidth=1
  1504. )
  1505. ax_rel.axhline(0, color='black', linestyle='-', linewidth=1.2)
  1506. ax_rel.tick_params(axis='both', which='major', direction='out', length=6, width=1.5)
  1507. plt.ylim(-50, 50)
  1508. plt.title(f"Relative Deviation: {metric_name} ({t})")
  1509. plt.ylabel("Relative Value (%)")
  1510. plt.xlabel("Communities ID")
  1511. plt.tight_layout()
  1512. plt.savefig(f"{output_path}/validation/{mode}_{metric_name}_relative_deviation_{t}.png", dpi=600, bbox_inches='tight')
  1513. plt.close()
  1514. all_stats_node_records = []
  1515. for t in types:
  1516. iter_data = result_node[t]['half_stack']
  1517. obs_vals = result_node[t]['obs_values']
  1518. for icom in range(n_comm):
  1519. for inode, net_name in enumerate(net_yeo7):
  1520. x_bar = np.mean(iter_data[icom, :, inode])
  1521. x_0 = obs_vals[icom, inode]
  1522. rel_val = ((x_bar - x_0) / np.abs(x_0)) * 100 if x_0 != 0 else 0
  1523. all_stats_node_records.append({
  1524. 'Metric': f'Node_LI_{t}',
  1525. 'Community': f'C{icom+1}',
  1526. 'Network': net_name,
  1527. 'Original': x_0,
  1528. 'Resampled_Mean': x_bar,
  1529. 'Relative_Error_%': rel_val
  1530. })
  1531. all_stats_coupling_records = []
  1532. for metric_name, result_dict in zip(['slope', 'strength'], [result_slope, result_strength]):
  1533. for t in types:
  1534. iter_data = result_dict[t]['half_stack']
  1535. obs_vals = result_dict[t]['obs_values']
  1536. for icom in range(n_comm):
  1537. x_bar = np.mean(iter_data[icom, :])
  1538. x_0 = obs_vals[icom]
  1539. rel_val = ((x_bar - x_0) / np.abs(x_0)) * 100 if x_0 != 0 else 0
  1540. all_stats_coupling_records.append({
  1541. 'Metric': f'{metric_name}_{t}',
  1542. 'Community': f'C{icom+1}',
  1543. 'Network': 'N/A',
  1544. 'Original': x_0,
  1545. 'Resampled_Mean': x_bar,
  1546. 'Relative_Error_%': rel_val
  1547. })
  1548. df_node = pd.DataFrame(all_stats_node_records)
  1549. df_node['Absolute_Dev'] = df_node['Resampled_Mean'] - df_node['Original']
  1550. df_node['Network'] = pd.Categorical(
  1551. df_node['Network'],
  1552. categories=['DMN', 'FPN', 'LIMB', 'VAN', 'DAN', 'SMN', 'VIS'],
  1553. ordered=True
  1554. )
  1555. for metric in df_node['Metric'].unique():
  1556. subset = df_node[df_node['Metric'] == metric]
  1557. pivot_abs = subset.pivot(index='Community', columns='Network', values='Absolute_Dev')
  1558. pivot_obs = subset.pivot(index='Community', columns='Network', values='Original')
  1559. pivot_rel = subset.pivot(index='Community', columns='Network', values='Relative_Error_%')
  1560. annot_matrix = []
  1561. for i in range(pivot_abs.shape[0]):
  1562. row_annot = []
  1563. for j in range(pivot_abs.shape[1]):
  1564. a_dev = pivot_abs.iloc[i, j]
  1565. r_err = pivot_rel.iloc[i, j]
  1566. o_val = pivot_obs.iloc[i, j]
  1567. if abs(o_val) > 0.05:
  1568. row_annot.append(f"$\Delta$:{a_dev:.3f}\n{r_err:.1f}%")
  1569. else:
  1570. row_annot.append(f"$\Delta$:{a_dev:.3f}")
  1571. annot_matrix.append(row_annot)
  1572. plt.figure(figsize=(8, 6))
  1573. sns.heatmap(
  1574. pivot_abs,
  1575. annot=np.array(annot_matrix),
  1576. fmt="",
  1577. cmap="RdBu_r",
  1578. center=0,
  1579. annot_kws={"size": 9, "va": "center"},
  1580. linewidths=.5,
  1581. cbar_kws={'label': 'Absolute Deviation (Resampled Mean - Original)'}
  1582. )
  1583. plt.title(f"Stability Heatmap: {metric}\n(Relative Error only shown for |X0| > 0.05)")
  1584. plt.savefig(f"{output_path}/validation/{mode}_Heatmap_{metric}_stability.png", dpi=600, bbox_inches='tight')
  1585. plt.show()
  1586. df_coupling = pd.DataFrame(all_stats_coupling_records)
  1587. df_coupling['Absolute_Dev'] = df_coupling['Resampled_Mean'] - df_coupling['Original']
  1588. target_metric_order = [
  1589. 'slope_inte', 'slope_segre',
  1590. 'strength_inte', 'strength_segre'
  1591. ]
  1592. df_coupling['Metric'] = pd.Categorical(
  1593. df_coupling['Metric'],
  1594. categories=target_metric_order,
  1595. ordered=True
  1596. )
  1597. pivot_abs = df_coupling.pivot(index='Community', columns='Metric', values='Absolute_Dev')
  1598. pivot_obs = df_coupling.pivot(index='Community', columns='Metric', values='Original')
  1599. pivot_rel = df_coupling.pivot(index='Community', columns='Metric', values='Relative_Error_%')
  1600. annot_matrix = []
  1601. for i in range(pivot_abs.shape[0]):
  1602. row_annot = []
  1603. for j in range(pivot_abs.shape[1]):
  1604. a_dev = pivot_abs.iloc[i, j]
  1605. r_err = pivot_rel.iloc[i, j]
  1606. o_val = pivot_obs.iloc[i, j]
  1607. if abs(o_val) > 0.05:
  1608. row_annot.append(f"$\Delta$:{a_dev:.3f}\n{r_err:.1f}%")
  1609. else:
  1610. row_annot.append(f"$\Delta$:{a_dev:.3f}")
  1611. annot_matrix.append(row_annot)
  1612. plt.figure(figsize=(8, 6))
  1613. ax = sns.heatmap(
  1614. pivot_abs,
  1615. annot=np.array(annot_matrix),
  1616. fmt="",
  1617. cmap="RdBu_r",
  1618. center=0,
  1619. annot_kws={"size": 9, "va": "center"},
  1620. linewidths=1,
  1621. linecolor='white',
  1622. cbar_kws={'label': 'Absolute Deviation ($\Delta$)'}
  1623. )
  1624. plt.title("Stability of Arousal-Coupling Metrics across Communities", pad=20)
  1625. plt.xlabel("Coupling Metrics")
  1626. plt.ylabel("Communities")
  1627. plt.tight_layout()
  1628. plt.savefig(f"{output_path}/validation/{mode}_Combined_Coupling_Stability_Heatmap.png", dpi=600)
  1629. plt.show()
  1630. def Hungarian_alignment(source_labels, target_labels):
  1631. """
  1632. Align source_labels to target_labels using the Hungarian algorithm (Kuhn-Munkres).
  1633. Purpose: Maximizes the overlap (Dice/Accuracy) between two clustering results
  1634. to solve the label switching problem.
  1635. """
  1636. unique_source = np.unique(source_labels)
  1637. unique_target = np.unique(target_labels)
  1638. # Calculate contingency table (intersection counts)
  1639. # Rows = Target, Columns = Source
  1640. cm = confusion_matrix(target_labels, source_labels)
  1641. # Linear sum assignment finds the optimal mapping by minimizing cost.
  1642. # We use -cm because the algorithm minimizes cost, and we want to maximize overlap.
  1643. row_ind, col_ind = linear_sum_assignment(-cm)
  1644. # Create mapping dictionary: {original_source_label: aligned_target_label}
  1645. mapping = {source_lab: target_lab for target_lab, source_lab in zip(row_ind, col_ind)}
  1646. # Map the labels. Use original label as fallback if mapping is missing.
  1647. aligned_labels = np.array([mapping.get(lab, lab) for lab in source_labels])
  1648. return aligned_labels
  1649. def _run_single_iteration(data, sub_ids, k, original_labels):
  1650. """
  1651. Helper function for parallel execution:
  1652. 1. Subsamples 80% of subjects.
  1653. 2. Averages within-subject runs to keep (subject x edge) feature stability.
  1654. 3. Performs KMeans on edges.
  1655. 4. Aligns results and calculates stability metrics (ARI and Dice).
  1656. """
  1657. unique_subs = np.unique(sub_ids)
  1658. # Randomly select 80% of subjects without replacement
  1659. selected_subs = np.random.choice(unique_subs, size=int(0.8 * len(unique_subs)), replace=False)
  1660. n_edges = data.shape[1]
  1661. # Construct feature matrix: rows = selected subjects, columns = edges
  1662. sub_avg_matrix = np.zeros((len(selected_subs), n_edges))
  1663. for i, sub in enumerate(selected_subs):
  1664. sub_mask = (sub_ids == sub)
  1665. sub_avg_matrix[i, :] = np.mean(data[sub_mask, :], axis=0)
  1666. # Transpose to cluster edges: (n_edges, n_subjects_as_features)
  1667. edge_features = sub_avg_matrix.T
  1668. # Clustering with reasonable initialization to ensure quality
  1669. kmeans_sub = KMeans(n_clusters=k, n_init=10, random_state=None).fit(edge_features)
  1670. labels_sub = kmeans_sub.labels_
  1671. # Re-align subsampled labels to original labels to ensure C1 corresponds to C1
  1672. aligned_sub = Hungarian_alignment(labels_sub, original_labels)
  1673. # Global stability metric
  1674. ari = adjusted_rand_score(original_labels, aligned_sub)
  1675. # Community-wise stability metrics
  1676. unique_clusters = np.sort(np.unique(original_labels))
  1677. per_cluster_dice = []
  1678. for cluster_id in unique_clusters:
  1679. intersect = np.sum((original_labels == cluster_id) & (aligned_sub == cluster_id))
  1680. total = np.sum(original_labels == cluster_id) + np.sum(aligned_sub == cluster_id)
  1681. dice = (2.0 * intersect / total) if total > 0 else 0.0
  1682. per_cluster_dice.append(dice)
  1683. return {
  1684. "ari": ari,
  1685. "mean_dice": np.mean(per_cluster_dice),
  1686. "per_cluster_dice": per_cluster_dice, # Needed for Raincloud plot
  1687. "aligned_labels": aligned_sub
  1688. }
  1689. def main_validation_workflow_parallel(data, sub_ids, original_labels, k, n_iterations=50, n_jobs=-1):
  1690. """
  1691. Main workflow for parallelized stability analysis.
  1692. Compares two strategies:
  1693. Strategy 1: Full-subject averaging (Deterministic baseline).
  1694. Strategy 2: Bootstrapped subsampling (Distribution of stability).
  1695. """
  1696. print(f"Starting Stability Analysis: k={k}, iterations={n_iterations}")
  1697. # --- Strategy 1: All-subject Averaging ---
  1698. unique_subs = np.unique(sub_ids)
  1699. n_edges = data.shape[1]
  1700. full_sub_matrix = np.zeros((len(unique_subs), n_edges))
  1701. for i, sub in enumerate(unique_subs):
  1702. full_sub_matrix[i, :] = np.mean(data[sub_ids == sub], axis=0)
  1703. # Cluster on full dataset
  1704. kmeans_strat1 = KMeans(n_clusters=k, n_init=20, random_state=42).fit(full_sub_matrix.T)
  1705. strat1_aligned = Hungarian_alignment(kmeans_strat1.labels_, original_labels)
  1706. # Calculate Strategy 1 metrics
  1707. ari_strat1 = adjusted_rand_score(original_labels, strat1_aligned)
  1708. dice_strat1 = []
  1709. for cid in np.unique(original_labels):
  1710. it = np.sum((original_labels == cid) & (strat1_aligned == cid))
  1711. tt = np.sum(original_labels == cid) + np.sum(strat1_aligned == cid)
  1712. dice_strat1.append(2.0 * it / tt if tt > 0 else 0.0)
  1713. # --- Strategy 2: Parallelized Subsampling ---
  1714. print(f"Running Strategy 2 on {n_jobs if n_jobs != -1 else 'all'} CPU cores...")
  1715. results_list = Parallel(n_jobs=n_jobs)(
  1716. delayed(_run_single_iteration)(data, sub_ids, k, original_labels)
  1717. for _ in range(n_iterations)
  1718. )
  1719. # Aggregate results for distribution analysis
  1720. ari_list = [r["ari"] for r in results_list]
  1721. mean_dice_list = [r["mean_dice"] for r in results_list]
  1722. # Extract (n_iterations, n_clusters) matrix for Raincloud plots
  1723. all_per_cluster_dice = np.array([r["per_cluster_dice"] for r in results_list])
  1724. # Find a 'typical' run (the one closest to median ARI) for visualization
  1725. ari_array = np.array(ari_list)
  1726. median_val = np.median(ari_array)
  1727. median_idx = np.argmin(np.abs(ari_array - median_val))
  1728. typical_labels = results_list[median_idx]["aligned_labels"]
  1729. print("Stability analysis workflow completed.")
  1730. return {
  1731. "original_labels": original_labels,
  1732. "strat1_results": {
  1733. "ari": ari_strat1,
  1734. "mean_dice": np.mean(dice_strat1),
  1735. "per_cluster_dice": dice_strat1,
  1736. "aligned_labels": strat1_aligned
  1737. },
  1738. "strat2_metrics": {
  1739. "ari": ari_list,
  1740. "mean_dice": mean_dice_list,
  1741. "per_cluster_dice": all_per_cluster_dice # Array for visualization
  1742. },
  1743. "strat2_typical_labels": typical_labels
  1744. }
  1745. def plot_subject_stability_metrics(result_subject, cmap, output_path):
  1746. """
  1747. Visualize stability analysis results:
  1748. Panel A: Confusion Matrix for Strategy 1 (Subject Averaging vs. Original)
  1749. Panel B: Raincloud Plot for Strategy 2 (Per-cluster Dice across iterations)
  1750. """
  1751. # 1. Extract data from results dictionary
  1752. orig_labels = result_subject["original_labels"]
  1753. strat1_aligned = result_subject["strat2_metrics"]["all_edge_labels"]
  1754. # Expecting per_cluster_dice shape: (n_iterations, n_clusters)
  1755. # e.g., (500, 7) for 7 communities across 500 subsampling iterations
  1756. per_cluster_dice = result_subject["strat2_metrics"]["per_cluster_dice"]
  1757. n_clusters = per_cluster_dice.shape[1]
  1758. n_iterations = per_cluster_dice.shape[0]
  1759. # 2. Initialize figure with two subplots
  1760. sns.set_theme(style="whitegrid")
  1761. fig, axes = plt.subplots(2, 2, figsize=(15, 12))
  1762. axes = axes.flatten()
  1763. # --- Panel A: Normalized Confusion Matrix (Strategy 1) ---
  1764. # Rows represent original community size normalization
  1765. cm_norm = np.zeros((7,7,500))
  1766. for i in range(len(strat1_aligned)):
  1767. cm = confusion_matrix(orig_labels, strat1_aligned[i])
  1768. cm_norm[:,:,i] = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]
  1769. cm_norm = np.mean(cm_norm,axis=2)
  1770. cluster_ticks = [f"C{i}" for i in np.unique(orig_labels)]
  1771. sns.heatmap(cm_norm, annot=True, fmt=".2f", cmap="Blues", ax=axes[0],
  1772. xticklabels=cluster_ticks,
  1773. yticklabels=cluster_ticks,
  1774. cbar_kws={'label': 'Proportion of Samples'})
  1775. axes[0].set_title("Panel A: Strategy 1 Alignment Accuracy\n(Subject Averaging vs. Original)",
  1776. fontsize=12, fontweight='bold', pad=15)
  1777. axes[0].set_xlabel("Aligned Strategy 1 Labels", fontsize=12)
  1778. axes[0].set_ylabel("Original Communities Labels", fontsize=12)
  1779. # --- Panel B: Community-wise Raincloud Plot (Strategy 2) ---
  1780. # Convert per-cluster dice to long-form DataFrame for Seaborn
  1781. cluster_names = [f'C{i+1}' for i in range(n_clusters)]
  1782. df_dice = pd.DataFrame(per_cluster_dice, columns=cluster_names)
  1783. df_melted = df_dice.melt(var_name='Community', value_name='Dice Score')
  1784. # Define color palette based on clusters
  1785. if cmap is None:
  1786. colors = sns.color_palette("husl", n_clusters)
  1787. else:
  1788. colors = [cmap(i) for i in range(n_clusters)]
  1789. # 1. Draw the "Cloud" (Half-Violin)
  1790. # inner=None to keep it clean, density_norm='width' to equalize visual weights
  1791. v = sns.violinplot(
  1792. x='Community', y='Dice Score', data=df_melted,
  1793. palette=colors, alpha=0.3, inner=None, density_norm='width',
  1794. ax=axes[1], hue='Community', legend=False
  1795. )
  1796. # Clip the violin plot to only show the right half
  1797. for violin in [c for c in axes[1].collections if isinstance(c, plt.matplotlib.collections.PolyCollection)]:
  1798. for path in violin.get_paths():
  1799. m = path.vertices[:, 0].mean()
  1800. path.vertices[:, 0] = np.clip(path.vertices[:, 0], m, np.inf)
  1801. # 2. Draw the "Core" (Narrow Boxplot)
  1802. # Placed on top of the violin to show statistical summary
  1803. sns.boxplot(
  1804. x='Community', y='Dice Score', data=df_melted,
  1805. width=0.15, palette=colors, showfliers=False,
  1806. boxprops={'alpha': 0.8, 'zorder': 10},
  1807. medianprops={'color': 'black', 'linewidth': 2},
  1808. ax=axes[1]
  1809. )
  1810. # 3. Draw the "Rain" (Jittered Stripplot)
  1811. # Offset slightly to the left to avoid overlapping the boxplot/violin
  1812. sns.stripplot(
  1813. x='Community', y='Dice Score', data=df_melted,
  1814. palette=colors, size=3, jitter=0.15, alpha=0.5,
  1815. dodge=False, ax=axes[1], hue='Community', legend=False
  1816. )
  1817. # Adjusting title and labels for Panel B
  1818. axes[1].set_title(f"Panel B: Community-wise Stability (Strategy 2)\n(Dice Distribution across {n_iterations} Iterations)",
  1819. fontsize=12, fontweight='bold', pad=15)
  1820. axes[1].set_ylim(0, 1.05)
  1821. axes[1].set_ylabel("Dice Coefficient", fontsize=12)
  1822. axes[1].set_xlabel("Communities ID", fontsize=12)
  1823. axes[1].grid(axis='y', linestyle='--', alpha=0.4)
  1824. # Final aesthetic touches
  1825. sns.despine(ax=axes[1], left=True)
  1826. plt.tight_layout()
  1827. if output_path:
  1828. plt.savefig(f'{output_path}/validation/sbj-level_result.png', dpi=600, bbox_inches='tight')
  1829. plt.show()
  1830. def plot_stability_validation(data_dict, config, cmap=None, output_path=None):
  1831. sns.set_theme(style="whitegrid")
  1832. mode = config.get('mode', 'subject')
  1833. k = config.get('best_k', 7)
  1834. prefix = config.get('prefix', 'stability')
  1835. n_rows, n_cols = (1, 2) if mode == 'subject' else (2, 2)
  1836. fig, axes = plt.subplots(n_rows, n_cols, figsize=(8*n_cols, 6*n_rows))
  1837. axes = np.atleast_1d(axes).flatten()
  1838. cluster_names = [f'C{i+1}' for i in range(k)]
  1839. colors = [cmap(i) for i in range(k)] if cmap else sns.color_palette("husl", k)
  1840. def draw_raincloud(ax, data, title):
  1841. df_dice = pd.DataFrame(data, columns=cluster_names)
  1842. df_melted = df_dice.melt(var_name='Communities ID', value_name='Dice Score')
  1843. # Cloud: Violin
  1844. sns.violinplot(x='Communities ID', y='Dice Score', data=df_melted, palette=colors,
  1845. alpha=0.3, inner=None, density_norm='width', ax=ax, hue='Communities ID', legend=False)
  1846. for violin in [c for c in ax.collections if isinstance(c, plt.matplotlib.collections.PolyCollection)]:
  1847. for path in violin.get_paths():
  1848. m = path.vertices[:, 0].mean()
  1849. path.vertices[:, 0] = np.clip(path.vertices[:, 0], m, np.inf)
  1850. # Core: Boxplot
  1851. sns.boxplot(x='Communities ID', y='Dice Score', data=df_melted, width=0.15, palette=colors,
  1852. showfliers=False, boxprops={'alpha': 0.8, 'zorder': 10},
  1853. medianprops={'color': 'black', 'linewidth': 2}, ax=ax)
  1854. # Rain: Strip
  1855. sns.stripplot(x='Communities ID', y='Dice Score', data=df_melted, palette=colors,
  1856. size=3, jitter=0.15, alpha=0.5, dodge=False, ax=ax, hue='Communities ID', legend=False)
  1857. ax.set_title(title, fontsize=14, fontweight='bold', pad=15)
  1858. ax.set_ylim(0, 1.05)
  1859. ax.set_xlabel("Communities ID", fontsize=14)
  1860. ax.set_ylabel("Dice Score", fontsize=14)
  1861. sns.despine(ax=ax, left=True)
  1862. # ¶¨Òå»æÍ¼ÄÚ²¿¸¨Öúº¯Êý£ºHeatmap
  1863. def draw_heatmap(ax, cm_data, title, ylabel="Reference Labels"):
  1864. cm_norm = cm_data.astype('float') / (cm_data.sum(axis=1)[:, np.newaxis] + 1e-10)
  1865. sns.heatmap(cm_norm, annot=True, fmt=".2f", cmap="Blues", ax=ax,
  1866. xticklabels=cluster_names, yticklabels=cluster_names,
  1867. square=True, cbar_kws={'label': 'Proportion of Samples'})
  1868. ax.set_title(title, fontsize=14, fontweight='bold', pad=15)
  1869. ax.set_xlabel("Aligned Labels", fontsize=14)
  1870. ax.set_ylabel(ylabel, fontsize=14)
  1871. if mode == 'subject':
  1872. # Panel A: Strategy 1 Alignment
  1873. cm_raw = np.zeros((k, k, 500))
  1874. orig_labels = data_dict["original_labels"]
  1875. for i, aligned in enumerate(data_dict["strat2_metrics"]["all_edge_labels"]):
  1876. from sklearn.metrics import confusion_matrix
  1877. cm_raw[:,:,i] = confusion_matrix(orig_labels, aligned)
  1878. draw_heatmap(axes[0], np.mean(cm_raw, axis=2), "Panel A: Alignment Accuracy\n(Subject Averaging vs. Original)")
  1879. # Panel B: Dice Distribution
  1880. draw_raincloud(axes[1], data_dict["strat2_metrics"]["per_cluster_dice"], f"Panel B: {k} Communities Stability")
  1881. elif mode == 'split-half':
  1882. # Panel A: H1 vs H2
  1883. cm_h1h2 = np.mean(np.array(data_dict['conf_mat_h1h2']), axis=0)
  1884. draw_heatmap(axes[0], cm_h1h2, f"Panel A: Alignment (K={k})\n(Half 1 vs. Aligned Half 2)")
  1885. # Panel B: Dice
  1886. draw_raincloud(axes[1], data_dict['per_cluster_dice'], "Panel B: Dice Distribution")
  1887. # Panel C: Origin vs Half
  1888. cm_orighalf_data = np.concatenate((data_dict['conf_mat_origh1'], data_dict['conf_mat_origh2']))
  1889. draw_heatmap(axes[2], np.mean(cm_orighalf_data, axis=0), f"Panel C: Alignment (K={k})\n(Origin vs. Aligned Half)")
  1890. # Panel D: Repeat Dice or Placeholder
  1891. draw_raincloud(axes[3], data_dict['per_cluster_dice'], "Panel D: Stability Summary")
  1892. # ------------------ ±£´æÂß¼­ ------------------
  1893. plt.tight_layout()
  1894. if output_path:
  1895. os.makedirs(os.path.join(output_path, 'validation'), exist_ok=True)
  1896. fname = f'{mode}_stability_k{k}_{prefix}.png'
  1897. plt.savefig(os.path.join(output_path, 'validation', fname), dpi=600, bbox_inches='tight')
  1898. plt.show()

validation.py at commit 0760eb3, under MIT · at the source

Overview

  1. State Key Laboratory of Cognitive Neuroscience and Learning and IDG/McGovern Institute for Brain Research, Beijing Normal University, Beijing, China
  2. Beijing Key Laboratory of Brain Imaging and Connectomics, Beijing Normal University, Beijing, China
  3. Chinese Institute for Brain Research, Beijing, China
Journal: eLife, volume 15, article RP110294
Dates: published online 24 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.7554/elife.110294 · PMID 42339870 · PMCID PMC13293607 · OpenAlex W7131869714
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: fMRI (modality), human (organism)
Methods: Spectral & time-frequency, Connectivity, Statistics, Machine learning, Preprocessing, Graphs, fMRI & imaging, Physiology & signal measures
Keywords: Human
MeSH: Arousal*, Brain*, Connectome*, Wakefulness*, Adult, Female, Humans, Magnetic Resonance Imaging, Male, Young Adult (* major topic)
Topic: Functional Brain Connectivity Studies (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: STI 2030-Major Projects (Nos. 2021ZD0200500, Nos. 2021ZD0201701); National Natural Science Foundation of China (82021004, Nos. T2325006); Fundamental Research Funds for the Central Universities (No. 2233200020)
Citations: not cited yet (Europe PMC); 65 references in the paper

Abstract

Arousal fluctuates continuously during wakefulness, yet how these moment-to-moment variations shape large-scale functional connectivity (FC) remains unclear. Here, we combined 7T fMRI with concurrent pupillometry to quantify, for every functional connection, how time-varying FC covaries with spontaneous arousal in the awake human brain. Rather than exerting a uniform influence across the connectome, arousal organized FC into a low-dimensional set of seven connectivity communities, each defined by characteristic network compositions. These communities exhibited systematic hemispheric asymmetries, specifically identifying a ‘left-hemisphere centripetal architecture’ where the left hemisphere serves as a structural sink for the asymmetric convergence of arousal-modulated signals. Importantly, hemispheric asymmetry did not arise from global shifts in connectivity strength but instead reflected structured spatial heterogeneity embedded within community architecture. This modular and asymmetric organization was highly preserved during naturalistic movie watching, indicating that arousal-related modulation of FC reflects intrinsic principles that generalize across awake cognitive contexts. Together, these findings demonstrate that moment-to-moment arousal fluctuations shape large-scale FC through structured, hemispherically asymmetric network organization during wakefulness.

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

kongxy6478/Arousal-modulates-functional-connectivity

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 0760eb307e6ef88473b01384346da57d06bd1dee, 17 June 2026
Languages: Python (4)
Size: 6 files, 4 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (4 files), NumPy (4 files), pandas (4 files), SciPy (4 files), seaborn (4 files), scikit-learn (2 files), statsmodels (2 files), neuromaps (1 file), NiBabel (1 file), Pingouin (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
6 files

Code availability

All analyses were implemented in Python (NumPy, SciPy, scikit-learn) using custom scripts. Visualization was performed with Matplotlib, Seaborn, and Surfplot. Computer codes used to calculate the communities, analyse results, and reproduce the figures of the study are openly available at https://github.com/kongxy6478/Arousal-modulates-functional-connectivity (copy archived at Kong, 2026).

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;
  • 4 scripts, each with its path and the digest of its content;
  • 7 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

This study used publicly available data from HCP (https://www.humanconnectome.org/). The processed data and analysis code that support the findings of this study are openly available at https://github.com/kongxy6478/Arousal-modulates-functional-connectivity (copy archived at Kong, 2026).

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, 3 authors, 1 keyword, 10 MeSH terms, 3 funders, 64 references.

Cite

This paper

Kong, X., Li, S., & Gong, G. (2026). Arousal modulates functional connectivity through structured and hemispherically asymmetric community architecture during wakefulness. eLife, 15, RP110294. https://doi.org/10.7554/elife.110294

BibTeX

@article{kong2026arousal,
author = {Kong, Xiangyu and Li, Siyu and Gong, Gaolang},
title = {{Arousal modulates functional connectivity through structured and hemispherically asymmetric community architecture during wakefulness}},
journal = {eLife},
year = {2026},
month = jun,
volume = {15},
pages = {RP110294},
publisher = {eLife Sciences Publications, Ltd},
issn = {2050-084X},
doi = {10.7554/elife.110294},
url = {https://doi.org/10.7554/elife.110294},
pmid = {42339870},
pmcid = {PMC13293607}
}

RIS

TY - JOUR
AU - Kong, Xiangyu
AU - Li, Siyu
AU - Gong, Gaolang
TI - Arousal modulates functional connectivity through structured and hemispherically asymmetric community architecture during wakefulness
T2 - eLife
J2 - Elife
PY - 2026
DA - 2026/06/24
VL - 15
SP - RP110294
SN - 2050-084X
PB - eLife Sciences Publications, Ltd
DO - 10.7554/elife.110294
UR - https://doi.org/10.7554/elife.110294
LA - en
ER -

CSL-JSON

{
"id": "10.7554/elife.110294",
"type": "article-journal",
"title": "Arousal modulates functional connectivity through structured and hemispherically asymmetric community architecture during wakefulness",
"container-title": "eLife",
"author": [
{
"family": "Kong",
"given": "Xiangyu"
},
{
"family": "Li",
"given": "Siyu"
},
{
"family": "Gong",
"given": "Gaolang"
}
],
"container-title-short": "Elife",
"volume": "15",
"page": "RP110294",
"DOI": "10.7554/elife.110294",
"PMID": "42339870",
"PMCID": "PMC13293607",
"ISSN": "2050-084X",
"publisher": "eLife Sciences Publications, Ltd",
"URL": "https://doi.org/10.7554/elife.110294",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
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/s41467-026-75959-w [code]
Charting higher-order models of brain function beyond pairwise interactions.
Journal: Nature communications
In common: neuromaps, NiBabel, statsmodels, 6 other tools, 8 references
[2] doi:10.1038/s41467-026-71270-w [code]
Spatiotemporal dynamics of the human cortical functional hierarchy across the lifespan.
Journal: Nature communications
In common: NiBabel, seaborn, scikit-learn, 4 other tools, fMRI, 5 references, author Gaolang Gong
[3] doi:10.1038/s41467-026-74466-2 [code]
Neuromorphic hierarchical modular reservoirs.
Journal: Nature communications
In common: neuromaps, Pingouin, NiBabel, 7 other tools, 3 references
[4] doi:10.1038/s41398-026-04025-2 [code]
Brain energetic landscapes shape state dysregulation in major depressive disorder: a morphological network controllability perspective.
Journal: Translational psychiatry
In common: neuromaps, Pingouin, NiBabel, 6 other tools, 4 references
[5] doi:10.7554/elife.103097 [code]
Canonical neurodevelopmental trajectories of structural and functional manifolds.
Journal: eLife
In common: neuromaps, Pingouin, NiBabel, 5 other tools, 4 references
[6] doi:10.21203/rs.3.rs-9326213/v1 [code]
Multi-task fMRI outperforms resting-state fMRI for revealing task-invariant organization of the human brain
Journal: Research Square (preprint)
In common: neuromaps, NiBabel, seaborn, 5 other tools, fMRI, 5 references
[7] doi:10.64898/2026.03.09.710558 [code]
Multi-task fMRI outperforms resting-state fMRI for revealing task-invariant organization of the human brain
Journal: bioRxiv (preprint)
In common: neuromaps, NiBabel, seaborn, 5 other tools, fMRI, 5 references
[8] doi:10.1371/journal.pbio.3003684 [code]
The retrieval of previously learned motor memories is facilitated by the reinstatement of default mode network manifold structures.
Journal: PLoS biology
In common: neuromaps, Pingouin, NiBabel, 6 other tools, fMRI, 3 references
[9] doi:10.1162/imag.a.1248 [code]
Estimating fMRI timescale maps.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: neuromaps, NiBabel, statsmodels, 5 other tools, fMRI, 4 references
[10] doi:10.1371/journal.pbio.3003916 [code]
Arousal-driven critical roaming reproduces human functional connectivity dynamics.
Journal: PLoS biology
In common: SciPy, Matplotlib, NumPy, fMRI, 7 references

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.