OSCR

Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains.

Code ↔ Paper

18 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 18 matches
  1. [1] § Methods and materials › Design, optimization, and training of Convolutional Neural Networks › Training procedures. ↔ scripts/07_train_weight_cnn.py, lines 76–111 · score 0.87 · cosine annealing learning, weight decay, Huber loss, rate scheduler, Optimization, dual
  2. [2] § Methods and materials › Design, optimization, and training of Convolutional Neural Networks › Training procedures. ↔ scripts/06_train_connectivity_cnn.py, lines 80–117 · score 0.75 · logits loss, training epoch, training loader, patience, scheduler, optimized
  3. [3] § Methods and materials › Neural simulation framework › Poisson spike train generation and bursting activity. ↔ PointNeuron_Simulation/signals.py, lines 229–257 · score 0.71 · geometric distribution, bursting activity, burst spiking, spike trains, ISI, uniform
  4. [4] § Methods and materials › Neural simulation framework › Poisson spike train generation and bursting activity. ↔ PointNeuron_Simulation/signals.py, lines 229–257 · score 0.68 · longer bursts, shorter bursts, initiations, spike train, smaller, window
  5. [5] § Methods and materials › Design, optimization, and training of Convolutional Neural Networks › Training procedures. ↔ scripts/07_train_weight_cnn.py, lines 76–111 · score 0.64 · training epoch, training loader, patience, scheduler, optimized, batches
  6. [6] § Methods and materials › Design, optimization, and training of Convolutional Neural Networks › Training sample generation. ↔ PointNeuron_Simulation/connectivity.py, lines 31–178 · score 0.63 · connectivity matrix, postsynaptic neuron, presynaptic neuron, class, network, simulation
  7. [7] § Results › Diagnostic analysis of CNN connectivity inference across perturbations ↔ scripts/4_0_indi_confi_fmm_Calu.ipynb, lines 1063–1191 · score 0.62 · focused window, normalized entropy, dip depth, KL divergence, peak height, threshold
  8. [8] § Methods and materials › In vitro HD-MEA/patch-clamp dataset › Spike-train feature extraction. ↔ PointNeuron_Simulation/utils.py, lines 525–593 · score 0.60 · Inter spike interval, spike trains, variation, quantified, ISI, coefficient
  9. [9] § Methods and materials › Neural simulation framework › Network simulation. ↔ PointNeuron_Simulation/neuron.py, lines 95–134 · score 0.58 · added bursting activity, Poisson processes, spike trains, synaptic, neurons, simulated
  10. [10] § Methods and materials › Neural simulation framework › Poisson spike train generation and bursting activity. ↔ PointNeuron_Simulation/neuron.py, lines 95–134 · score 0.57 · absolute refractory period, spike train, bursts, Poisson, window, activity
  11. [11] § Results › Diagnostic analysis of CNN-based synaptic weight inference across perturbations and the effect of pooled training ↔ scripts/4_0_indi_confi_fmm_Calu.ipynb, lines 313–441 · score 0.56 · weightCNN, connCNN, feature map, baseline models, FMM, predicted
  12. [12] § Methods and materials › Neural simulation framework › Leaky integrate-and-fire model. ↔ PointNeuron_Simulation/simulation.py, lines 100–147 · score 0.55 · refractory period, transient, reversal, reset, leak, postsynaptic
  13. [13] § Methods and materials › Model ranking by balanced perturbation performance ↔ scripts/4-1-FMM_Indi_Conn.ipynb, lines 70–183 · score 0.55 · binomial resampling, bootstrap, CI, ranking, perturbation, model
  14. [14] § Methods and materials › Neural simulation framework › Poisson spike train generation and bursting activity. ↔ PointNeuron_Simulation/utils.py, lines 525–593 · score 0.54 · inter spike intervals, spike trains, ISIs, Poisson, burst
  15. [15] § Results › Illustration of the controlled benchmarking of CNN robustness across biophysical perturbations ↔ scripts/4-1-FMM_Indi_Conn.ipynb, lines 1–66 · score 0.53 · BP0.5, dip depth, KL divergence, training baseline, peak height, baseline model
  16. [16] § Methods and materials › Neural simulation framework › Leaky integrate-and-fire model. ↔ PointNeuron_Simulation/neuron.py, lines 20–93 · score 0.53 · membrane capacitance, reversal, reset, leak, fire, neuron
  17. [17] § Methods and materials › CCG indicators ↔ scripts/4_0_indi_confi_fmm_Calu.ipynb, lines 1063–1191 · score 0.50 · focus window, KL divergence, tails, bins, excitatory, dip
  18. [18] § Results › Diagnostic analysis of CNN-based synaptic weight inference across perturbations and the effect of pooled training ↔ scripts/4-1-FMM_Indi_Conn.ipynb, lines 1276–1310 · score 0.50 · XGBoost, Boosting, SHAP, interactions, predict, baseline

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Jupyter notebook · 1,200 lines · 60 KB · no license · 3 matches

  1. # %% [markdown]
  2. # # Info
  3. #
  4. # Author: Xiaoqian Sun </br>
  5. # Created: 07/28/2025 </br>
  6. # Function: Analysis on Signal Level
  7. # - create a vector of indicators to describe CCGs
  8. # - correlate each indicator with perturbation strength
  9. # - regress model performance on these indicators
  10. # - cluster perturbations by CCG types to identify which combinations are most harmful
  11. #
  12. # ### Entropy
  13. # Entropy doesn't care whether bins are positive or negative. It only cares how unevenly values are distributed after normalization to a probability distribution.
  14. # Measures the overall unpredictability of a distribution
  15. # - A flat (uniform) CCG → high entropy (close to log₂(N))
  16. # - A sharp peak or dip → lower entropy (because values are unevenly distributed)
  17. #
  18. # ### KL Divergence
  19. # KL divergence highlights how far the actual shape is from "flat randomness." Unlike entropy, KL requires defining a reference distribution (here, uniform).
  20. # - If bins are all equal → KL divergence is 0
  21. # - If a few bins dominate (e.g., a peak or dip near 0 ms) → KL divergence increases
  22. #
  23. # ### Intuition
  24. # - Entropy is absolute — it measures how spread out the probability mass is.
  25. # - KL divergence is relative — it measures how far the current distribution is from a baseline (like uniform).
  26. # %%
  27. import os
  28. import numpy as np
  29. import pandas as pd
  30. import seaborn as sns
  31. from Functions import *
  32. import matplotlib.pyplot as plt
  33. from scipy.stats import entropy
  34. from scipy.special import rel_entr
  35. from scipy.ndimage import gaussian_filter1d
  36. import warnings
  37. warnings.filterwarnings('ignore')
  38. # %% [markdown]
  39. # # Path & Common File & Parameters
  40. # %%
  41. # path --------------------------------------------------
  42. root_path = r'/Users/xiaoqiansun/Desktop/InferConnectivity'
  43. work_path = os.getcwd()
  44. dataloader_path = os.path.join(work_path, 'dataloader')
  45. result_path = os.path.join(work_path, 'Result')
  46. allRuns_path = os.path.join(result_path, 'all120Runs')
  47. resultAnalysis_path = os.path.join(work_path, 'resultAnalysis')
  48. indiSave_path = os.path.join(resultAnalysis_path, 'indi'); os.makedirs(indiSave_path, exist_ok=True)
  49. confiSave_path = os.path.join(resultAnalysis_path, 'fmm_confi'); os.makedirs(confiSave_path, exist_ok=True)
  50. data_path = os.path.join(root_path, 'Dataset', 'm2post')
  51. info_df = pd.read_csv(os.path.join(dataloader_path, 'binInfo.csv')).iloc[:, 1:]
  52. bin_size=1; timebins = info_df['t']
  53. # Para --------------------------------------------------
  54. # baseline_types = ['E3', 'E5', 'E8', 'g0.5', 'g1.5']
  55. baseline_types = ['E3', 'E5', 'E8', 'g0.5', 'g1.5', 'BP0.5', 'mu165', 'mu180', 'tau5', 'tau40']
  56. perturbations = ['mu', 'tau', 'BP', 'EI', 'Weight']
  57. indicators_exc = ['peak_height', 'peak_lag', 'peak_halfMax_width', 'peak_to_noise', 'temporal_span', 'entropy', 'norm_entropy', 'norm_entropy_window', 'kl_divergence_window']
  58. indicators_inh = ['dip_depth', 'dip_lag', 'dip_halfMax_width', 'dip_to_noise', 'temporal_span', 'entropy', 'norm_entropy', 'norm_entropy_window', 'kl_divergence_window']
  59. cmap= 'coolwarm'; cm_cmap='PuBu';
  60. mC='firebrick'; bC='steelblue'; ccgColor='k'
  61. coolwarm_blue = '#3B4CC0'; coolwarm_red = '#B40426'
  62. coolwarm_softblue = '#6BAED6'; coolwarm_softred = '#E07A7A’; correctGREEN= '#2E7D32'
  63. # %% [markdown]
  64. # # Indicator - Just Ran Once
  65. # %%
  66. # bs_Idx = 1
  67. # baselineType = baseline_types[bs_Idx]
  68. # baselineName = 'baseline'#'baseline_E3'
  69. # base_mu = 150
  70. # base_tau = 15
  71. # base_BP = 1.0
  72. # base_EI = 0.5
  73. # base_g = 1.0
  74. # dataPrefix = '_'+baselineType
  75. # bs_featureDic = [{'mu':base_mu, 'tau':base_tau, 'BP':base_BP, 'EI':base_EI, 'g': base_g}]
  76. # print(' - baseline', dataPrefix, 'with base values:', bs_featureDic)
  77. # %%
  78. # mus = [150+2*i for i in range(16)]
  79. # mu_prefix = ['mu'+str(mu) for mu in mus]
  80. # mu_tests = ['mu'+dataPrefix+str(mu) for mu in mus] if ('mu' not in dataPrefix and 'E5' not in dataPrefix) else mu_prefix
  81. # mu_featureDics = [{'mu':mu, 'tau':base_tau, 'BP':base_BP, 'EI':base_EI, 'g':base_g } for mu in mus]
  82. # taus = [5*i for i in range(1, 11)]
  83. # tau_prefix = ['tau'+str(tau) for tau in taus]
  84. # tau_tests = ['tau'+dataPrefix+str(tau) for tau in taus] if ('tau' not in dataPrefix and 'E5' not in dataPrefix) else tau_prefix
  85. # tau_featureDics = [{'mu':base_mu, 'tau':tau, 'BP':base_BP, 'EI':base_EI, 'g':base_g } for tau in taus]
  86. # BPs = [0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
  87. # BP_prefix = ['BP'+str(BP) for BP in BPs]
  88. # BP_tests = ['BP'+dataPrefix+str(BP) for BP in BPs] if ('BP' not in dataPrefix and 'E5' not in dataPrefix) else BP_prefix
  89. # BP_featureDics = [{'mu':base_mu, 'tau':base_tau, 'BP':BP, 'EI':base_EI, 'g':base_g } for BP in BPs]
  90. # EIs = [0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]
  91. # EI_prefix = ['EI'+str(EI) for EI in EIs]
  92. # EI_tests = ['EI'+dataPrefix+str(EI) for EI in EIs] if ('E' not in dataPrefix and 'E5' not in dataPrefix) else EI_prefix
  93. # EI_featureDics = [{'mu':base_mu, 'tau':base_tau, 'BP':base_BP, 'EI':EI, 'g':base_g } for EI in EIs]
  94. # gs = [0.4, 0.6, 0.8, 1.0, 1.2, 1.4, 1.6, 1.8, 2.0]
  95. # gbar_prefix = ['g'+str(g) for g in gs]
  96. # gbar_tests = ['g'+dataPrefix+str(g) for g in gs] if ('g' not in dataPrefix and 'E5' not in dataPrefix) else gbar_prefix
  97. # g_featureDics = [{'mu':base_mu, 'tau':base_tau, 'BP':base_BP, 'EI':base_EI, 'g': g} for g in gs]
  98. # file_test_list = [baselineName] + mu_tests + tau_tests + BP_tests + EI_tests + gbar_tests
  99. # file_prefix_list = ['baseline'] + mu_prefix + tau_prefix + BP_prefix + EI_prefix + gbar_prefix
  100. # file_featureDics = bs_featureDic + mu_featureDics + tau_featureDics + BP_featureDics + EI_featureDics + g_featureDics
  101. # print(len(file_prefix_list), '= 1 baseline + 48 perturbations')
  102. # %%
  103. # indicator_exc_list = []; indicator_inh_list = []
  104. # for i in range(len(file_test_list)):
  105. # fileTest = file_test_list[i]
  106. # featureDic = file_featureDics[i]
  107. # try:
  108. # ccg, ws, ls, ss = read_ccg(data_path, fileTest+'_CCG.csv')
  109. # for nIdx in range(len(ccg)):
  110. # w = ws[nIdx]; sname = ss[nIdx]
  111. # if w > 0:
  112. # indi = ccg_indicators(ccg[nIdx], connection_type='exc', smoothSigma=0.5, bin_size=1, peak_window_ms=10, ifVerbose=False,ifPlot=False,ifSave=False)
  113. # indi.update(featureDic); indi['Weight'] = w; indi['sample']=sname; indi['test']=fileTest
  114. # indicator_exc_list.append(indi)
  115. # elif w < 0:
  116. # indi = ccg_indicators(ccg[nIdx], connection_type='inh', smoothSigma=0.5, bin_size=1, peak_window_ms=10, ifVerbose=False,ifPlot=False,ifSave=False)
  117. # indi.update(featureDic); indi['Weight'] = w; indi['sample']=sname; indi['test']=fileTest
  118. # indicator_inh_list.append(indi)
  119. # except Exception as e:
  120. # print(fileTest, '--', e)
  121. # break
  122. # indicatorExc_df=pd.DataFrame(indicator_exc_list);# indicatorExc_df.to_csv(os.path.join(indiSave_path, baselineType+'_indicator_exc.csv'))
  123. # indicatorInh_df=pd.DataFrame(indicator_inh_list);# indicatorInh_df.to_csv(os.path.join(indiSave_path, baselineType+'_indicator_inh.csv'))
  124. # print(indicatorExc_df.shape, indicatorInh_df.shape); indicatorInh_df.tail(5)
  125. # %% [markdown]
  126. # # ----------------------------------------------------------------------
  127. #
  128. # # FMM, Confi, Indi
  129. # %% [markdown]
  130. # ### loader best M
  131. # %%
  132. baseline_types = ['E3', 'E5', 'E8', 'g0.5', 'g1.5', 'BP0.5', 'mu165', 'mu180', 'tau5', 'tau40']
  133. model, optimizer, scheduler, criterion = init_conn_model()
  134. connCNNE3,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_E3_100.0.pth')
  135. model, optimizer, scheduler, criterion = init_conn_model()
  136. connCNNE5,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_E5_100.0.pth')
  137. model, optimizer, scheduler, criterion = init_conn_model()
  138. connCNNE8,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_E8_100.0.pth')
  139. model, optimizer, scheduler, criterion = init_conn_model()
  140. connCNNg05,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_g0.5_100.0.pth')
  141. model, optimizer, scheduler, criterion = init_conn_model()
  142. connCNNg15,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_g1.5_100.0.pth')
  143. model, optimizer, scheduler, criterion = init_conn_model()
  144. connCNNBP05,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_BP0.5_100.0.pth')
  145. model, optimizer, scheduler, criterion = init_conn_model()
  146. connCNNmu165,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_mu165_100.0.pth')
  147. model, optimizer, scheduler, criterion = init_conn_model()
  148. connCNNmu180,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_mu180_100.0.pth')
  149. model, optimizer, scheduler, criterion = init_conn_model()
  150. connCNNtau5,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_tau5_100.0.pth')
  151. model, optimizer, scheduler, criterion = init_conn_model()
  152. connCNNtau40,_ = load_conn_model(model,optimizer,scheduler,savePath=allRuns_path, filename='ConnCNN_tau40_100.0.pth')
  153. baseline_connCNN_models = [connCNNE3, connCNNE5, connCNNE8, connCNNg05, connCNNg15,
  154. connCNNBP05, connCNNmu165, connCNNmu180, connCNNtau5, connCNNtau40]
  155. # %%
  156. model, optimizer, scheduler, criterion = init_weight_model()
  157. weighCNNE3,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_E3_0.0096.pth')
  158. model, optimizer, scheduler, criterion = init_weight_model()
  159. weighCNNE5,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_E5_0.0016.pth')
  160. model, optimizer, scheduler, criterion = init_weight_model()
  161. weighCNNE8,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_E8_0.0003.pth')
  162. model, optimizer, scheduler, criterion = init_weight_model()
  163. weighCNNg05,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_g0.5_0.0001.pth')
  164. model, optimizer, scheduler, criterion = init_weight_model()
  165. weighCNNg15,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_g1.5_0.0024.pth')
  166. model, optimizer, scheduler, criterion = init_weight_model()
  167. weighCNNBP05,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_BP0.5_0.0001.pth')
  168. model, optimizer, scheduler, criterion = init_weight_model()
  169. weighCNNmu165,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_mu165_0.0033.pth')
  170. model, optimizer, scheduler, criterion = init_weight_model()
  171. weighCNNmu180,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_mu180_0.0005.pth')
  172. model, optimizer, scheduler, criterion = init_weight_model()
  173. weighCNNtau5,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_tau5_0.0001.pth')
  174. model, optimizer, scheduler, criterion = init_weight_model()
  175. weighCNNtau40,_ = load_model(model,optimizer,scheduler,savePath=allRuns_path, filename='w_tau40_0.0002.pth')
  176. baseline_weightCNN_models = [weighCNNE3, weighCNNE5, weighCNNE8, weighCNNg05, weighCNNg15,
  177. weighCNNBP05, weighCNNmu165, weighCNNmu180, weighCNNtau5, weighCNNtau40]
  178. # %% [markdown]
  179. # ### perbs list
  180. # %%
  181. mu_prefix = ['mu'+str(150+2*i) for i in range(16)]
  182. muE3_prefix = ['mu_E3'+str(150+2*i) for i in range(16)]
  183. muE8_prefix = ['mu_E8'+str(150+2*i) for i in range(16)]
  184. mug05_prefix = ['mu_g0.5'+str(150+2*i) for i in range(16)]
  185. mug15_prefix = ['mu_g1.5'+str(150+2*i) for i in range(16)]
  186. muBP05_prefix = ['mu_BP0.5'+str(150+2*i) for i in range(16)]
  187. mutau5_prefix = ['mu_tau5'+str(150+2*i) for i in range(16)]
  188. mutau40_prefix = ['mu_tau40'+str(150+2*i) for i in range(16)]
  189. tau_prefix = ['tau'+str(5*i) for i in range(1, 11)]
  190. tauE3_prefix = ['tau_E3'+str(5*i) for i in range(1, 11)]
  191. tauE8_prefix = ['tau_E8'+str(5*i) for i in range(1, 11)]
  192. taug05_prefix = ['tau_g0.5'+str(5*i) for i in range(1, 11)]
  193. taug15_prefix = ['tau_g1.5'+str(5*i) for i in range(1, 11)]
  194. tauBP05_prefix = ['tau_BP0.5'+str(5*i) for i in range(1, 11)]
  195. taumu165_prefix = ['tau_mu165'+str(5*i) for i in range(1, 11)]
  196. taumu180_prefix = ['tau_mu180'+str(5*i) for i in range(1, 11)]
  197. BPs = [0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
  198. BP_prefix = ['BP'+str(bp) for bp in BPs]
  199. BPE3_prefix = ['BP_E3'+str(bp) for bp in BPs]
  200. BPE8_prefix = ['BP_E8'+str(bp) for bp in BPs]
  201. BPg05_prefix = ['BP_g0.5'+str(bp) for bp in BPs]
  202. BPg15_prefix = ['BP_g1.5'+str(bp) for bp in BPs]
  203. BPmu165_prefix = ['BP_mu165'+str(bp) for bp in BPs]
  204. BPmu180_prefix = ['BP_mu180'+str(bp) for bp in BPs]
  205. BPtau5_prefix = ['BP_tau5'+str(bp) for bp in BPs]
  206. BPtau40_prefix = ['BP_tau40'+str(bp) for bp in BPs]
  207. EIatio=[0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]
  208. EI_prefix = ['EI'+str(ei) for ei in EIatio]
  209. EIg05_prefix = ['EI_g0.5'+str(ei) for ei in EIatio] ; EIg05_prefix.remove('EI_g0.50.3')
  210. EIg15_prefix = ['EI_g1.5'+str(ei) for ei in EIatio]
  211. EIBP05_prefix = ['EI_BP0.5'+str(ei) for ei in EIatio]
  212. EImu165_prefix = ['EI_mu165'+str(ei) for ei in EIatio]
  213. EImu180_prefix = ['EI_mu180'+str(ei) for ei in EIatio]
  214. EItau5_prefix = ['EI_tau5'+str(ei) for ei in EIatio]
  215. EItau40_prefix = ['EI_tau40'+str(ei) for ei in EIatio]
  216. gbars = [0.4, 0.6, 0.8, 1.0, 1.2, 1.4, 1.6, 1.8, 2.0]
  217. gbar_prefix = ['g'+str(g) for g in gbars]
  218. gbarE3_prefix = ['g_E3'+str(g) for g in gbars] ; gbarE3_prefix.remove('g_E30.4')
  219. gbarE8_prefix = ['g_E8'+str(g) for g in gbars]
  220. gbarBP05_prefix = ['g_BP0.5'+str(g) for g in gbars]; gbarBP05_prefix.remove('g_BP0.50.4')
  221. gbarmu165_prefix = ['g_mu165'+str(g) for g in gbars]; gbarmu165_prefix.remove('g_mu1650.4')
  222. gbarmu180_prefix = ['g_mu180'+str(g) for g in gbars]; gbarmu180_prefix.remove('g_mu1800.4')
  223. gbartau5_prefix = ['g_tau5'+str(g) for g in gbars]; gbartau5_prefix.remove('g_tau50.4')
  224. gbartau40_prefix = ['g_tau40'+str(g) for g in gbars]; gbartau40_prefix.remove('g_tau400.4')
  225. file_prefix_E5_list = mu_prefix + tau_prefix + BP_prefix + EI_prefix + gbar_prefix
  226. file_prefix_E3_list = muE3_prefix + tauE3_prefix + BPE3_prefix + EI_prefix + gbarE3_prefix
  227. file_prefix_E8_list = muE8_prefix + tauE8_prefix + BPE8_prefix + EI_prefix + gbarE8_prefix
  228. file_prefix_g05_list = mug05_prefix + taug05_prefix + BPg05_prefix + EIg05_prefix + gbar_prefix
  229. file_prefix_g15_list = mug15_prefix + taug15_prefix + BPg15_prefix + EIg15_prefix + gbar_prefix
  230. file_prefix_BP05_list = muBP05_prefix + tauBP05_prefix + BP_prefix + EIBP05_prefix + gbarBP05_prefix
  231. file_prefix_mu165_list = mu_prefix + taumu165_prefix + BPmu165_prefix + EImu165_prefix + gbarmu165_prefix
  232. file_prefix_mu180_list = mu_prefix + taumu180_prefix + BPmu180_prefix + EImu180_prefix + gbarmu180_prefix
  233. file_prefix_tau5_list = mutau5_prefix + tau_prefix + BPtau5_prefix + EItau5_prefix + gbartau5_prefix
  234. file_prefix_tau40_list = mutau40_prefix + tau_prefix + BPtau40_prefix + EItau40_prefix + gbartau40_prefix
  235. file_prefix_lists = [file_prefix_E3_list, file_prefix_E5_list, file_prefix_E8_list, file_prefix_g05_list, file_prefix_g15_list,
  236. file_prefix_BP05_list, file_prefix_mu165_list, file_prefix_mu180_list, file_prefix_tau5_list, file_prefix_tau40_list]
  237. perturbation_loader_lists = [ [load_oneLoader(dataloader_path, file_prefix+'_loader.pkl') for file_prefix in file_prefix_list] for file_prefix_list in file_prefix_lists ]
  238. # %% [markdown]
  239. # ### calculation
  240. # %%
  241. criterionC = nn.BCEWithLogitsLoss()
  242. criterionW = HuberLossWithWeight()
  243. fm_mean_cols = ['conv1_fm0_mean', 'conv1_fm1_mean', 'conv1_fm2_mean', 'conv1_fm3_mean', 'conv2_fm0_mean', 'conv2_fm1_mean']
  244. for bIdx in range(len(baseline_types)):
  245. predSummary_list = []
  246. # retrive the best baseline model ---------------------------------------------------------------
  247. bsType = baseline_types[bIdx]
  248. connCNN_Model = baseline_connCNN_models[bIdx]
  249. weightCNN_Model = baseline_weightCNN_models[bIdx]
  250. # predict on baseline-test set -------------------------------------------------------------------
  251. train_loader, val_loader, test_loader = load_dataloaders(dataloader_path, 'baseline_'+bsType+'_loader.pkl')
  252. # pred with connCNN ------------
  253. avg_connect_loss, accu, pred_list,confi_list, label_list, w_list, sCInfo_list = evaluateConnCNN_wConfi_model(connCNN_Model, test_loader, criterionC)
  254. # pred with weightCNN ----------
  255. avg_wegiht_loss, predW_list, gtW_list, label0_loss, label1_loss, sWInfo_list = evaluateWeightCNN_model(weightCNN_Model, test_loader, criterionW)
  256. # feature map mean -------------
  257. conn_fmMean_list = []; weight_fmMean_list = []
  258. conn_fmMean_pre_list = []; weight_fmMean_pre_list = []
  259. conn_fmMean_postBN_list = []; weight_fmMean_postBN_list = []
  260. for ccgs, cs, ws, ss in test_loader:
  261. for m in range(len(ccgs)):
  262. # pre activation (raw featureMaps)
  263. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(connCNN_Model, ccgs[m]) # connCNN
  264. conn_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  265. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(weightCNN_Model, ccgs[m])# weightCNN
  266. weight_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  267. # postBN (featureMaps after BatchNorm, but before tanh)
  268. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(connCNN_Model, ccgs[m]) # connCNN
  269. conn_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  270. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(weightCNN_Model, ccgs[m])# weightCNN
  271. weight_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  272. # post activation (featureMaps after BatchNorm & tanh)
  273. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(connCNN_Model, ccgs[m]) # connCNN
  274. conn_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  275. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(weightCNN_Model, ccgs[m])# weightCNN
  276. weight_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  277. predSummary = pd.DataFrame([label_list, pred_list, confi_list, gtW_list, predW_list, sCInfo_list, ['baseline_'+bsType]*len(sCInfo_list)],
  278. index=['connLabel', 'predLabel', 'confi', 'gtWeights', 'predWeights', 'sample', 'test']).T
  279. connPre_FMM = pd.DataFrame(np.array(conn_fmMean_pre_list), columns=['pre_conn_'+col for col in fm_mean_cols])
  280. connPost_FMM = pd.DataFrame(np.array(conn_fmMean_list), columns=['post_conn_'+col for col in fm_mean_cols])
  281. connPostBN_FMM = pd.DataFrame(np.array(conn_fmMean_postBN_list), columns=['postBN_conn_'+col for col in fm_mean_cols])
  282. weightPre_FMM = pd.DataFrame(np.array(weight_fmMean_pre_list), columns=['pre_weight_'+col for col in fm_mean_cols])
  283. weightPost_FMM = pd.DataFrame(np.array(weight_fmMean_list), columns=['post_weight_'+col for col in fm_mean_cols])
  284. weightPostBN_FMM = pd.DataFrame(np.array(weight_fmMean_postBN_list), columns=['postBN_weight_'+col for col in fm_mean_cols])
  285. predSummary = pd.concat([predSummary, connPre_FMM, connPost_FMM, connPostBN_FMM, weightPre_FMM, weightPost_FMM, weightPostBN_FMM], axis=1); predSummary_list.append(predSummary)
  286. # predict on associated perturbations ------------------------------------------------------------
  287. file_prefix_list = file_prefix_lists[bIdx]
  288. perturbation_loader = perturbation_loader_lists[bIdx]
  289. for fileIdx in range(len(file_prefix_list)):
  290. perb_type = file_prefix_list[fileIdx]
  291. perb_loader = perturbation_loader[fileIdx]
  292. # pred with connCNN ------------
  293. avg_connect_loss, accu, pred_list,confi_list, label_list, w_list, sCInfo_list = evaluateConnCNN_wConfi_model(connCNN_Model, perb_loader, criterionC)
  294. # pred with weightCNN ----------
  295. avg_wegiht_loss, predW_list, gtW_list, label0_loss, label1_loss, sWInfo_list = evaluateWeightCNN_model(weightCNN_Model, perb_loader, criterionW)
  296. assert np.array_equal(sCInfo_list, sWInfo_list), bsType+' '+perb_type+" sInfo Arrays are not equal"
  297. # get feature map mean on each ccg -------------
  298. conn_fmMean_list = []; weight_fmMean_list = []
  299. conn_fmMean_pre_list = []; weight_fmMean_pre_list = []
  300. conn_fmMean_postBN_list = []; weight_fmMean_postBN_list = []
  301. for ccgs, cs, ws, ss in perb_loader:
  302. for m in range(len(ccgs)):
  303. # pre activation (raw featureMaps)
  304. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(connCNN_Model, ccgs[m]) # connCNN
  305. conn_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  306. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(weightCNN_Model, ccgs[m])# weightCNN
  307. weight_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  308. # postBN (featureMaps after BatchNorm, but before tanh)
  309. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(connCNN_Model, ccgs[m]) # connCNN
  310. conn_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  311. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(weightCNN_Model, ccgs[m])# weightCNN
  312. weight_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  313. # post activation (featureMaps after BatchNorm & tanh)
  314. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(connCNN_Model, ccgs[m])
  315. conn_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  316. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(weightCNN_Model, ccgs[m])
  317. weight_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  318. predSummary = pd.DataFrame([label_list, pred_list, confi_list, gtW_list, predW_list, sCInfo_list, [perb_type]*len(sCInfo_list)],
  319. index=['connLabel', 'predLabel', 'confi', 'gtWeights', 'predWeights', 'sample', 'test']).T
  320. connPre_FMM = pd.DataFrame(np.array(conn_fmMean_pre_list), columns=['pre_conn_'+col for col in fm_mean_cols])
  321. connPost_FMM = pd.DataFrame(np.array(conn_fmMean_list), columns=['post_conn_'+col for col in fm_mean_cols])
  322. connPostBN_FMM = pd.DataFrame(np.array(conn_fmMean_postBN_list), columns=['postBN_conn_'+col for col in fm_mean_cols])
  323. weightPre_FMM = pd.DataFrame(np.array(weight_fmMean_pre_list), columns=['pre_weight_'+col for col in fm_mean_cols])
  324. weightPost_FMM = pd.DataFrame(np.array(weight_fmMean_list), columns=['post_weight_'+col for col in fm_mean_cols])
  325. weightPostBN_FMM = pd.DataFrame(np.array(weight_fmMean_postBN_list), columns=['postBN_weight_'+col for col in fm_mean_cols])
  326. predSummary = pd.concat([predSummary, connPre_FMM, connPost_FMM, connPostBN_FMM, weightPre_FMM, weightPost_FMM, weightPostBN_FMM], axis=1); predSummary_list.append(predSummary)
  327. bspredSummary = pd.concat(predSummary_list)
  328. bspredSummary.to_csv(os.path.join(confiSave_path, bsType+'_confi_w_perbs.csv'))
  329. # baseline_types = ['E3', 'E5', 'E8', 'g0.5', 'g1.5', 'BP0.5', 'mu165', 'mu180', 'tau5', 'tau40']
  330. print('Done')
  331. bspredSummary.tail(2)
  332. # %% [markdown]
  333. # ### combine FMM, Confi, Indictors
  334. # %%
  335. # baseline_types = ['E3', 'E5', 'E8', 'g0.5', 'g1.5', 'BP0.5', 'mu165', 'mu180', 'tau5', 'tau40']
  336. for bIdx in range(len(baseline_types)):
  337. bsType = baseline_types[bIdx]
  338. bspredSummary = pd.read_csv(os.path.join(confiSave_path, bsType+'_confi_w_perbs.csv')).iloc[:, 1:]
  339. exc_indicator_df = pd.read_csv(os.path.join(indiSave_path, bsType+'_indicator_exc.csv')).iloc[:, 1:]
  340. inh_indicator_df = pd.read_csv(os.path.join(indiSave_path, bsType+'_indicator_inh.csv')).iloc[:, 1:]
  341. # merge
  342. pred_indi_exc = pd.merge(bspredSummary, exc_indicator_df, on='sample', how='inner')
  343. pred_indi_exc = pred_indi_exc.drop(columns=['test_y']); pred_indi_exc = pred_indi_exc.rename(columns={'test_x': 'test'})
  344. pred_indi_exc.to_csv(os.path.join(resultAnalysis_path, bsType+'_confi_fmm_indi_exc.csv'))
  345. pred_indi_inh = pd.merge(bspredSummary, inh_indicator_df, on='sample', how='inner')
  346. pred_indi_inh = pred_indi_inh.drop(columns=['test_y']); pred_indi_inh = pred_indi_inh.rename(columns={'test_x': 'test'})
  347. pred_indi_inh.to_csv(os.path.join(resultAnalysis_path, bsType+'_confi_fmm_indi_inh.csv'))
  348. print('Done')
  349. print(pred_indi_exc.shape); pred_indi_exc.head(2)
  350. # %% [markdown]
  351. # ## Best Model Accu/Loss
  352. # %%
  353. import re
  354. def clean_test_value(val):
  355. # Remove patterns like _g0.5 or _BP0.5
  356. return re.sub(r'_g0\.5|_BP0\.5|_g1\.5|_mu165|_mu180|_tau5|_tau40|_E3|_E8', '', val)
  357. mus = [150+2*i for i in range(16)]; mu_prefix = ['mu'+str(mu) for mu in mus]
  358. taus = [5*i for i in range(1, 11)]; tau_prefix = ['tau'+str(tau) for tau in taus]
  359. BPs = [0.5, 0.6, 0.7, 0.8, 0.9, 1.0]; BP_prefix = ['BP'+str(BP) for BP in BPs]
  360. EIs = [0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]; EI_prefix = ['EI'+str(EI) for EI in EIs]
  361. gs = [0.4, 0.6, 0.8, 1.0, 1.2, 1.4, 1.6, 1.8, 2.0]; gbar_prefix = ['g'+str(g) for g in gs]
  362. perturbs_list = mu_prefix + tau_prefix + BP_prefix + EI_prefix + gbar_prefix
  363. model_list = ['mu150', 'mu165', 'mu180', 'tau5', 'tau15', 'tau40', 'BP0.5', 'BP1.0', 'E3', 'E5', 'E8', 'g0.5', 'g1.0', 'g1.5']
  364. # %%
  365. accu_list = []; loss_list = []
  366. for bIdx in range(len(baseline_types)):
  367. bsType = baseline_types[bIdx]
  368. bspredSummary = pd.read_csv(os.path.join(confiSave_path, bsType+'_confi_w_perbs.csv')).iloc[:, 1:]
  369. bspredSummary['hit'] = (bspredSummary['connLabel']==bspredSummary['predLabel']).astype(int)
  370. bspredSummary['mse'] = (bspredSummary['gtWeights'] - bspredSummary['predWeights']) ** 2
  371. bspredSummary['test'] = bspredSummary['test'].apply(clean_test_value)
  372. mean_accu = bspredSummary.groupby('test')['hit'].mean()*100
  373. mean_accu = mean_accu[~mean_accu.index.str.contains('baseline')]
  374. mean_accu = mean_accu.reindex(perturbs_list)
  375. accu_list.append(mean_accu)
  376. mean_loss = bspredSummary.groupby('test')['mse'].mean()
  377. mean_loss = mean_loss[~mean_loss.index.str.contains('baseline')]
  378. mean_loss = mean_loss.reindex(perturbs_list)
  379. loss_list.append(mean_loss)
  380. bestM_accu = pd.concat(accu_list, axis=1); bestM_accu.columns = baseline_types; bestM_accu = bestM_accu.T
  381. bestM_accu.loc['mu150'] = bestM_accu.loc['E5']; bestM_accu.loc['tau15'] = bestM_accu.loc['E5'];
  382. bestM_accu.loc['BP1.0'] = bestM_accu.loc['E5']; bestM_accu.loc['g1.0'] = bestM_accu.loc['E5'];
  383. bestM_accu.loc['g1.5'] = bestM_accu.loc['E5']; bestM_accu = bestM_accu.loc[model_list]
  384. bestM_accu.to_csv(os.path.join(allRuns_path, 'bestM_heatmap','bestM_accu.csv'))
  385. bestM_loss = pd.concat(loss_list, axis=1); bestM_loss.columns = baseline_types; bestM_loss = bestM_loss.T
  386. bestM_loss.loc['mu150'] = bestM_loss.loc['E5']; bestM_loss.loc['tau15'] = bestM_loss.loc['E5'];
  387. bestM_loss.loc['BP1.0'] = bestM_loss.loc['E5']; bestM_loss.loc['g1.0'] = bestM_loss.loc['E5'];
  388. bestM_loss.loc['g1.5'] = bestM_loss.loc['E5']; bestM_loss = bestM_loss.loc[model_list]
  389. bestM_loss.to_csv(os.path.join(allRuns_path, 'bestM_heatmap','bestM_loss.csv'))
  390. # %%
  391. cmap_accu = 'coolwarm_r'; cmap_loss = 'coolwarm';
  392. plt.figure(figsize=(12, 8))
  393. sns.heatmap(bestM_accu, annot=False, cmap=cmap_accu, cbar_kws={'label': 'Accuracy (%)'}, linecolor='lightgrey', linewidths=0.5 )
  394. plt.xlabel("Perturbation", fontweight='bold'); plt.ylabel("Best CNN", fontweight='bold')
  395. plt.tight_layout(); plt.savefig(os.path.join(allRuns_path, 'bestM_heatmap', 'accu.png'), bbox_inches='tight', pad_inches=0.1); plt.close()
  396. fig, ax = plt.subplots(figsize=(12, 8))
  397. sns.heatmap(bestM_loss, annot=False, cmap=cmap_loss, cbar_kws={'label': 'Loss'}, linecolor='lightgrey', linewidths=0.5, vmax=0.01 )
  398. plt.xlabel("Perturbation", fontweight='bold'); plt.ylabel("Best CNN", fontweight='bold')
  399. plt.tight_layout(); plt.savefig(os.path.join(allRuns_path, 'bestM_heatmap', 'loss.png'), bbox_inches='tight', pad_inches=0.1); plt.close()
  400. # %% [markdown]
  401. # # ----------------------------------------------------------------------
  402. # # FMM, Confi, Indi, allBaselines
  403. # %% [markdown]
  404. # ### load model
  405. # %%
  406. model, optimizer, scheduler, criterion = init_conn_model()
  407. connCNN,_ = load_conn_model(model,optimizer,scheduler,savePath=os.path.join(result_path, 'allBaselineTrain_Testing'), filename='c_100.0.pth')
  408. model, optimizer, scheduler, criterion = init_weight_model()
  409. weighCNN,_ = load_model(model,optimizer,scheduler,savePath=os.path.join(result_path, 'allBaselineTrain_Testing'), filename='w_0.006.pth')
  410. # %% [markdown]
  411. # ### perb list
  412. # %%
  413. mus = [150+2*i for i in range(16)]
  414. mu_prefix = ['mu'+str(mu) for mu in mus]; mu_featureDic = [{'mu':mu, 'tau':15, 'BP':1, 'EI':0.5, 'g':1} for mu in mus]
  415. muE3_prefix = ['mu_E3'+str(mu) for mu in mus]; muE3_featureDic = [{'mu':mu, 'tau':15, 'BP':1, 'EI':0.3, 'g':1} for mu in mus]
  416. muE8_prefix = ['mu_E8'+str(mu) for mu in mus]; muE8_featureDic = [{'mu':mu, 'tau':15, 'BP':1, 'EI':0.8, 'g':1} for mu in mus]
  417. mug05_prefix = ['mu_g0.5'+str(mu) for mu in mus]; mug05_featureDic = [{'mu':mu, 'tau':15, 'BP':1, 'EI':0.5, 'g':0.5} for mu in mus]
  418. mug15_prefix = ['mu_g1.5'+str(mu) for mu in mus]; mug15_featureDic = [{'mu':mu, 'tau':15, 'BP':1, 'EI':0.5, 'g':1.5} for mu in mus]
  419. muBP05_prefix = ['mu_BP0.5'+str(mu) for mu in mus]; muBP05_featureDic = [{'mu':mu, 'tau':15, 'BP':0.5, 'EI':0.5, 'g':1} for mu in mus]
  420. mutau5_prefix = ['mu_tau5'+str(mu) for mu in mus]; mutau5_featureDic = [{'mu':mu, 'tau':5, 'BP':1, 'EI':0.5, 'g':1} for mu in mus]
  421. mutau40_prefix = ['mu_tau40'+str(mu) for mu in mus]; mutau40_featureDic = [{'mu':mu, 'tau':40, 'BP':1, 'EI':0.5, 'g':1} for mu in mus]
  422. taus = [5*i for i in range(1, 11)]
  423. tau_prefix = ['tau'+str(taus) for taus in taus]; tau_featureDic = [{'mu':150, 'tau':tau, 'BP':1, 'EI':0.5, 'g':1} for tau in taus]
  424. tauE3_prefix = ['tau_E3'+str(taus) for taus in taus]; tauE3_featureDic = [{'mu':150, 'tau':tau, 'BP':1, 'EI':0.3, 'g':1} for tau in taus]
  425. tauE8_prefix = ['tau_E8'+str(taus) for taus in taus]; tauE8_featureDic = [{'mu':150, 'tau':tau, 'BP':1, 'EI':0.8, 'g':1} for tau in taus]
  426. taug05_prefix = ['tau_g0.5'+str(taus) for taus in taus]; taug05_featureDic = [{'mu':150, 'tau':tau, 'BP':1, 'EI':0.5, 'g':0.5} for tau in taus]
  427. taug15_prefix = ['tau_g1.5'+str(taus) for taus in taus]; taug15_featureDic = [{'mu':150, 'tau':tau, 'BP':1, 'EI':0.5, 'g':1.5} for tau in taus]
  428. tauBP05_prefix = ['tau_BP0.5'+str(taus) for taus in taus]; tauBP05_featureDic = [{'mu':150, 'tau':tau, 'BP':0.5, 'EI':0.5, 'g':1} for tau in taus]
  429. taumu165_prefix = ['tau_mu165'+str(taus) for taus in taus]; taumu165_featureDic = [{'mu':165, 'tau':tau, 'BP':1, 'EI':0.5, 'g':1} for tau in taus]
  430. taumu180_prefix = ['tau_mu180'+str(taus) for taus in taus]; taumu180_featureDic = [{'mu':180, 'tau':tau, 'BP':1, 'EI':0.5, 'g':1} for tau in taus]
  431. BPs = [0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
  432. BP_prefix = ['BP'+str(bp) for bp in BPs]; BP_featureDic = [{'mu':150, 'tau':15, 'BP':bp, 'EI':0.5, 'g':1} for bp in BPs]
  433. BPE3_prefix = ['BP_E3'+str(bp) for bp in BPs]; BPE3_featureDic = [{'mu':150, 'tau':15, 'BP':bp, 'EI':0.3, 'g':1} for bp in BPs]
  434. BPE8_prefix = ['BP_E8'+str(bp) for bp in BPs]; BPE8_featureDic = [{'mu':150, 'tau':15, 'BP':bp, 'EI':0.8, 'g':1} for bp in BPs]
  435. BPg05_prefix = ['BP_g0.5'+str(bp) for bp in BPs]; BPg05_featureDic = [{'mu':150, 'tau':15, 'BP':bp, 'EI':0.5, 'g':0.5} for bp in BPs]
  436. BPg15_prefix = ['BP_g1.5'+str(bp) for bp in BPs]; BPg15_featureDic = [{'mu':150, 'tau':15, 'BP':bp, 'EI':0.5, 'g':1.5} for bp in BPs]
  437. BPmu165_prefix = ['BP_mu165'+str(bp) for bp in BPs]; BPmu165_featureDic = [{'mu':165, 'tau':15, 'BP':bp, 'EI':0.5, 'g':1} for bp in BPs]
  438. BPmu180_prefix = ['BP_mu180'+str(bp) for bp in BPs]; BPmu180_featureDic = [{'mu':180, 'tau':15, 'BP':bp, 'EI':0.5, 'g':1} for bp in BPs]
  439. BPtau5_prefix = ['BP_tau5'+str(bp) for bp in BPs]; BPtau5_featureDic = [{'mu':150, 'tau':5, 'BP':bp, 'EI':0.5, 'g':1} for bp in BPs]
  440. BPtau40_prefix = ['BP_tau40'+str(bp) for bp in BPs]; BPtau40_featureDic = [{'mu':150, 'tau':40, 'BP':bp, 'EI':0.5, 'g':1} for bp in BPs]
  441. EIatio=[0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]
  442. EI_prefix = ['EI'+str(ei) for ei in EIatio]; EI_featureDic = [{'mu':150, 'tau':15, 'BP':1, 'EI':ei, 'g':1} for ei in EIatio]
  443. EIg05_prefix = ['EI_g0.5'+str(ei) for ei in EIatio if ei !=0.3]; EIg05_featureDic = [{'mu':150, 'tau':15, 'BP':1, 'EI':ei, 'g':0.5} for ei in EIatio if ei !=0.3]
  444. EIg15_prefix = ['EI_g1.5'+str(ei) for ei in EIatio]; EIg15_featureDic = [{'mu':150, 'tau':15, 'BP':1, 'EI':ei, 'g':1.5} for ei in EIatio]
  445. EIBP05_prefix = ['EI_BP0.5'+str(ei) for ei in EIatio]; EIBP05_featureDic = [{'mu':150, 'tau':15, 'BP':0.5, 'EI':ei, 'g':1} for ei in EIatio]
  446. EImu165_prefix = ['EI_mu165'+str(ei) for ei in EIatio]; EImu165_featureDic = [{'mu':165, 'tau':15, 'BP':1, 'EI':ei, 'g':1} for ei in EIatio]
  447. EImu180_prefix = ['EI_mu180'+str(ei) for ei in EIatio]; EImu180_featureDic = [{'mu':180, 'tau':15, 'BP':1, 'EI':ei, 'g':1} for ei in EIatio]
  448. EItau5_prefix = ['EI_tau5'+str(ei) for ei in EIatio]; EItau5_featureDic = [{'mu':150, 'tau':5, 'BP':1, 'EI':ei, 'g':1} for ei in EIatio]
  449. EItau40_prefix = ['EI_tau40'+str(ei) for ei in EIatio]; EItau40_featureDic = [{'mu':150, 'tau':40, 'BP':1, 'EI':ei, 'g':1} for ei in EIatio]
  450. gbars = [0.4, 0.6, 0.8, 1.0, 1.2, 1.4, 1.6, 1.8, 2.0]
  451. gbar_prefix = ['g'+str(g) for g in gbars]; gbar_featureDic = [{'mu':150, 'tau':15, 'BP':1, 'EI':0.5, 'g':g} for g in gbars]
  452. gbarE3_prefix = ['g_E3'+str(g) for g in gbars if g!= 0.4]; gbarE3_featureDic = [{'mu':150, 'tau':15, 'BP':1, 'EI':0.3, 'g':g} for g in gbars if g!= 0.4]
  453. gbarE8_prefix = ['g_E8'+str(g) for g in gbars]; gbarE8_featureDic = [{'mu':150, 'tau':15, 'BP':1, 'EI':0.8, 'g':g} for g in gbars]
  454. gbarBP05_prefix = ['g_BP0.5'+str(g) for g in gbars if g!= 0.4]; gbarBP05_featureDic = [{'mu':150, 'tau':15, 'BP':0.5, 'EI':0.5, 'g':g} for g in gbars if g!= 0.4]
  455. gbarmu165_prefix = ['g_mu165'+str(g) for g in gbars if g!= 0.4]; gbarmu165_featureDic = [{'mu':165, 'tau':15, 'BP':1, 'EI':0.5, 'g':g} for g in gbars if g!= 0.4]
  456. gbarmu180_prefix = ['g_mu180'+str(g) for g in gbars if g!= 0.4]; gbarmu180_featureDic = [{'mu':180, 'tau':15, 'BP':1, 'EI':0.5, 'g':g} for g in gbars if g!= 0.4]
  457. gbartau5_prefix = ['g_tau5'+str(g) for g in gbars if g!= 0.4]; gbartau5_featureDic = [{'mu':150, 'tau':5, 'BP':1, 'EI':0.5, 'g':g} for g in gbars if g!= 0.4]
  458. gbartau40_prefix = ['g_tau40'+str(g) for g in gbars if g!= 0.4]; gbartau40_featureDic = [{'mu':150, 'tau':40, 'BP':1, 'EI':0.5, 'g':g} for g in gbars if g!= 0.4]
  459. file_prefix_E5_list = mu_prefix + tau_prefix + BP_prefix + EI_prefix + gbar_prefix
  460. file_prefix_E3_list = muE3_prefix + tauE3_prefix + BPE3_prefix + EI_prefix + gbarE3_prefix
  461. file_prefix_E8_list = muE8_prefix + tauE8_prefix + BPE8_prefix + EI_prefix + gbarE8_prefix
  462. file_prefix_g05_list = mug05_prefix + taug05_prefix + BPg05_prefix + EIg05_prefix + gbar_prefix
  463. file_prefix_g15_list = mug15_prefix + taug15_prefix + BPg15_prefix + EIg15_prefix + gbar_prefix
  464. file_prefix_BP05_list = muBP05_prefix + tauBP05_prefix + BP_prefix + EIBP05_prefix + gbarBP05_prefix
  465. file_prefix_mu165_list = mu_prefix + taumu165_prefix + BPmu165_prefix + EImu165_prefix + gbarmu165_prefix
  466. file_prefix_mu180_list = mu_prefix + taumu180_prefix + BPmu180_prefix + EImu180_prefix + gbarmu180_prefix
  467. file_prefix_tau5_list = mutau5_prefix + tau_prefix + BPtau5_prefix + EItau5_prefix + gbartau5_prefix
  468. file_prefix_tau40_list = mutau40_prefix + tau_prefix + BPtau40_prefix + EItau40_prefix + gbartau40_prefix
  469. file_featureDic_E5_list = mu_featureDic + tau_featureDic + BP_featureDic + EI_featureDic + gbar_featureDic
  470. file_featureDic_E3_list = muE3_featureDic + tauE3_featureDic + BPE3_featureDic + EI_featureDic + gbarE3_featureDic
  471. file_featureDic_E8_list = muE8_featureDic + tauE8_featureDic + BPE8_featureDic + EI_featureDic + gbarE8_featureDic
  472. file_featureDic_g05_list = mug05_featureDic + taug05_featureDic + BPg05_featureDic + EIg05_featureDic + gbar_featureDic
  473. file_featureDic_g15_list = mug15_featureDic + taug15_featureDic + BPg15_featureDic + EIg15_featureDic + gbar_featureDic
  474. file_featureDic_BP05_list = muBP05_featureDic + tauBP05_featureDic + BP_featureDic + EIBP05_featureDic + gbarBP05_featureDic
  475. file_featureDic_mu165_list = mu_featureDic + taumu165_featureDic + BPmu165_featureDic + EImu165_featureDic + gbarmu165_featureDic
  476. file_featureDic_mu180_list = mu_featureDic + taumu180_featureDic + BPmu180_featureDic + EImu180_featureDic + gbarmu180_featureDic
  477. file_featureDic_tau5_list = mutau5_featureDic + tau_featureDic + BPtau5_featureDic + EItau5_featureDic + gbartau5_featureDic
  478. file_featureDic_tau40_list = mutau40_featureDic + tau_featureDic + BPtau40_featureDic + EItau40_featureDic + gbartau40_featureDic
  479. all_file_prefix_list = file_prefix_E3_list + file_prefix_E5_list + file_prefix_E8_list + file_prefix_g05_list + file_prefix_g15_list +\
  480. file_prefix_BP05_list + file_prefix_mu165_list + file_prefix_mu180_list + file_prefix_tau5_list + file_prefix_tau40_list
  481. print(len(all_file_prefix_list), 'perturbations in total')
  482. all_perturbation_loader = [load_oneLoader(dataloader_path, file_prefix+'_loader.pkl') for file_prefix in all_file_prefix_list]
  483. # %% [markdown]
  484. # ### calculation
  485. # %%
  486. criterionC = nn.BCEWithLogitsLoss()
  487. criterionW = HuberLossWithWeight()
  488. fm_mean_cols = ['conv1_fm0_mean', 'conv1_fm1_mean', 'conv1_fm2_mean', 'conv1_fm3_mean', 'conv2_fm0_mean', 'conv2_fm1_mean']
  489. predSummary_list = []
  490. # retrive the best baseline model ---------------------------------------------------------------
  491. bsType = 'all'
  492. connCNN_Model = connCNN
  493. weightCNN_Model = weighCNN
  494. # predict on baseline-test set -------------------------------------------------------------------
  495. train_loader, val_loader, test_loader = load_dataloaders(dataloader_path, 'baseline_'+bsType+'_loader.pkl')
  496. # pred with connCNN ------------
  497. avg_connect_loss, accu, pred_list,confi_list, label_list, w_list, sCInfo_list = evaluateConnCNN_wConfi_model(connCNN_Model, test_loader, criterionC)
  498. # pred with weightCNN ----------
  499. avg_wegiht_loss, predW_list, gtW_list, label0_loss, label1_loss, sWInfo_list = evaluateWeightCNN_model(weightCNN_Model, test_loader, criterionW)
  500. # feature map mean -------------
  501. conn_fmMean_list = []; weight_fmMean_list = []
  502. conn_fmMean_pre_list = []; weight_fmMean_pre_list = []
  503. conn_fmMean_postBN_list = []; weight_fmMean_postBN_list = []
  504. for ccgs, cs, ws, ss in test_loader:
  505. for m in range(len(ccgs)):
  506. # pre activation (raw featureMaps)
  507. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(connCNN_Model, ccgs[m]) # connCNN
  508. conn_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  509. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(weightCNN_Model, ccgs[m])# weightCNN
  510. weight_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  511. # postBN (featureMaps after BatchNorm, but before tanh)
  512. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(connCNN_Model, ccgs[m]) # connCNN
  513. conn_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  514. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(weightCNN_Model, ccgs[m])# weightCNN
  515. weight_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  516. # post activation (featureMaps after BatchNorm & tanh)
  517. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(connCNN_Model, ccgs[m]) # connCNN
  518. conn_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  519. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(weightCNN_Model, ccgs[m])# weightCNN
  520. weight_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  521. predSummary = pd.DataFrame([label_list, pred_list, confi_list, gtW_list, predW_list, sCInfo_list, ['baseline_'+bsType]*len(sCInfo_list)],
  522. index=['connLabel', 'predLabel', 'confi', 'gtWeights', 'predWeights', 'sample', 'test']).T
  523. connPre_FMM = pd.DataFrame(np.array(conn_fmMean_pre_list), columns=['pre_conn_'+col for col in fm_mean_cols])
  524. connPost_FMM = pd.DataFrame(np.array(conn_fmMean_list), columns=['post_conn_'+col for col in fm_mean_cols])
  525. connPostBN_FMM = pd.DataFrame(np.array(conn_fmMean_postBN_list), columns=['postBN_conn_'+col for col in fm_mean_cols])
  526. weightPre_FMM = pd.DataFrame(np.array(weight_fmMean_pre_list), columns=['pre_weight_'+col for col in fm_mean_cols])
  527. weightPost_FMM = pd.DataFrame(np.array(weight_fmMean_list), columns=['post_weight_'+col for col in fm_mean_cols])
  528. weightPostBN_FMM = pd.DataFrame(np.array(weight_fmMean_postBN_list), columns=['postBN_weight_'+col for col in fm_mean_cols])
  529. predSummary = pd.concat([predSummary, connPre_FMM, connPost_FMM, connPostBN_FMM, weightPre_FMM, weightPost_FMM, weightPostBN_FMM], axis=1); predSummary_list.append(predSummary)
  530. # predict on associated perturbations ------------------------------------------------------------
  531. for fileIdx in range(len(all_file_prefix_list)):
  532. perb_type = all_file_prefix_list[fileIdx]
  533. perb_loader = all_perturbation_loader[fileIdx]
  534. # pred with connCNN ------------
  535. avg_connect_loss, accu, pred_list,confi_list, label_list, w_list, sCInfo_list = evaluateConnCNN_wConfi_model(connCNN_Model, perb_loader, criterionC)
  536. # pred with weightCNN ----------
  537. avg_wegiht_loss, predW_list, gtW_list, label0_loss, label1_loss, sWInfo_list = evaluateWeightCNN_model(weightCNN_Model, perb_loader, criterionW)
  538. assert np.array_equal(sCInfo_list, sWInfo_list), bsType+' '+perb_type+" sInfo Arrays are not equal"
  539. # get feature map mean on each ccg -------------
  540. conn_fmMean_list = []; weight_fmMean_list = []
  541. conn_fmMean_pre_list = []; weight_fmMean_pre_list = []
  542. conn_fmMean_postBN_list = []; weight_fmMean_postBN_list = []
  543. for ccgs, cs, ws, ss in perb_loader:
  544. for m in range(len(ccgs)):
  545. # pre activation (raw featureMaps)
  546. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(connCNN_Model, ccgs[m]) # connCNN
  547. conn_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  548. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps(weightCNN_Model, ccgs[m])# weightCNN
  549. weight_fmMean_pre_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  550. # postBN (featureMaps after BatchNorm, but before tanh)
  551. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(connCNN_Model, ccgs[m]) # connCNN
  552. conn_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  553. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_postBN(weightCNN_Model, ccgs[m])# weightCNN
  554. weight_fmMean_postBN_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  555. # post activation (featureMaps after BatchNorm & tanh)
  556. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(connCNN_Model, ccgs[m])
  557. conn_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  558. fm_conv1, fm_conv2, activation_1, activation_2 = cal_feature_maps_post_activation(weightCNN_Model, ccgs[m])
  559. weight_fmMean_list.append(fm_conv1.mean(dim=1).tolist() + fm_conv2.mean(dim=1).tolist())
  560. predSummary = pd.DataFrame([label_list, pred_list, confi_list, gtW_list, predW_list, sCInfo_list, [perb_type]*len(sCInfo_list)],
  561. index=['connLabel', 'predLabel', 'confi', 'gtWeights', 'predWeights', 'sample', 'test']).T
  562. connPre_FMM = pd.DataFrame(np.array(conn_fmMean_pre_list), columns=['pre_conn_'+col for col in fm_mean_cols])
  563. connPost_FMM = pd.DataFrame(np.array(conn_fmMean_list), columns=['post_conn_'+col for col in fm_mean_cols])
  564. connPostBN_FMM = pd.DataFrame(np.array(conn_fmMean_postBN_list), columns=['postBN_conn_'+col for col in fm_mean_cols])
  565. weightPre_FMM = pd.DataFrame(np.array(weight_fmMean_pre_list), columns=['pre_weight_'+col for col in fm_mean_cols])
  566. weightPost_FMM = pd.DataFrame(np.array(weight_fmMean_list), columns=['post_weight_'+col for col in fm_mean_cols])
  567. weightPostBN_FMM = pd.DataFrame(np.array(weight_fmMean_postBN_list), columns=['postBN_weight_'+col for col in fm_mean_cols])
  568. predSummary = pd.concat([predSummary, connPre_FMM, connPost_FMM, connPostBN_FMM, weightPre_FMM, weightPost_FMM, weightPostBN_FMM], axis=1); predSummary_list.append(predSummary)
  569. allBaseline_bspredSummary = pd.concat(predSummary_list)
  570. allBaseline_bspredSummary.to_csv(os.path.join(confiSave_path, 'allBaselines_confi_w_perbs.csv'))
  571. allBaseline_bspredSummary.head(2)
  572. # %% [markdown]
  573. # ### combine FMM, Confi, Indictors
  574. # %%
  575. # perb_exc_all, perb_inh_all = [], []
  576. # for bIdx in range(len(baseline_types)):
  577. # bsType = baseline_types[bIdx]
  578. # bspredSummary = pd.read_csv(os.path.join(confiSave_path, bsType+'_confi_w_perbs.csv')).iloc[:, 1:]
  579. # exc_indicator_df = pd.read_csv(os.path.join(indiSave_path, bsType+'_indicator_exc.csv')).iloc[:, 1:]
  580. # perb_exc_all.append(exc_indicator_df)
  581. # inh_indicator_df = pd.read_csv(os.path.join(indiSave_path, bsType+'_indicator_inh.csv')).iloc[:, 1:]
  582. # perb_inh_all.append(inh_indicator_df)
  583. perb_exc_indis = pd.concat(perb_exc_all, ignore_index=True); print(perb_exc_indis.shape)
  584. perb_inh_indis = pd.concat(perb_inh_all, ignore_index=True); print(perb_einhindis.shape)
  585. allBaseline_pred_indi_exc = pd.merge(allBaseline_bspredSummary, perb_exc_indis, on='sample', how='inner')
  586. allBaseline_pred_indi_exc = allBaseline_pred_indi_exc.drop(columns=['test_y']); allBaseline_pred_indi_exc = allBaseline_pred_indi_exc.rename(columns={'test_x': 'test'})
  587. allBaseline_pred_indi_inh = pd.merge(allBaseline_bspredSummary, perb_inh_indis, on='sample', how='inner')
  588. allBaseline_pred_indi_inh = allBaseline_pred_indi_inh.drop(columns=['test_y']); allBaseline_pred_indi_inh = allBaseline_pred_indi_inh.rename(columns={'test_x': 'test'})
  589. print('allBaseline_pred_indi_exc.shape:', allBaseline_pred_indi_exc.shape)
  590. print('allBaseline_pred_indi_inh.shape:', allBaseline_pred_indi_inh.shape)
  591. allBaseline_pred_indi_exc.to_csv(os.path.join(resultAnalysis_path, 'allBaselines_confi_fmm_indi_exc.csv'))
  592. allBaseline_pred_indi_inh.to_csv(os.path.join(resultAnalysis_path, 'allBaselines_confi_fmm_indi_inh.csv'))
  593. print('Done')
  594. allBaseline_pred_indi_inh.head(2)
  595. # %%
  596. # %% [markdown]
  597. # # ----------------------------------------------------------------------
  598. # # Test - Indicator Cal
  599. # %%
  600. i=0
  601. fileTest = file_test_list[i]
  602. featureDic = file_featureDics[i]
  603. bs_ccg, ws, ls, ss = read_ccg(data_path, fileTest+'_CCG.csv')
  604. # %%
  605. _ = ccg_indicators(bs_ccg[0], connection_type='exc', smoothSigma=0.5, bin_size=1, peak_window_ms=10, ifVerbose=True, ifPlot=True, timebins=timebins, ifSave=False)
  606. print(_); print('---------------------------------------')
  607. _ = ccg_indicators(bs_ccg[20], connection_type='inh', smoothSigma=0.5, bin_size=1, peak_window_ms=10, ifVerbose=True, ifPlot=True, timebins=timebins, ifSave=False)
  608. print(_); print('---------------------------------------')
  609. # %% [markdown]
  610. # ### Digest - exc
  611. # %%
  612. i=0
  613. fileTest = file_test_list[i]
  614. featureDic = file_featureDics[i]
  615. bs_ccg, ws, ls, ss = read_ccg(data_path, fileTest+'_CCG.csv')
  616. print('looking at', fileTest, ccg.shape)
  617. ws[20]
  618. # %%
  619. bin_size = 1
  620. peak_window_ms = 10
  621. center = len(ccg) // 2
  622. baseline = np.mean(np.concatenate([ccg[:10], ccg[-10:]]))
  623. ccg = bs_ccg[0] #327
  624. # peak
  625. search_bins = int(peak_window_ms / bin_size)
  626. search_region = ccg[center-search_bins : center+search_bins+1]
  627. peak_idx_rel = np.argmax(search_region)
  628. peak_idx = center - search_bins + peak_idx_rel
  629. peak_val = ccg[peak_idx]
  630. half_val = peak_val / 2
  631. peak_lag = (peak_idx - center) * bin_size
  632. print('peak happens at bin', peak_idx, '=', peak_val, 'with time lag =', peak_lag )
  633. plt.figure(figsize=(4,2))
  634. plt.bar(timebins, ccg, color=ccgColor, width=1); plt.title('raw'); plt.show()
  635. # %%
  636. # peak width (full width at half max)
  637. half_val = peak_height / 2
  638. left, right = peak_idx, peak_idx
  639. while left > 0 and ccg[left] > half_val:
  640. left -= 1
  641. while right < len(ccg) - 1 and ccg[right] > half_val:
  642. right += 1
  643. peak_width = (right - left) * bin_size
  644. print('peak drop to half in', peak_width, 'bins')
  645. print(ccg[peak_idx:peak_idx+peak_width])
  646. # %%
  647. # noise estimation from tails ----------------------------------
  648. tail_bins = np.r_[np.arange(25), np.arange(len(ccg) - 25, len(ccg))]
  649. noise_floor = np.mean(ccg[tail_bins])
  650. noise_std = np.std(ccg[tail_bins])
  651. peak_to_noise = ((peak_val - noise_floor) / (noise_std + 1e-10))
  652. print('peak_to_noise =', peak_to_noise)
  653. # %%
  654. # entropy
  655. ccg_smooth = gaussian_filter1d(ccg, sigma=0.5)
  656. ccg_prob = ccg_smooth / (np.sum(ccg_smooth) + 1e-10)
  657. ccg_entropy = entropy(ccg_prob, base=2)
  658. ccg_entropy_norm = ccg_entropy / np.log2(len(ccg))
  659. print('entropy =', ccg_entropy, 'norm entropy =', ccg_entropy_norm)
  660. fig, ax = plt.subplots(1, 3, figsize=(12, 2))
  661. ax[0].bar(timebins, ccg, color=ccgColor, width=1); ax[0].set_title('raw')
  662. ax[1].bar(timebins, ccg_smooth, color=ccgColor, width=1); ax[1].set_title('smoothed')
  663. ax[2].bar(timebins, ccg_prob, color=ccgColor, width=1); ax[2].set_title('prob')
  664. plt.tight_layout(); plt.show()
  665. # %%
  666. # KL divergence only within ±10 bins of the center
  667. peak_window_ms = 10
  668. window_bins = int(peak_window_ms / bin_size)
  669. kl_window = ccg[center - window_bins:center + window_bins + 1]
  670. P = kl_window / (np.sum(kl_window) + 1e-10)
  671. U = np.ones_like(P) / len(P)
  672. kl_div = np.sum(rel_entr(P, U)) / np.log(2) # in bits
  673. print('kl_div =', kl_div)
  674. fig, ax = plt.subplots(1, 2, figsize=(12, 2))
  675. ax[0].bar(timebins, ccg, color=ccgColor, width=1); ax[0].set_title('raw')
  676. ax[1].bar(timebins, ccg, color=ccgColor, width=1);
  677. ax[1].bar(timebins[center - window_bins:center + window_bins + 1], kl_window, color='r', width=1);ax[1].set_title('shifted-focused')
  678. plt.tight_layout(); plt.show()
  679. # %%
  680. # entropy only within ±10 bins of the center
  681. ccg_entropy_focus = entropy(P, base=2)
  682. ccg_entropy_focus_norm = ccg_entropy_focus / np.log2(len(kl_window))
  683. print('ccg_entropy_focus =', ccg_entropy_focus, 'ccg_entropy_focus_norm =', ccg_entropy_focus_norm)
  684. # %%
  685. # Temporal span above noise threshold
  686. thresh = noise_floor + 2 * noise_std
  687. left_span, right_span = peak_idx, peak_idx
  688. while left_span > 0 and ccg_smooth[left_span] > thresh:
  689. left_span -= 1
  690. while right_span < len(ccg_smooth) - 1 and ccg_smooth[right_span] > thresh:
  691. right_span += 1
  692. temporal_span = (right_span - left_span) * bin_size
  693. print('from', left_span, 'to', right_span, 'bins, we have ccg above 0.25*peak/dip')
  694. fig, ax = plt.subplots(1, 2, figsize=(8, 2))
  695. ax[0].bar(timebins, ccg, color=ccgColor, width=1); ax[0].set_title('raw')
  696. ax[1].bar(timebins, ccg_smooth, color=ccgColor, width=1); ax[1].set_title('smoothed')
  697. ax[1].bar(timebins[left_span:right_span], ccg_smooth[left_span:right_span],
  698. color=mC, width=1, alpha=0.4, label='temporal span'); ax[1].legend(loc=1)
  699. plt.tight_layout(); plt.show()
  700. # %% [markdown]
  701. # ### Digest - inh
  702. # %%
  703. bin_size = 1
  704. ccg = bs_ccg[20]
  705. peak_window_ms = 10
  706. search_bins = int(peak_window_ms / bin_size)
  707. search_region = ccg[center-search_bins : center+search_bins+1]
  708. baseline = np.mean(np.concatenate([ccg[:10], ccg[-10:]]))
  709. # peak
  710. # center = len(ccg) // 2
  711. # peak_height = np.min(ccg)
  712. # peak_idx = np.argmin(ccg)
  713. # peak_lag = (peak_idx - center) * bin_size
  714. peak_idx_rel = np.argmin(search_region)
  715. peak_idx = center - search_bins + peak_idx_rel
  716. peak_val = ccg[peak_idx]
  717. peak_lag = (peak_idx - center) * bin_size
  718. print('dip peak happens at bin', peak_idx, '=', peak_val, 'with time lag =', peak_lag )
  719. plt.figure(figsize=(4,2))
  720. plt.bar(timebins, ccg, color=ccgColor, width=1); plt.title('raw'); plt.show()
  721. # %%
  722. # peak width (full width at half max)
  723. # half_val = peak_height / 2
  724. # left, right = peak_idx, peak_idx
  725. # while left > 0 and ccg[left] < half_val:
  726. # left -= 1
  727. # while right < len(ccg) - 1 and ccg[right] < half_val:
  728. # right += 1
  729. # peak_width = (right - left) * bin_size
  730. half_val = (peak_val + baseline) / 2
  731. left, right = peak_idx, peak_idx
  732. while left > 0 and (ccg[left] < half_val):
  733. left -= 1
  734. while right < len(ccg) - 1 and (ccg[right] < half_val):
  735. right += 1
  736. peak_width = (right - left) * bin_size
  737. print('peak drop to half in', peak_width, 'bins')
  738. print(ccg[peak_idx:peak_idx+peak_width])
  739. # %%
  740. # noise estimation from tails ----------------------------------
  741. tail_bins = np.r_[np.arange(25), np.arange(len(ccg) - 25, len(ccg))]
  742. noise_floor = np.mean(ccg[tail_bins])
  743. noise_std = np.std(ccg[tail_bins])
  744. peak_to_noise = ((peak_val - noise_floor) / (noise_std + 1e-10)) if connection_type == 'exc' else ((noise_floor - peak_val) / (noise_std + 1e-10))
  745. # %%
  746. # %%
  747. # entropy
  748. ccg_smooth = gaussian_filter1d(ccg, sigma=0.5)
  749. ccg_shifted = ccg_smooth - np.min(ccg_smooth) # shift to make in nonnegative
  750. ccg_prob = ccg_shifted / (np.sum(ccg_shifted) + 1e-10)
  751. ccg_entropy = entropy(ccg_prob, base=2)
  752. ccg_entropy_norm =ccg_entropy / np.log2(len(ccg_smooth))
  753. print('entropy =', ccg_entropy)
  754. print('normed entropy =', ccg_entropy_norm)
  755. fig, ax = plt.subplots(1, 4, figsize=(12, 2))
  756. ax[0].bar(timebins, ccg, color=ccgColor, width=1); ax[0].set_title('raw')
  757. ax[1].bar(timebins, ccg_smooth, color=ccgColor, width=1); ax[1].set_title('smoothed')
  758. ax[2].bar(timebins, ccg_shifted, color=ccgColor, width=1); ax[2].set_title('shifted')
  759. ax[3].bar(timebins, ccg_prob, color=ccgColor, width=1); ax[3].set_title('prob')
  760. plt.tight_layout(); plt.show()
  761. # %%
  762. peak_window_ms = 10
  763. window_bins = int(peak_window_ms / bin_size)
  764. kl_window = ccg[center - window_bins:center + window_bins + 1]
  765. # inh
  766. kl_window =np.max(kl_window) - kl_window
  767. P = kl_window / (np.sum(kl_window) + 1e-10)
  768. U = np.ones_like(P) / len(P)
  769. kl_div = np.sum(rel_entr(P, U)) # in bits
  770. print('kl_div =', kl_div)
  771. fig, ax = plt.subplots(1, 3, figsize=(12, 2))
  772. ax[0].bar(timebins, ccg, color=ccgColor, width=1); ax[0].set_title('raw')
  773. ax[1].bar(timebins, ccg, color=ccgColor, width=1);
  774. ax[1].bar(timebins[center - window_bins:center + window_bins + 1], ccg[center - window_bins:center + window_bins + 1], color='r', width=1);ax[1].set_title('focused')
  775. ax[2].bar(timebins[center - window_bins:center + window_bins + 1], kl_window, color='r', width=1);ax[2].set_title('shifted-focused')
  776. plt.tight_layout(); plt.show()
  777. # %%
  778. # Temporal span above noise threshold
  779. thresh = 0.25 * peak_height
  780. left_span, right_span = peak_idx, peak_idx
  781. while left_span > 0 and ccg_smooth[left_span] < thresh:
  782. left_span -= 1
  783. while right_span < len(ccg_smooth) - 1 and ccg_smooth[right_span] < thresh:
  784. right_span += 1
  785. temporal_span = (right_span - left_span) * bin_size
  786. print('from', left_span, 'to', right_span, 'bins, we have ccg above 0.25*peak/dip')
  787. fig, ax = plt.subplots(1, 2, figsize=(8, 2))
  788. ax[0].bar(timebins, ccg, color=ccgColor, width=1); ax[0].set_title('raw')
  789. ax[1].bar(timebins, ccg_smooth, color=ccgColor, width=1); ax[1].set_title('smoothed')
  790. ax[1].bar(timebins[left_span:right_span], ccg_smooth[left_span:right_span],
  791. color=mC, width=1, alpha=0.4, label='temporal span'); ax[1].legend(loc=1)
  792. plt.tight_layout(); plt.show()
  793. # %%
  794. # %% [markdown]
  795. # ### Function
  796. # %%
  797. def ccg_indicators(ccg, connection_type='exc', smoothSigma=0.5, bin_size=1, peak_window_ms=10, ifVerbose=False, ifPlot=False, timebins=None, figsize=(8,2), barColor='lightslategrey', spanColor='firebrick', ifSave=True, savePath=None, filename=None):
  798. '''
  799. Calculate signal indicators for CCGs.
  800. For excitatory connections (peaks) and inhibitory connections (dips), adjusts indicators accordingly.
  801. Returns: a dictionary with standardized keys:
  802. - peak_height or dip_depth
  803. - peak_to_noise or dip_to_noise
  804. - peak_halfMax_width or dip_halfMax_width
  805. - peak_lag or dip_lag
  806. - entropy
  807. - temporal_span
  808. '''
  809. ccg = np.array(ccg, dtype=np.float32)
  810. N = len(ccg); center = N // 2
  811. baseline = np.mean(np.concatenate([ccg[:10], ccg[-10:]]))
  812. # search for peak/dip -----------------------------------------
  813. # restrict to [-peak_window_ms, peak_window_ms] range around the center, e.g., -10ms ~ 10ms
  814. search_bins = int(peak_window_ms / bin_size)
  815. search_region = ccg[center-search_bins : center+search_bins+1]
  816. if connection_type == 'exc':
  817. peak_idx_rel = np.argmax(search_region)
  818. peak_idx = center - search_bins + peak_idx_rel
  819. peak_val = ccg[peak_idx]
  820. direction = 'peak'
  821. elif connection_type == 'inh':
  822. peak_idx_rel = np.argmin(search_region)
  823. peak_idx = center - search_bins + peak_idx_rel
  824. peak_val = ccg[peak_idx]
  825. direction = 'dip'
  826. else:
  827. raise ValueError("connection_type must be 'exc' or 'inh'")
  828. peak_lag = (peak_idx - center) * bin_size
  829. # half-max width -----------------------------------------------
  830. half_val = peak_val / 2 if connection_type == 'exc' else (peak_val + baseline) / 2
  831. left, right = peak_idx, peak_idx
  832. while left > 0 and ((ccg[left] > half_val) if connection_type == 'exc' else (ccg[left] < half_val)):
  833. left -= 1
  834. while right < len(ccg) - 1 and ((ccg[right] > half_val) if connection_type == 'exc' else (ccg[right] < half_val)):
  835. right += 1
  836. peak_width = (right - left) * bin_size
  837. # noise estimation from tails ----------------------------------
  838. tail_bins = np.r_[np.arange(25), np.arange(len(ccg) - 25, len(ccg))]
  839. noise_floor = np.mean(ccg[tail_bins])
  840. noise_std = np.std(ccg[tail_bins])
  841. peak_to_noise = ((peak_val - noise_floor) / (noise_std + 1e-10)) if connection_type == 'exc' else ((noise_floor - peak_val) / (noise_std + 1e-10))
  842. # smooth for entropy and temporal span --------------------------
  843. ccg_smooth = gaussian_filter1d(ccg, sigma=smoothSigma)
  844. ccg_prob = ccg_smooth / (np.sum(ccg_smooth) + 1e-10)
  845. ccg_entropy = entropy(ccg_prob, base=2)
  846. ccg_entropy_norm = ccg_entropy / np.log2(N)
  847. # KL Divergence (+- 10ms window) ---------------------------------
  848. window_bins = int(peak_window_ms / bin_size)
  849. kl_window = ccg[center - window_bins:center + window_bins + 1]
  850. P = kl_window / (np.sum(kl_window) + 1e-10)
  851. U = np.ones_like(P) / len(P)
  852. kl_div = np.sum(rel_entr(P, U)) / np.log(2) # in bits
  853. # entropy (+- 10ms window) -----------------------------------------
  854. window_entropy = entropy(P, base=2)
  855. window_entropy_norm = window_entropy / np.log2(len(P))
  856. # temporal span based on threshold
  857. thresh = noise_floor + 2 * noise_std if connection_type == 'exc' else noise_floor - 2 * noise_std
  858. left_span, right_span = peak_idx, peak_idx
  859. while left_span > 0 and ((ccg_smooth[left_span] > thresh) if connection_type == 'exc' else (ccg_smooth[left_span] < thresh)):
  860. left_span -= 1
  861. while right_span < len(ccg_smooth) - 1 and ((ccg_smooth[right_span] > thresh) if connection_type == 'exc' else (ccg_smooth[right_span] < thresh)):
  862. right_span += 1
  863. temporal_span = (right_span - left_span) * bin_size
  864. if ifVerbose:
  865. print(f"{direction} occurs at bin {peak_idx} = {peak_val} with lag = {peak_lag}")
  866. print(f"{direction} drops to half in {peak_width} bins")
  867. print(f"Tails: mean = {round(noise_floor, 3)}, std = {round(noise_std, 3)}, {direction}_to_noise = {round(peak_to_noise, 3)}")
  868. print(f"Temporal span above threshold ({thresh:.2f}) = {temporal_span} bins around center")
  869. print(f"Entropy = {round(ccg_entropy, 3)}")
  870. print(f"Normalized entropy = {ccg_entropy_norm:.3f}")
  871. print(f"Normalized entropy (±{peak_window_ms} ms) = {window_entropy_norm:.3f}")
  872. print(f"KL divergence (±{peak_window_ms} ms) = {kl_div:.3f}")
  873. if ifPlot:
  874. fig, ax = plt.subplots(1, 3, figsize=figsize)
  875. ax[0].bar(timebins, ccg, color=barColor, width=1); ax[0].set_title('raw')
  876. ax[1].bar(timebins, ccg_smooth, color=barColor, width=1); ax[1].set_title('smoothed')
  877. ax[1].bar(timebins[left_span:right_span], ccg_smooth[left_span:right_span],
  878. color=spanColor, width=1, alpha=0.4, label='temporal span'); ax[1].legend(loc=1)
  879. ax[2].bar(timebins[center-window_bins : center+window_bins+1], kl_window, color='g', width=1);ax[2].set_title('+-10 focus window')
  880. plt.tight_layout()
  881. if ifSave:
  882. if not os.path.exists(savePath):
  883. os.makedirs(savePath)
  884. plt.savefig(os.path.join(savePath, filename))
  885. plt.close()
  886. else:
  887. plt.show()
  888. # Return unified dictionary with consistent key naming
  889. result = {
  890. f"{direction}_height" if direction == 'peak' else "dip_depth": peak_val,
  891. f"{direction}_lag": peak_lag,
  892. f"{direction}_halfMax_width": peak_width,
  893. f"{direction}_to_noise": peak_to_noise,
  894. "temporal_span": temporal_span,
  895. "entropy": ccg_entropy,
  896. "norm_entropy": ccg_entropy_norm,
  897. 'norm_entropy_window': window_entropy_norm,
  898. "kl_divergence_window": kl_div
  899. }
  900. return result
  901. # %%
  902. direction='dip'
  903. f"{direction}_height" if direction == 'peak' else "dip_depth"
  904. # %%

4_0_indi_confi_fmm_Calu.ipynb at commit 6bcacec, no license · at the source

Overview

Authors: Xiaoqian Sun1, Hui Lu2,3, Chen Zeng4, Rahul Simha1
  1. Department of Computer Science, School of Engineering and Applied Science, The George Washington University, Washington, District of Columbia, United States of America
  2. The GW Institute for Neuroscience, The George Washington University, Washington, District of Columbia, United States of America
  3. Department of Pharmacology and Physiology, School of Medicine and Health Sciences, The George Washington University, Washington, District of Columbia, United States of America
  4. Department of Physics, Columbian College of Arts and Sciences, The George Washington University, Washington, District of Columbia, United States of America
Institutions: George Washington University (United States)
Journal: PLoS computational biology, volume 22, issue 8, article e1014615
Dates: received 9 February 2026; accepted 22 July 2026; published online 10 August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pcbi.1014615 · PMID 42574490 · PMCID PMC13475989 · OpenAlex W7202112642
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: extracellular electrophysiology (units, LFP) (modality), human (organism)
Methods: Connectivity, Machine learning, Statistics, Single-unit activity, calcium imaging
MeSH: Action Potentials*, Convolutional Neural Networks*, Machine Learning*, Models, Neurological*, Nerve Net*, Animals, Computational Biology, Computer Simulation, Humans, Neurons, Synapses (* major topic)
Topic: Advanced Memory and Neural Computing (Electrical and Electronic Engineering, Engineering), according to OpenAlex
Funding: George Washington University (2018–2023 Cross-Disciplinary Research Fund); National Institutes of Health (R01NS118197); NINDS NIH HHS (R01 NS118197)
Citations: not cited yet (Europe PMC); 51 references in the paper

Abstract

Understanding neuronal topology—how neurons are connected—is essential for uncovering neural computation principles and functional organization. However, accurately reconstructing such connectivity remains challenging due to the indirect nature of neural recordings and the complexity of network dynamics. As a first step towards this problem, a growing body of work has explored inferring monosynaptic connectivity directly from spike data. Among these, convolutional neural networks have shown promise when applied to spike-train cross-correlograms. Nevertheless, their ability to generalize across realistic experimental variability and the internal features that drive their predictions remain poorly understood. In this paper, we present a systematic benchmarking and diagnostic study of neural-network-based synaptic inference using simulations across a broad range of biophysical regimes. We show that connectivity classification and synaptic weight estimation, though often combined, rely on distinct internal representations and exhibit markedly different generalization behavior: robust connectivity models emphasize global structure in spike-train correlations, whereas weight estimation models are more sensitive to local signal amplitude and generalize less predictably. Importantly, we find that training on pooled, biologically grounded simulation data substantially improves robustness across parameter perturbations, outperforming models trained under narrow conditions. We further validate these findings in both simulated network data and an in vitro dataset from high‑density microelectrode array recordings with patch‑clamp‑verified ground‑truth connections. Models trained on diverse simulated circuits generalize effectively to novel network architectures and the experimental dataset. Together, these results demonstrate that incorporating biologically realistic diversity during training is critical for developing reliable machine-learning tools for large-scale synaptic inference from neural recordings.

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

Repositories

Its files are read in the Code ↔ Paper reader above, with 18 matches between paragraphs and lines of code.

XiaoqianSun0104/OmniCNN_Infer_Connectivity

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 6bcacec66af1098cbbc01c6bd95a0b73512dd4ff, 16 May 2026
Languages: Python (17), Shell (2), Jupyter (2)
Size: 77 files, 21 scripts
Software Heritage: not archived
Found in: the text, “Computational resources and implementation”
Holds: README, environment (environment.yml, requirements.txt), documentation, 2 notebooks
Not found: license file, CITATION.cff, tests, continuous integration
Tools: NumPy (16 files), pandas (14 files), Matplotlib (10 files), seaborn (6 files), NEURON (4 files), PyTorch (4 files), scikit-learn (4 files), SciPy (4 files), SHAP (2 files), statsmodels (1 file), XGBoost (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
22 files

shigerushinomoto/CoNNECT

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: cf768743818419f3d12c9ba7a6e33d9fff3c0ab7, 12 August 2022
Languages: Python (4)
Size: 11 files, 4 scripts
Software Heritage: not archived
Found in: the text, “Comparison with existing methods on HD-MEA data”
Holds: README, license file, environment (modules/setup.py)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (4 files), Keras (1 file), TensorFlow (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
6 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 25 scripts, each with its path and the digest of its content;
  • 18 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

All author-generated code used for network simulations, cross-correlogram computation, model training, and benchmarking analyses is publicly available at: https://github.com/XiaoqianSun0104/OmniCNN_Infer_Connectivity. Simulation data supporting the findings of this study can be regenerated using the provided scripts and parameter configurations. Processed example outputs and scripts required to reproduce the analyses are available within the repository. The in vitro high-density microelectrode array dataset analyzed in this study was obtained from the previously published dataset of Donner et al. (2024) and is publicly available at: https://renkulab.io/projects/christian.donner/deepephys-data.

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, issue, pages, dates, 4 authors, 11 MeSH terms, 3 funders, 47 references.

Cite

This paper

Sun, X., Lu, H., Zeng, C., & Simha, R. (2026). Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains. PLoS computational biology, 22(8), e1014615. https://doi.org/10.1371/journal.pcbi.1014615

BibTeX

@article{sun2026toward,
author = {Sun, Xiaoqian and Lu, Hui and Zeng, Chen and Simha, Rahul},
title = {{Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains}},
journal = {PLoS computational biology},
year = {2026},
month = aug,
volume = {22},
number = {8},
pages = {e1014615},
publisher = {PLOS},
issn = {1553-734X},
doi = {10.1371/journal.pcbi.1014615},
url = {https://doi.org/10.1371/journal.pcbi.1014615},
pmid = {42574490},
pmcid = {PMC13475989}
}

RIS

TY - JOUR
AU - Sun, Xiaoqian
AU - Lu, Hui
AU - Zeng, Chen
AU - Simha, Rahul
TI - Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains
T2 - PLoS computational biology
J2 - PLoS Comput Biol
PY - 2026
DA - 2026/08/10
VL - 22
IS - 8
SP - e1014615
SN - 1553-734X
PB - PLOS
DO - 10.1371/journal.pcbi.1014615
UR - https://doi.org/10.1371/journal.pcbi.1014615
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pcbi.1014615",
"type": "article-journal",
"title": "Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains",
"container-title": "PLoS computational biology",
"author": [
{
"family": "Sun",
"given": "Xiaoqian"
},
{
"family": "Lu",
"given": "Hui"
},
{
"family": "Zeng",
"given": "Chen"
},
{
"family": "Simha",
"given": "Rahul"
}
],
"container-title-short": "PLoS Comput Biol",
"volume": "22",
"issue": "8",
"page": "e1014615",
"DOI": "10.1371/journal.pcbi.1014615",
"PMID": "42574490",
"PMCID": "PMC13475989",
"ISSN": "1553-734X",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pcbi.1014615",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
10
]
]
}
}

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/s41598-026-42580-2 [code]
Comparing effective and functional connectivity.
Journal: Scientific reports
In common: pandas, SciPy, Matplotlib, 1 other tool, 9 references
[2] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: SHAP, XGBoost, Keras, 9 other tools
[3] doi:10.3389/fnsys.2026.1822122 [code]
Convergence-divergence circuits for multimodal integration of innate and learned opponent valences.
Journal: Frontiers in systems neuroscience
In common: NEURON, Keras, TensorFlow, 7 other tools, 2 references
[4] doi:10.3390/s26175327 [code]
Subject Identity Confounds qEEG Emotion Recognition on DEAP and DREAMER.
Journal: Sensors (Basel, Switzerland)
In common: SHAP, XGBoost, Keras, 8 other tools
[5] doi:10.1038/s41598-026-48613-0 [code]
An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging.
Journal: Scientific reports
In common: SHAP, XGBoost, Keras, 8 other tools
[6] doi:10.1126/sciadv.aed3650 [code]
Truthful visualizations for mass spectrometry imaging enable high-spatial-resolution interactive &lt;i&gt;m/z&lt;/i&gt; mapping and exploration.
Journal: Science advances
In common: SHAP, XGBoost, Keras, 7 other tools
[7] doi:10.1038/s41467-026-72057-9 [code]
Sex-specific behavioral feedback modulates sensorimotor processing and drives flexible social behavior.
Journal: Nature communications
In common: Keras, TensorFlow, statsmodels, 7 other tools, 1 reference
[8] doi:10.1038/s41467-026-75700-7 [code]
Gene regulatory innovations from transposable elements in primate cerebellum development.
Journal: Nature communications
In common: SHAP, Keras, TensorFlow, 7 other tools
[9] doi:10.1186/s13059-026-04125-8 [code]
MLMarker: a machine learning framework for tissue inference and biomarker discovery.
Journal: Genome biology
In common: SHAP, XGBoost, statsmodels, 7 other tools
[10] doi:10.1016/j.bpj.2026.06.005 [code]
Curvature-based machine-learning method for automated segmentation of dendritic spines.
Journal: Biophysical journal
In common: SHAP, XGBoost, Keras, 6 other tools

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.