OSCR

Application of Machine Learning Models to Identify Differences in Neural Electrophysiological Properties Across Estrous Cycle Phases.

Code ↔ Paper

4 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 4 matches
  1. [1] § Materials and Methods › Data Preprocessing and Feature Extraction: Passive Properties ↔ Estrous_Classification.ipynb, lines 657–658 · score 0.94 · voltage deflection onset, inward rectification ratio, steady state voltage, rebound voltage, sag amplitude, voltage change
  2. [2] § Materials and Methods › Model Training and Classification ↔ Estrous_Classification.ipynb, lines 702–751 · score 0.91 · gradient boosting classifier, decision tree classifier, logistic regression, random forest classifier, neighbors classifier, score
  3. [3] § Materials and Methods › Data Preprocessing and Feature Extraction: AP Properties ↔ Estrous_Classification.ipynb, lines 134–169 · score 0.85 · find_peaks, spike amplitude, spike peak, distance, prominence, Scipy
  4. [4] § Materials and Methods › Data Preprocessing and Feature Extraction: mEPSC Properties ↔ Estrous_Classification.ipynb, lines 590–591 · score 0.81 · 10–90, inter event, half width, slope, noise, decay

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 · 2,588 lines · 88 KB · MIT · 4 matches

  1. # %% [markdown]
  2. # <a href="https://colab.research.google.com/github/Armaan-Raina/Estrous-Phase-Classification/blob/main/Estrous_Classification.ipynb" target="_parent"><img src="https://colab.research.google.com/assets/colab-badge.svg" alt="Open In Colab"/></a>
  3. # %% [markdown]
  4. # # **Estrous Cycle Phase Classification**
  5. # %% [markdown]
  6. # ## **Preprocessing/ Feature Extraction**
  7. # %% [markdown]
  8. # ### **Importing Packages**
  9. # %%
  10. !pip install pyABF efel
  11. # %%
  12. import pyabf
  13. import numpy
  14. import pandas
  15. import matplotlib.pyplot
  16. import os
  17. import glob
  18. import scipy.signal
  19. import scipy
  20. import seaborn
  21. import efel
  22. import sklearn.manifold.TSNE
  23. import sklearn.preprocessing.StandardScaler
  24. import sklearn.model_selection.train_test_split
  25. import sklearn.ensemble.RandomForestClassifier
  26. import sklearn.ensemble.GradientBoostingClassifier
  27. import sklearn.linear_model.LogisticRegression
  28. import sklearn.neighbors.KNeighborsClassifier
  29. import sklearn.neural_network.MLPClassifier
  30. import sklearn.svm.SVC
  31. import sklearn.tree.DecisionTreeClassifier
  32. import sklearn.decomposition.PCA
  33. import sklearn.metrics.accuracy_score
  34. import sklearn.metrics.classification_report
  35. import statsmodels.formula.api
  36. import statsmodels.api
  37. import matplotlib.colors
  38. import pickle
  39. import statsmodels.stats.multicomp.pairwise_tukeyhsd
  40. import networkx as nx
  41. import itertools
  42. from PIL import Image
  43. # %%
  44. !pip freeze > requirements.txt
  45. # %% [markdown]
  46. # ### **Uploading raw data**
  47. # %%
  48. raw_mEPSC_diestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Diestrus/*.abf')
  49. raw_mEPSC_estrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Estrus/*.abf')
  50. raw_mEPSC_Eproestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Early Proestrus/*.abf')
  51. raw_mEPSC_Lproestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Late Proestrus/*.abf')
  52. # %%
  53. raw_spike_diestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Diestrus Spiking/*.abf')
  54. raw_spike_estrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Estrus Spiking/*.abf')
  55. raw_spike_Eproestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Early Proestrus Spiking/*.abf')
  56. raw_spike_Lproestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Late Proestrus Spiking/*.abf')
  57. # %%
  58. raw_passive_diestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Diestrus Passive/*.abf')
  59. raw_passive_estrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Estrus Passive/*.abf')
  60. raw_passive_Eproestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Early Proestrus Passive/*.abf')
  61. raw_passive_Lproestrus = glob.glob('/content/drive/MyDrive/Estrous Classification Project/Late Proestrus Passive/*.abf')
  62. # %% [markdown]
  63. # ### **Defining functions for preprocessing**
  64. # %%
  65. def read_abf(filename:str):
  66. abf = pyabf.ABF(filename)
  67. time = np.array([], dtype=float)
  68. current = np.array([], dtype=float)
  69. sweep_len = abf.sweepX[-1]
  70. for sweep in range(abf.sweepCount):
  71. abf.setSweep(sweep)
  72. current = np.concatenate((current, abf.sweepY))
  73. time = np.concatenate((time, abf.sweepX + sweep * sweep_len))
  74. return current #No longer returning time
  75. # %%
  76. def filter_signal(current, freq=20000):
  77. raw_signal = current
  78. adjusted_signal = raw_signal - np.median(raw_signal)
  79. b_lowpass, a_lowpass = signal.bessel(4, 1000, 'low', analog=False, norm='phase', fs=freq)
  80. b_notch, a_notch = signal.iirnotch(60.0, 30.0, fs=freq)
  81. b_multiband = signal.convolve(b_lowpass, b_notch)
  82. a_multiband = signal.convolve(a_lowpass, a_notch)
  83. filtered_signal = signal.filtfilt(b_multiband, a_multiband, adjusted_signal)
  84. return filtered_signal
  85. # %% [markdown]
  86. # ### **Preprocessing for Spiking Data**
  87. # %%
  88. abf = pyabf.ABF(raw_spike_diestrus[0])
  89. for sweep in abf.sweepList:
  90. abf.setSweep(sweep)
  91. plt.plot(abf.sweepX, abf.sweepY)
  92. # %%
  93. #abf -> list of np.ndarrays
  94. clean_diestrus = []
  95. clean_estrus = []
  96. clean_Eproestrus = []
  97. clean_Lproestrus = []
  98. for file in raw_spike_diestrus:
  99. clean_diestrus.append(read_abf(file))
  100. for file in raw_spike_estrus:
  101. clean_estrus.append(read_abf(file))
  102. for file in raw_spike_Eproestrus:
  103. clean_Eproestrus.append(read_abf(file))
  104. for file in raw_spike_Lproestrus:
  105. clean_Lproestrus.append(read_abf(file))
  106. # %% [markdown]
  107. # #### **Following code snippets modified based on implemented from [link text](https://)https://spikesandbursts.wordpress.com/2023/08/24/patch-clamp-data-analysis-in-python-bursts/**
  108. # %%
  109. def spike_detection(sig, fs):
  110. thresh_min = -25
  111. thresh_prominence = 15
  112. thresh_min_width = 0.5 * (fs/1000)
  113. distance_min = 1 * (fs/1000)
  114. peaks, peaks_dict = scipy.signal.find_peaks(sig,
  115. height=thresh_min,
  116. threshold=thresh_min,
  117. distance=distance_min,
  118. prominence=thresh_prominence,
  119. width=thresh_min_width,
  120. wlen=None, # Window length to calculate prominence
  121. rel_height=0.5, # Relative height at which the peak width is measured
  122. plateau_size=None)
  123. spikes_table = pd.DataFrame(columns = ['spike', 'spike_index', 'spike_time',
  124. 'inst_freq', 'isi_s',
  125. 'width', 'rise_half_ms', 'decay_half_ms',
  126. 'spike_peak', 'spike_amplitude'])
  127. spikes_table.spike = np.arange(1, len(peaks) + 1)
  128. spikes_table.spike_index = peaks
  129. spikes_table.spike_time = peaks / fs # Divided by fs to get s
  130. spikes_table.isi_s = np.diff(peaks, axis=0, prepend=peaks[0]) / fs
  131. spikes_table.inst_freq = 1 / spikes_table.isi_s
  132. spikes_table.width = peaks_dict['widths']/(fs/1000) # Width (ms) at half-height
  133. spikes_table.rise_half_ms = (peaks - peaks_dict['left_ips'])/(fs/1000)
  134. spikes_table.decay_half_ms = (peaks_dict['right_ips'] - peaks)/(fs/1000)
  135. spikes_table.spike_peak = peaks_dict['peak_heights'] # height parameter is needed
  136. spikes_table.spike_amplitude = peaks_dict['prominences'] # prominence parameter is needed
  137. return spikes_table
  138. # %%
  139. abf = pyabf.ABF(raw_spike_diestrus[0])
  140. print(abf.dataPointsPerMs*1000)
  141. # %%
  142. ##Testing on one file
  143. df = spike_detection(clean_diestrus[0][1], 20000)
  144. # %%
  145. mod_df = df.drop(["spike", "spike_index", "spike_time"], axis=1)
  146. # %%
  147. mod_df = mod_df.replace([-np.inf, np.inf], np.nan).dropna()
  148. # %%
  149. plt.matshow(mod_df.corr())
  150. cb = plt.colorbar()
  151. cb.ax.tick_params(labelsize=14)
  152. plt.title('Correlation Matrix', fontsize=16);
  153. # %%
  154. #Estimating parameters based on diestrus trace
  155. hist_data = mod_df['isi_s']
  156. hist_stats = pd.DataFrame()
  157. bin_size = 10 #ms
  158. isi_range = np.ptp(hist_data)
  159. bins = int((isi_range * 1000 / bin_size) + 0.5) # Round to the nearest integer
  160. hist = np.histogram(hist_data, bins=bins)
  161. hist_counts = hist[0]
  162. hist_bins = hist[1]
  163. # Cumulative moving average
  164. cum = np.cumsum(hist_counts) # Cumulative sum
  165. cma = cum / np.arange(1, len(cum) + 1)
  166. # Calculate peaks and valleys of the cma
  167. cma_peaks_indexes = scipy.signal.argrelextrema(cma, np.greater)
  168. cma_valleys_indexes = scipy.signal.argrelextrema(cma, np.less)
  169. # Select the peak you're interested in
  170. peak_index = cma_peaks_indexes[0][0] # Change second number to select the peak
  171. alpha = cma[peak_index] * 0.5 # Half-peak, adapt the value to your threshold criterion
  172. # Calculate cma_threshold_index relative to the selected cma_peak
  173. cma_threshold = (np.argmin(cma[peak_index:] >= alpha) + peak_index) * bin_size/1000
  174. # Dataframe with histogram statistics
  175. length = len(hist_stats)
  176. hist_stats.loc[length, 'mean_isi'] = np.mean(hist_data)
  177. hist_stats.loc[length, 'median_isi'] = np.median(hist_data)
  178. hist_stats.loc[length, 'kurtosis'] = scipy.stats.kurtosis(hist_counts)
  179. hist_stats.loc[length, 'skewness'] = scipy.stats.skew(hist_counts, bias=True)
  180. hist_stats.loc[length, 'cma_threshold'] = cma_threshold
  181. hist_stats.loc[length, 'cma_valley_time'] = cma_valleys_indexes[0][1] * bin_size/1000 # Change peak index as needed
  182. hist_stats.loc[length, 'cma_peak_time'] = cma_peaks_indexes[0][0] * bin_size/1000 # Change peak index as needed
  183. # Plot ISI histogram
  184. fig, ax = plt.subplots(figsize=(8, 4))
  185. ax.set_title("ISI histogram")
  186. ax.hist(hist_data, bins=bins, alpha=0.6)
  187. # Plot CMA
  188. cma_x = np.linspace(np.min(hist_bins), np.max(hist_bins), bins)
  189. ax.plot(cma_x, cma)
  190. # Plot CMA threshold line
  191. ax.axvline(cma_threshold, linestyle="dotted", color="gray")
  192. # Plot CMA valleys
  193. ax.plot(cma_x[cma_valleys_indexes], cma[cma_valleys_indexes], 'ko')
  194. ax.plot(cma_x[cma_peaks_indexes], cma[cma_peaks_indexes], 'mo')
  195. # ax.set_xscale('log') # Logarithmic scale may be easier to set the threshold
  196. ax.set_xlabel("Time bins (s)")
  197. ax.set_ylabel("Count")
  198. ax.set_xlim([0,1])
  199. # Show graph and table
  200. plt.show()
  201. hist_stats
  202. # %%
  203. def burst_detection(df, spike_times, spike_amplitudes, spike_peaks, n_spikes, max_isi, min_ibi=None):
  204. df = df.sort_values(by=spike_times)
  205. df['burst'] = np.nan
  206. burst_num = 0
  207. burst_start = None
  208. last_spike = None
  209. for i, row in df.iterrows(): # Loop through DataFrame rows
  210. spike = row[spike_times] # Extract the spike position
  211. if burst_start is None: # It checks if it is the first spike
  212. burst_start = spike # It marks the current spike position as the start of a burst
  213. last_spike = spike # Update the last_spike position to the current spike position
  214. df.at[i, 'burst'] = burst_num # Assign burst number
  215. elif spike - last_spike <= max_isi: # It checks if the current spike is within max isi
  216. df.at[i, 'burst'] = burst_num
  217. last_spike = spike
  218. elif spike - last_spike > min_ibi: # It checks if the interburst interval has been reached
  219. burst_num += 1
  220. burst_start = spike
  221. last_spike = spike
  222. df.at[i, 'burst'] = burst_num
  223. # Filter bursts with less than min_spikes
  224. df = df[df.groupby('burst')[spike_times].transform('count') >= n_spikes]
  225. bursts = df.groupby('burst')[spike_times].agg(['min', 'max', 'count'])
  226. bursts.columns = ['burst_start', 'burst_end', 'spikes_in_bursts']
  227. bursts['burst_length'] = bursts['burst_end'] - bursts['burst_start']
  228. bursts['avg_spike_amplitude'] = df.groupby('burst')[spike_amplitudes].mean()
  229. bursts['avg_spike_peaks'] = df.groupby('burst')[spike_peaks].mean()
  230. bursts['spikes_frequency'] = bursts['spikes_in_bursts'] / bursts['burst_length']
  231. bursts = bursts.reset_index()
  232. bursts['burst_number'] = bursts.index + 1
  233. return bursts[['burst_number', 'burst_start', 'burst_end',
  234. 'burst_length', 'spikes_in_bursts', 'avg_spike_amplitude',
  235. 'avg_spike_peaks', 'spikes_frequency']]
  236. # %%
  237. diestrus_spiking = pd.DataFrame()
  238. estrus_spiking = pd.DataFrame()
  239. Eproestrus_spiking = pd.DataFrame()
  240. Lproestrus_spiking = pd.DataFrame()
  241. for trace in clean_diestrus:
  242. spikes_table = spike_detection(trace[1], 20000)
  243. diestrus_spiking = pd.concat([diestrus_spiking, burst_detection(spikes_table, 'spike_time', 'spike_amplitude', 'spike_peak', n_spikes = 2, max_isi=0.5, min_ibi=0.1)])
  244. for trace in clean_estrus:
  245. spikes_table = spike_detection(trace[1], 20000)
  246. estrus_spiking = pd.concat([estrus_spiking, burst_detection(spikes_table, 'spike_time', 'spike_amplitude', 'spike_peak', n_spikes = 2, max_isi=0.5, min_ibi=0.1)])
  247. for trace in clean_Eproestrus:
  248. spikes_table = spike_detection(trace[1], 20000)
  249. Eproestrus_spiking = pd.concat([Eproestrus_spiking, burst_detection(spikes_table, 'spike_time', 'spike_amplitude', 'spike_peak', n_spikes = 2, max_isi=0.5, min_ibi=0.1)])
  250. for trace in clean_Lproestrus:
  251. spikes_table = spike_detection(trace[1], 20000)
  252. Lproestrus_spiking = pd.concat([Lproestrus_spiking, burst_detection(spikes_table, 'spike_time', 'spike_amplitude', 'spike_peak', n_spikes = 2, max_isi=0.5, min_ibi=0.1)])
  253. # %%
  254. diestrus_spiking.to_csv("/content/drive/MyDrive/Estrous Classification Project/diestrus_spiking_feats")
  255. estrus_spiking.to_csv("/content/drive/MyDrive/Estrous Classification Project/estrus_spiking_feats")
  256. Eproestrus_spiking.to_csv("/content/drive/MyDrive/Estrous Classification Project/Earlyproestrus_spiking_feats")
  257. Lproestrus_spiking.to_csv("/content/drive/MyDrive/Estrous Classification Project/Lateproestrus_spiking_feats")
  258. # %% [markdown]
  259. # ### **Preprocessing for Passive Data**
  260. # %%
  261. #Exploration of abf file to determine constants
  262. abf = pyabf.ABF(raw_passive_Lproestrus[0])
  263. sweeps = abf.sweepList
  264. for sweep in sweeps:
  265. abf.setSweep(sweep)
  266. plt.plot(abf.sweepX, abf.sweepC)
  267. abf.setSweep(4)
  268. ##Finding the start and end indices of the stimulation
  269. stim_start = 0
  270. stim_end = 0
  271. idx = 0
  272. startFlag = False
  273. val = abf.sweepC
  274. while val[idx] == 0:
  275. idx+=1
  276. stim_start += 1
  277. print(stim_start)
  278. for i in range(stim_start, len(val)+1):
  279. if val[i] == 0:
  280. stim_end = i
  281. break
  282. print(stim_end)
  283. print(abf.sweepX[9562])
  284. print(abf.sweepX[stim_end])
  285. plt.plot(abf.sweepX, abf.sweepC)
  286. plt.axvline(0.4781, color='r')
  287. plt.axvline(1.0781, color='g')
  288. print(abf.sweepEpochs)
  289. # %%
  290. ##CONSTANTS - determined in last block
  291. stim_start = [4781 / abf.dataRate * 1000] # ms
  292. stim_end = [10781 / abf.dataRate * 1000] # ms
  293. ir_current_start = [i for i in range(-140, 10, 10)]
  294. ir_current_stop = 0
  295. curr_start_idx = 4781
  296. curr_stop_idx = 10781
  297. # %%
  298. ir_current_start, abf.dataPointsPerMs * 1000 #20 kHz sampling rate
  299. # %%
  300. def init_features(filename):
  301. abf = pyabf.ABF(filename)
  302. sweeps = abf.sweepList
  303. table = pd.DataFrame(columns=[
  304. 'rmp_mV',
  305. 'steady_state_voltage_mV',
  306. 'vmin_mV',
  307. 'vdeflection_begin_mV',
  308. 'vrebound_mV',
  309. 'voltage_delta_mV',
  310. 'current_pA',
  311. 'input_resistance_Gohm',
  312. 'time_constant_ms',
  313. 'capacitance_pF',
  314. 'sag_amplitude_mV',
  315. 'sag_ratio1',
  316. 'sag_ratio2',
  317. 'sweep',
  318. 'inward_rectification_ratio'
  319. ])
  320. inwardRec = []
  321. inwardRecRatio = np.nan
  322. for sweep in sweeps:
  323. abf.setSweep(sweep)
  324. stim_start_index = 4781
  325. stim_end_index = 10781
  326. stim_start = [abf.sweepX[stim_start_index] * 1000]
  327. stim_end = [abf.sweepX[stim_end_index] * 1000]
  328. current_pA = np.mean(abf.sweepC[stim_start_index:stim_end_index])
  329. trace = {
  330. 'T': abf.sweepX * 1000,
  331. 'V': abf.sweepY,
  332. 'stim_start': stim_start,
  333. 'stim_end': stim_end
  334. }
  335. feature_values = efel.get_feature_values(
  336. [trace],
  337. [
  338. 'voltage_base', 'steady_state_voltage_stimend',
  339. 'minimum_voltage', 'voltage_deflection_begin',
  340. 'voltage_deflection', 'voltage_deflection_vb_ssse',
  341. 'decay_time_constant_after_stim', 'sag_amplitude',
  342. 'sag_ratio1', 'sag_ratio2', 'voltage_after_stim'
  343. ]
  344. )[0]
  345. def safe_get(key, idx=0): #preventing nulls so it can be handled later
  346. return feature_values[key][idx] if (key in feature_values and feature_values[key] is not None) else np.nan
  347. base_v = safe_get('voltage_base')
  348. steady_v = safe_get('steady_state_voltage_stimend')
  349. vmin = safe_get('minimum_voltage')
  350. vdeflection_begin = safe_get('voltage_deflection_begin')
  351. vrebound = safe_get('voltage_after_stim')
  352. sag_amp = safe_get('sag_amplitude')
  353. sag_ratio1 = safe_get('sag_ratio1')
  354. sag_ratio2 = safe_get('sag_ratio2')
  355. tau = safe_get('decay_time_constant_after_stim')
  356. delta_v = steady_v - base_v if not np.isnan(steady_v) and not np.isnan(base_v) else np.nan
  357. Rin = delta_v / current_pA if current_pA != 0 else np.nan
  358. capacitance = tau / Rin if Rin and not np.isnan(tau) and Rin != 0 else np.nan
  359. # record specific sweeps for inward rectification ratio
  360. if sweep == 1:
  361. inwardRec.append(Rin)
  362. if sweep == 14:
  363. inwardRec.append(Rin)
  364. table.loc[len(table)] = [
  365. base_v, steady_v, vmin, vdeflection_begin, vrebound, delta_v,
  366. current_pA, Rin, tau, capacitance, sag_amp, sag_ratio1, sag_ratio2, sweep, np.nan
  367. ]
  368. if len(inwardRec) == 2 and not any(np.isnan(inwardRec)):
  369. inwardRecRatio = inwardRec[1] / inwardRec[0]
  370. table.iloc[-len(sweeps):, table.columns.get_loc('inward_rectification_ratio')] = inwardRecRatio
  371. return table
  372. # %%
  373. diestrus_passive_df = pd.DataFrame()
  374. estrus_passive_df = pd.DataFrame()
  375. Eproestrus_passive_df = pd.DataFrame()
  376. Lproestrus_passive_df = pd.DataFrame()
  377. for filename in raw_passive_diestrus:
  378. df = init_features(filename)
  379. diestrus_passive_df = pd.concat([df, diestrus_passive_df])
  380. diestrus_passive_df['label'] = 0
  381. for filename in raw_passive_estrus:
  382. df = init_features(filename)
  383. estrus_passive_df = pd.concat([df, estrus_passive_df])
  384. estrus_passive_df['label'] = 1
  385. for filename in raw_passive_Eproestrus:
  386. df = init_features(filename)
  387. Eproestrus_passive_df = pd.concat([df, Eproestrus_passive_df])
  388. Eproestrus_passive_df['label'] = 2
  389. for filename in raw_passive_Lproestrus:
  390. df = init_features(filename)
  391. Lproestrus_passive_df = pd.concat([df, Lproestrus_passive_df])
  392. Lproestrus_passive_df['label'] = 3
  393. # %%
  394. diestrus_passive_df.head(n = 45)
  395. # %%
  396. diestrus_passive_df = diestrus_passive_df.dropna()
  397. estrus_passive_df = estrus_passive_df.dropna()
  398. Eproestrus_passive_df = Eproestrus_passive_df.dropna()
  399. Lproestrus_passive_df = Lproestrus_passive_df.dropna()
  400. # %%
  401. diestrus_passive_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/diestrus_passive_feats")
  402. estrus_passive_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/estrus_passive_feats")
  403. Eproestrus_passive_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/Eproestrus_passive_feats")
  404. Lproestrus_passive_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/Lproestrus_passive_feats")
  405. # %% [markdown]
  406. # ## **Feature Loading**
  407. #
  408. # %% [markdown]
  409. # ### **mEPSC (MiniAnalysis) Features**
  410. # %%
  411. mEPSC_diestrus = pd.read_excel('/content/drive/MyDrive/Spreadsheets/diestrus.xlsx')
  412. mEPSC_estrus = pd.read_excel('/content/drive/MyDrive/Spreadsheets/estrus.xlsx')
  413. mEPSC_Eproestrus = pd.read_excel('/content/drive/MyDrive/Spreadsheets/earlyproestrus.xlsx')
  414. mEPSC_Lproestrus = pd.read_excel('/content/drive/MyDrive/Spreadsheets/lateproestrus.xlsx')
  415. # %%
  416. #Filtering out extraneous events likely due to experimental setup or other circumstances
  417. mEPSC_diestrus = mEPSC_diestrus[(mEPSC_diestrus['Amplitude'] < 50) & (mEPSC_diestrus['Baseline'] > 0)]
  418. mEPSC_estrus = mEPSC_estrus[(mEPSC_estrus['Amplitude'] < 50) & (mEPSC_estrus['Baseline'] > 0)]
  419. mEPSC_Eproestrus = mEPSC_Eproestrus[(mEPSC_Eproestrus['Amplitude'] < 50) & (mEPSC_Eproestrus['Baseline'] > 0)]
  420. mEPSC_Lproestrus = mEPSC_Lproestrus[(mEPSC_Lproestrus['Amplitude'] < 50) & (mEPSC_Lproestrus['Baseline'] > 0)]
  421. # %%
  422. def compute_iei(df: pd.DataFrame) -> pd.DataFrame:
  423. df["Inter-Event Interval (ms)"] = df["Time (ms)"].diff()
  424. return df
  425. # %%
  426. #Simple IEI computation
  427. mEPSC_diestrus = compute_iei(mEPSC_diestrus)
  428. mEPSC_estrus = compute_iei(mEPSC_estrus)
  429. mEPSC_Eproestrus = compute_iei(mEPSC_Eproestrus)
  430. mEPSC_Lproestrus = compute_iei(mEPSC_Lproestrus)
  431. # %%
  432. mEPSC_diestrus.iloc[0, 18] = mEPSC_diestrus.iloc[0, 1]
  433. mEPSC_estrus.iloc[0, 18] = mEPSC_estrus.iloc[0, 1]
  434. mEPSC_Eproestrus.iloc[0, 18] = mEPSC_Eproestrus.iloc[0, 1]
  435. mEPSC_Lproestrus.iloc[0, 18] = mEPSC_Lproestrus.iloc[0, 1]
  436. # %%
  437. mEPSC_min_count = len(min([mEPSC_diestrus, mEPSC_estrus, mEPSC_Eproestrus, mEPSC_Lproestrus], key=len))
  438. # %%
  439. mEPSC_diestrus['label'] = 0
  440. mEPSC_estrus['label'] = 1
  441. mEPSC_Eproestrus['label'] = 2
  442. mEPSC_Lproestrus['label'] = 3
  443. # %%
  444. mEPSC_diestrus = mEPSC_diestrus.dropna()
  445. mEPSC_estrus = mEPSC_estrus.dropna()
  446. mEPSC_Eproestrus = mEPSC_Eproestrus.dropna()
  447. mEPSC_Lproestrus = mEPSC_Lproestrus.dropna()
  448. # %%
  449. mEPSC = pd.concat([mEPSC_diestrus, mEPSC_estrus, mEPSC_Eproestrus, mEPSC_Lproestrus])
  450. # %%
  451. #Reducing count of each class to match for testing to determine if class imbalance significantly impacts model accuracy
  452. mEPSC_even = pd.concat([mEPSC_diestrus.head(mEPSC_min_count), mEPSC_estrus.head(mEPSC_min_count), mEPSC_Eproestrus.head(mEPSC_min_count), mEPSC_Lproestrus.head(mEPSC_min_count)])
  453. # %%
  454. mEPSC_even.shape
  455. # %%
  456. mEPSC_even = mEPSC_even.drop(["Obs Num", "Time (ms)", "Group", "Channel", "Peak Dir", "Burst#", "BurstE#", "Rel Time"], axis=1)
  457. # %%
  458. mEPSC = mEPSC.drop(["Obs Num", "Time (ms)", "Group", "Channel", "Peak Dir", "Burst#", "BurstE#", "Rel Time"], axis=1) #Dropping insignificant features
  459. # %%
  460. mEPSC_names = ["Amplitude", "Rise (ms)", "Decay (ms)", "Area", "Baseline", "Noise", "10-90Rise", "HalfWidth", "Rise50", "10-90Slope", "Inter-Event Interval (ms)"]
  461. # %% [markdown]
  462. # ### **AP Features**
  463. # %%
  464. spiking_diestrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/diestrus_spiking_feats")
  465. spiking_estrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/estrus_spiking_feats")
  466. spiking_Eproestrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/Earlyproestrus_spiking_feats")
  467. spiking_Lproestrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/Lateproestrus_spiking_feats")
  468. # %%
  469. spiking_min_count = len(min([spiking_diestrus, spiking_estrus, spiking_Eproestrus, spiking_Lproestrus], key=len))
  470. # %%
  471. spiking_diestrus.shape, spiking_estrus.shape, spiking_Eproestrus.shape, spiking_Lproestrus.shape
  472. # %%
  473. spiking_diestrus['label'] = 0
  474. spiking_estrus['label'] = 1
  475. spiking_Eproestrus['label'] = 2
  476. spiking_Lproestrus['label'] = 3
  477. # %%
  478. spiking = pd.concat([spiking_diestrus, spiking_estrus, spiking_Eproestrus, spiking_Lproestrus])
  479. # %%
  480. spiking = spiking.drop(["Unnamed: 0", "burst_number", "burst_start", "burst_end"], axis=1) #Dropping insignificant features
  481. # %%
  482. spiking_even = pd.concat([spiking_diestrus.head(spiking_min_count), spiking_estrus.head(spiking_min_count), spiking_Eproestrus.head(spiking_min_count), spiking_Lproestrus.head(spiking_min_count)])
  483. # %%
  484. spiking_even = spiking_even.drop(["Unnamed: 0", "burst_number", "burst_start", "burst_end"], axis=1) #Dropping insignificant features
  485. # %%
  486. spiking_names = ["burst_length", "spikes_in_bursts", "avg_spike_amplitude", "avg_spike_peaks", "spikes_frequency"]
  487. # %% [markdown]
  488. # ### **Passive Features**
  489. # %%
  490. passive_diestrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/diestrus_passive_feats")
  491. passive_estrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/estrus_passive_feats")
  492. passive_Eproestrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/Eproestrus_passive_feats")
  493. passive_Lproestrus = pd.read_csv("/content/drive/MyDrive/Estrous Classification Project/Lproestrus_passive_feats")
  494. #These already have labels saved into them
  495. # %%
  496. passive_min_count = len(min([passive_diestrus, passive_estrus, passive_Eproestrus, passive_Lproestrus], key=len))
  497. # %%
  498. passive = pd.concat([passive_diestrus, passive_estrus, passive_Eproestrus, passive_Lproestrus])
  499. # %%
  500. passive = passive.drop(["Unnamed: 0", "sag_ratio1", "sweep"], axis=1)
  501. # %%
  502. passive_even = pd.concat([passive_diestrus.head(passive_min_count), passive_estrus.head(passive_min_count), passive_Eproestrus.head(passive_min_count), passive_Lproestrus.head(passive_min_count)])
  503. # %%
  504. passive_even = passive_even.drop(["Unnamed: 0", "sag_ratio1", "sweep"], axis=1)
  505. # %%
  506. passive_names = ["rmp_mV", "steady_state_voltage_mV", "vmin_mV", "vdeflection_begin_mV", "vrebound_mV", "voltage_delta_mV", "current_pA", "input_resistance_Gohm", "time_constant_ms", "capacitance_pF", "sag_amplitude_mV", "sag_ratio1", "sag_ratio2", "sweep", "inward_rectification_ratio"]
  507. # %%
  508. passive_names_pretty = ["RMP (mV)", "Steady State Voltage (mV)", "Min. Voltage (mV)", "Voltage Deflection Onset (mV)", "Rebound Voltage (mV)", "Voltage Change (mV)", "Current (pA)", "Input Resistance (GOhms)", "time_constant_ms", "capacitance_pF", "sag_amplitude_mV","sag_ratio2", "inward_rectification_ratio"]
  509. # %% [markdown]
  510. # ## **Model Training**
  511. # %% [markdown]
  512. # #### **Feature Preparation**
  513. # %%
  514. X_mEPSC = mEPSC.drop('label', axis=1)
  515. y_mEPSC = mEPSC['label']
  516. # %%
  517. X_spiking = spiking.drop('label', axis=1)
  518. y_spiking = spiking['label']
  519. # %%
  520. X_passive = passive.drop('label', axis=1)
  521. y_passive = passive['label']
  522. # %%
  523. X_even_mEPSC = mEPSC_even.drop('label', axis=1)
  524. y_even_mEPSC = mEPSC_even['label']
  525. X_even_spiking = spiking_even.drop('label', axis=1)
  526. y_even_spiking = spiking_even['label']
  527. X_even_passive = passive_even.drop('label', axis=1)
  528. y_even_passive = passive_even['label']
  529. # %%
  530. sc = StandardScaler()
  531. X_mEPSC = sc.fit_transform(X_mEPSC)
  532. X_spiking = sc.fit_transform(X_spiking)
  533. X_passive = sc.fit_transform(X_passive)
  534. X_even_mEPSC = sc.fit_transform(X_even_mEPSC)
  535. X_even_spiking = sc.fit_transform(X_even_spiking)
  536. X_even_passive = sc.fit_transform(X_even_passive)
  537. # %% [markdown]
  538. # #### **mEPSC First**
  539. # %%
  540. mEPSC_classification_reports = {}
  541. mEPSC_feature_importances = {}
  542. # %%
  543. SEEDS = [42, 7, 13, 21, 99]
  544. mEPSC_seed_accuracies = {model_name: [] for model_name in [
  545. 'RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC'
  546. ]}
  547. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  548. # Store ALL predictions and true labels for ALL seeds' runs
  549. all_preds_mEPSC = {name: [] for name in model_names}
  550. all_y_tests_mEPSC = {name: [] for name in model_names}
  551. for seed in SEEDS:
  552. X_train_mEPSC, X_test_mEPSC, y_train_mEPSC, y_test_mEPSC = train_test_split(
  553. X_mEPSC, y_mEPSC, test_size=0.25, shuffle=True, random_state=seed
  554. )
  555. mEPSC_models = [
  556. RandomForestClassifier(random_state=seed),
  557. GradientBoostingClassifier(random_state=seed),
  558. LogisticRegression(max_iter=1000, random_state=seed),
  559. KNeighborsClassifier(), # No random_state
  560. MLPClassifier(max_iter=1000, random_state=seed),
  561. SVC(random_state=seed),
  562. DecisionTreeClassifier(random_state=seed),
  563. ]
  564. for name, model in zip(model_names, mEPSC_models):
  565. model.fit(X_train_mEPSC, y_train_mEPSC)
  566. pred = model.predict(X_test_mEPSC)
  567. acc = accuracy_score(y_test_mEPSC, pred)
  568. mEPSC_seed_accuracies[name].append(acc)
  569. # Store for all seeds
  570. all_preds_mEPSC[name].extend(pred)
  571. all_y_tests_mEPSC[name].extend(y_test_mEPSC)
  572. # Feature importances for each seed of the RFC
  573. if name == 'RFC':
  574. mEPSC_feature_importances[seed] = model.feature_importances_
  575. # Average across seeds
  576. mEPSC_accuracies = [np.mean(mEPSC_seed_accuracies[name]) for name in model_names]
  577. mEPSC_std = [np.std(mEPSC_seed_accuracies[name]) for name in model_names]
  578. for name in model_names:
  579. if name in all_preds_mEPSC and name in all_y_tests_mEPSC:
  580. rep = classification_report(all_y_tests_mEPSC[name], all_preds_mEPSC[name], zero_division=0, output_dict=True)
  581. mEPSC_classification_reports[name] = rep
  582. # %% [markdown]
  583. # #### **Spiking Next**
  584. # %%
  585. spiking_classification_reports = {}
  586. spiking_feature_importances = {}
  587. # %%
  588. SEEDS = [42, 7, 13, 21, 99] # Same seeds
  589. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  590. spiking_seed_accuracies = {name: [] for name in model_names}
  591. all_preds_spiking = {name: [] for name in model_names}
  592. all_y_tests_spiking = {name: [] for name in model_names}
  593. for seed in SEEDS:
  594. X_train_spiking, X_test_spiking, y_train_spiking, y_test_spiking = train_test_split(
  595. X_spiking, y_spiking, test_size=0.25, shuffle=True, random_state=seed
  596. )
  597. spiking_models = [
  598. RandomForestClassifier(random_state=seed),
  599. GradientBoostingClassifier(random_state=seed),
  600. LogisticRegression(max_iter=1000, random_state=seed),
  601. KNeighborsClassifier(), # No random_state
  602. MLPClassifier(max_iter=1000, random_state=seed),
  603. SVC(random_state=seed),
  604. DecisionTreeClassifier(random_state=seed),
  605. ]
  606. for name, model in zip(model_names, spiking_models):
  607. model.fit(X_train_spiking, y_train_spiking)
  608. pred = model.predict(X_test_spiking)
  609. acc = accuracy_score(y_test_spiking, pred)
  610. spiking_seed_accuracies[name].append(acc)
  611. # Store for all seeds
  612. all_preds_spiking[name].extend(pred)
  613. all_y_tests_spiking[name].extend(y_test_spiking)
  614. if name == 'RFC':
  615. spiking_feature_importances[seed] = model.feature_importances_
  616. # Average across seeds
  617. spiking_accuracies = [np.mean(spiking_seed_accuracies[name]) for name in model_names]
  618. spiking_std = [np.std(spiking_seed_accuracies[name]) for name in model_names]
  619. for name in model_names:
  620. if name in all_preds_spiking and name in all_y_tests_spiking:
  621. rep = classification_report(all_y_tests_spiking[name], all_preds_spiking[name], zero_division=0, output_dict=True)
  622. spiking_classification_reports[name] = rep
  623. # %% [markdown]
  624. # #### **Passive Last**
  625. # %%
  626. passive_classification_reports = {}
  627. passive_feature_importances = {}
  628. # %%
  629. SEEDS = [42, 7, 13, 21, 99]
  630. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  631. passive_seed_accuracies = {name: [] for name in model_names}
  632. # Store ALL predictions and true labels for ALL seeds' runs
  633. all_preds_passive = {name: [] for name in model_names}
  634. all_y_tests_passive = {name: [] for name in model_names}
  635. for seed in SEEDS:
  636. X_train_passive, X_test_passive, y_train_passive, y_test_passive = train_test_split(
  637. X_passive, y_passive, test_size=0.25, shuffle=True, random_state=seed
  638. )
  639. passive_models = [
  640. RandomForestClassifier(random_state=seed),
  641. GradientBoostingClassifier(random_state=seed),
  642. LogisticRegression(max_iter=1000, random_state=seed),
  643. KNeighborsClassifier(),
  644. MLPClassifier(max_iter=1000, random_state=seed),
  645. SVC(random_state=seed),
  646. DecisionTreeClassifier(random_state=seed),
  647. ]
  648. for name, model in zip(model_names, passive_models):
  649. model.fit(X_train_passive, y_train_passive)
  650. pred = model.predict(X_test_passive)
  651. acc = accuracy_score(y_test_passive, pred)
  652. passive_seed_accuracies[name].append(acc)
  653. # Store for all seeds
  654. all_preds_passive[name].extend(pred)
  655. all_y_tests_passive[name].extend(y_test_passive)
  656. if name == "RFC":
  657. passive_feature_importances[seed] = model.feature_importances_
  658. # Average across seeds
  659. passive_accuracies = [np.mean(passive_seed_accuracies[name]) for name in model_names]
  660. passive_std = [np.std(passive_seed_accuracies[name]) for name in model_names]
  661. for name in model_names:
  662. if name in all_preds_passive and name in all_y_tests_passive:
  663. rep = classification_report(all_y_tests_passive[name], all_preds_passive[name], zero_division=0, output_dict=True)
  664. passive_classification_reports[name] = rep
  665. # %% [markdown]
  666. # #### **Now again, addressing data inequality**
  667. # %%
  668. ###mEPSC###
  669. SEEDS = [42, 7, 13, 21, 99]
  670. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  671. mEPSC_even_seed_accuracies = {name: [] for name in model_names}
  672. all_preds_mEPSC_even = {name: [] for name in model_names}
  673. all_y_tests_mEPSC_even = {name: [] for name in model_names}
  674. for seed in SEEDS:
  675. X_train_mEPSC_even, X_test_mEPSC_even, y_train_mEPSC_even, y_test_mEPSC_even = train_test_split(
  676. X_even_mEPSC, y_even_mEPSC, test_size=0.25, shuffle=True, random_state=seed
  677. )
  678. mEPSC_even_models = [
  679. RandomForestClassifier(random_state=seed),
  680. GradientBoostingClassifier(random_state=seed),
  681. LogisticRegression(max_iter=1000, random_state=seed),
  682. KNeighborsClassifier(),
  683. MLPClassifier(max_iter=1000, random_state=seed),
  684. SVC(random_state=seed),
  685. DecisionTreeClassifier(random_state=seed),
  686. ]
  687. for name, model in zip(model_names, mEPSC_even_models):
  688. model.fit(X_train_mEPSC_even, y_train_mEPSC_even)
  689. pred = model.predict(X_test_mEPSC_even)
  690. acc = accuracy_score(y_test_mEPSC_even, pred)
  691. mEPSC_even_seed_accuracies[name].append(acc)
  692. # Store for all seeds
  693. all_preds_mEPSC_even[name].extend(pred)
  694. all_y_tests_mEPSC_even[name].extend(y_test_mEPSC_even)
  695. # Average across seeds
  696. mEPSC_even_accuracies = [np.mean(mEPSC_even_seed_accuracies[name]) for name in model_names]
  697. mEPSC_even_std = [np.std(mEPSC_even_seed_accuracies[name]) for name in model_names]
  698. # %%
  699. ###Spiking###
  700. SEEDS = [42, 7, 13, 21, 99]
  701. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  702. spiking_even_seed_accuracies = {name: [] for name in model_names}
  703. # Store ALL predictions and true labels for ALL seeds' runs
  704. all_preds_spiking_even = {name: [] for name in model_names}
  705. all_y_tests_spiking_even = {name: [] for name in model_names}
  706. for seed in SEEDS:
  707. X_train_spiking_even, X_test_spiking_even, y_train_spiking_even, y_test_spiking_even = train_test_split(
  708. X_even_spiking, y_even_spiking, test_size=0.25, shuffle=True, random_state=seed
  709. )
  710. spiking_even_models = [
  711. RandomForestClassifier(random_state=seed),
  712. GradientBoostingClassifier(random_state=seed),
  713. LogisticRegression(max_iter=1000, random_state=seed),
  714. KNeighborsClassifier(),
  715. MLPClassifier(max_iter=1000, random_state=seed),
  716. SVC(random_state=seed),
  717. DecisionTreeClassifier(random_state=seed),
  718. ]
  719. for name, model in zip(model_names, spiking_even_models):
  720. model.fit(X_train_spiking_even, y_train_spiking_even)
  721. pred = model.predict(X_test_spiking_even)
  722. acc = accuracy_score(y_test_spiking_even, pred)
  723. spiking_even_seed_accuracies[name].append(acc)
  724. # Store for all seeds
  725. all_preds_spiking_even[name].extend(pred)
  726. all_y_tests_spiking_even[name].extend(y_test_spiking_even)
  727. # Average across seeds
  728. spiking_even_accuracies = [np.mean(spiking_even_seed_accuracies[name]) for name in model_names]
  729. spiking_even_std = [np.std(spiking_even_seed_accuracies[name]) for name in model_names]
  730. # %%
  731. ###Passive###
  732. SEEDS = [42, 7, 13, 21, 99]
  733. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  734. passive_even_seed_accuracies = {name: [] for name in model_names}
  735. all_preds_passive_even = {name: [] for name in model_names}
  736. all_y_tests_passive_even = {name: [] for name in model_names}
  737. for seed in SEEDS:
  738. X_train_passive_even, X_test_passive_even, y_train_passive_even, y_test_passive_even = train_test_split(
  739. X_even_passive, y_even_passive, test_size=0.25, shuffle=True, random_state=seed
  740. )
  741. passive_even_models = [
  742. RandomForestClassifier(random_state=seed),
  743. GradientBoostingClassifier(random_state=seed),
  744. LogisticRegression(max_iter=1000, random_state=seed),
  745. KNeighborsClassifier(),
  746. MLPClassifier(max_iter=1000, random_state=seed),
  747. SVC(random_state=seed),
  748. DecisionTreeClassifier(random_state=seed),
  749. ]
  750. for name, model in zip(model_names, passive_even_models):
  751. model.fit(X_train_passive_even, y_train_passive_even)
  752. pred = model.predict(X_test_passive_even)
  753. acc = accuracy_score(y_test_passive_even, pred)
  754. passive_even_seed_accuracies[name].append(acc)
  755. # Store for all seeds
  756. all_preds_passive_even[name].extend(pred)
  757. all_y_tests_passive_even[name].extend(y_test_passive_even)
  758. # Average across seeds
  759. passive_even_accuracies = [np.mean(passive_even_seed_accuracies[name]) for name in model_names]
  760. passive_even_std = [np.std(passive_even_seed_accuracies[name]) for name in model_names]
  761. # %% [markdown]
  762. # #### **Once more with randomized labels**
  763. # %%
  764. rng = np.random.default_rng(seed=42)
  765. # %%
  766. SEEDS = [42, 7, 13, 21, 99]
  767. mEPSC_rnd_seed_accuracies = {model_name: [] for model_name in [
  768. 'RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC'
  769. ]}
  770. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  771. all_preds_mEPSC_rnd = {name: [] for name in model_names}
  772. all_y_tests_mEPSC_rnd = {name: [] for name in model_names}
  773. for seed in SEEDS:
  774. X_train_mEPSC, X_test_mEPSC, y_train_mEPSC, y_test_mEPSC = train_test_split(
  775. X_mEPSC, y_mEPSC, test_size=0.25, shuffle=True, random_state=seed
  776. )
  777. y_train_mEPSC_rnd = rng.permutation(y_train_mEPSC)
  778. mEPSC_models = [
  779. RandomForestClassifier(random_state=seed),
  780. GradientBoostingClassifier(random_state=seed),
  781. LogisticRegression(max_iter=1000, random_state=seed),
  782. KNeighborsClassifier(), # No random_state
  783. MLPClassifier(max_iter=1000, random_state=seed),
  784. SVC(random_state=seed),
  785. DecisionTreeClassifier(random_state=seed),
  786. ]
  787. for name, model in zip(model_names, mEPSC_models):
  788. model.fit(X_train_mEPSC, y_train_mEPSC_rnd)
  789. pred = model.predict(X_test_mEPSC)
  790. acc = accuracy_score(y_test_mEPSC, pred)
  791. mEPSC_rnd_seed_accuracies[name].append(acc)
  792. # Store for all seeds
  793. all_preds_mEPSC_rnd[name].extend(pred)
  794. all_y_tests_mEPSC_rnd[name].extend(y_test_mEPSC)
  795. # Average across seeds (accuracy)
  796. mEPSC_rnd_accuracies = [np.mean(mEPSC_rnd_seed_accuracies[name]) for name in model_names]
  797. mEPSC_rnd_std = [np.std(mEPSC_rnd_seed_accuracies[name]) for name in model_names]
  798. # %%
  799. SEEDS = [42, 7, 13, 21, 99] # Same seeds as mEPSC for consistency
  800. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  801. spiking_rnd_seed_accuracies = {name: [] for name in model_names}
  802. all_preds_spiking_rnd = {name: [] for name in model_names}
  803. all_y_tests_spiking_rnd = {name: [] for name in model_names}
  804. for seed in SEEDS:
  805. X_train_spiking, X_test_spiking, y_train_spiking, y_test_spiking = train_test_split(
  806. X_spiking, y_spiking, test_size=0.25, shuffle=True, random_state=seed
  807. )
  808. y_train_spiking_rnd = rng.permutation(y_train_spiking)
  809. spiking_models = [
  810. RandomForestClassifier(random_state=seed),
  811. GradientBoostingClassifier(random_state=seed),
  812. LogisticRegression(max_iter=1000, random_state=seed),
  813. KNeighborsClassifier(), # No random_state
  814. MLPClassifier(max_iter=1000, random_state=seed),
  815. SVC(random_state=seed),
  816. DecisionTreeClassifier(random_state=seed),
  817. ]
  818. for name, model in zip(model_names, spiking_models):
  819. model.fit(X_train_spiking, y_train_spiking_rnd)
  820. pred = model.predict(X_test_spiking)
  821. acc = accuracy_score(y_test_spiking, pred)
  822. spiking_rnd_seed_accuracies[name].append(acc)
  823. # Store for all seeds
  824. all_preds_spiking_rnd[name].extend(pred)
  825. all_y_tests_spiking_rnd[name].extend(y_test_spiking)
  826. # Average across seeds
  827. spiking_rnd_accuracies = [np.mean(spiking_rnd_seed_accuracies[name]) for name in model_names]
  828. spiking_rnd_std = [np.std(spiking_rnd_seed_accuracies[name]) for name in model_names]
  829. # %%
  830. SEEDS = [42, 7, 13, 21, 99]
  831. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  832. passive_rnd_seed_accuracies = {name: [] for name in model_names}
  833. # Store ALL predictions and true labels for ALL seeds' runs
  834. all_preds_passive_rnd = {name: [] for name in model_names}
  835. all_y_tests_passive_rnd = {name: [] for name in model_names}
  836. for seed in SEEDS:
  837. X_train_passive, X_test_passive, y_train_passive, y_test_passive = train_test_split(
  838. X_passive, y_passive, test_size=0.25, shuffle=True, random_state=seed
  839. )
  840. y_train_passive_rnd = rng.permutation(y_train_passive)
  841. passive_models = [
  842. RandomForestClassifier(random_state=seed),
  843. GradientBoostingClassifier(random_state=seed),
  844. LogisticRegression(max_iter=1000, random_state=seed),
  845. KNeighborsClassifier(),
  846. MLPClassifier(max_iter=1000, random_state=seed),
  847. SVC(random_state=seed),
  848. DecisionTreeClassifier(random_state=seed),
  849. ]
  850. for name, model in zip(model_names, passive_models):
  851. model.fit(X_train_passive, y_train_passive_rnd)
  852. pred = model.predict(X_test_passive)
  853. acc = accuracy_score(y_test_passive, pred)
  854. passive_rnd_seed_accuracies[name].append(acc)
  855. # Store for all seeds
  856. all_preds_passive_rnd[name].extend(pred)
  857. all_y_tests_passive_rnd[name].extend(y_test_passive)
  858. # Average across seeds
  859. passive_rnd_accuracies = [np.mean(passive_rnd_seed_accuracies[name]) for name in model_names]
  860. passive_rnd_std = [np.std(passive_rnd_seed_accuracies[name]) for name in model_names]
  861. # %% [markdown]
  862. # #### **Saving Results**
  863. # %%
  864. mEPSC_results = {
  865. "Accuracies": mEPSC_seed_accuracies,
  866. "Classification Reports": mEPSC_classification_reports,
  867. "Feature Importances": mEPSC_feature_importances
  868. }
  869. spiking_results = {
  870. "Accuracies": spiking_seed_accuracies,
  871. "Classification Reports": spiking_classification_reports,
  872. "Feature Importances": spiking_feature_importances
  873. }
  874. passive_results = {
  875. "Accuracies": passive_seed_accuracies,
  876. "Classification Reports": passive_classification_reports,
  877. "Feature Importances": passive_feature_importances
  878. }
  879. mEPSC_even_results = {
  880. "Accuracies": mEPSC_even_seed_accuracies
  881. }
  882. spiking_even_results = {
  883. "Accuracies": spiking_even_seed_accuracies
  884. }
  885. passive_even_results = {
  886. "Accuracies": passive_even_seed_accuracies
  887. }
  888. mEPSC_rnd_results = {
  889. "Accuracies": mEPSC_rnd_seed_accuracies
  890. }
  891. spiking_rnd_results = {
  892. "Accuracies": spiking_rnd_seed_accuracies
  893. }
  894. passive_rnd_results = {
  895. "Accuracies": passive_rnd_seed_accuracies
  896. }
  897. # Define the base path for saving to Google Drive
  898. base_path = "/content/drive/MyDrive/Estrous Classification Project/"
  899. with open(base_path + "mEPSC_results.pkl", "wb") as f:
  900. pickle.dump(mEPSC_results, f)
  901. with open(base_path + "spiking_results.pkl", "wb") as f:
  902. pickle.dump(spiking_results, f)
  903. with open(base_path + "passive_results.pkl", "wb") as f:
  904. pickle.dump(passive_results, f)
  905. with open(base_path + "mEPSC_even_results.pkl", "wb") as f:
  906. pickle.dump(mEPSC_even_results, f)
  907. with open(base_path + "spiking_even_results.pkl", "wb") as f:
  908. pickle.dump(spiking_even_results, f)
  909. with open(base_path + "passive_even_results.pkl", "wb") as f:
  910. pickle.dump(passive_even_results, f)
  911. with open(base_path + "mEPSC_rnd_results.pkl", "wb") as f:
  912. pickle.dump(mEPSC_rnd_results, f)
  913. with open(base_path + "spiking_rnd_results.pkl", "wb") as f:
  914. pickle.dump(spiking_rnd_results, f)
  915. with open(base_path + "passive_rnd_results.pkl", "wb") as f:
  916. pickle.dump(passive_rnd_results, f)
  917. # %% [markdown]
  918. # ## **Results/Analysis**
  919. # %% [markdown]
  920. # ### **Loading Model Training Results**
  921. # %%
  922. # Define the base path for loading from Google Drive
  923. base_path = "/content/drive/MyDrive/Estrous Classification Project/"
  924. with open(base_path + "mEPSC_results.pkl", "rb") as f:
  925. mEPSC_results = pickle.load(f)
  926. with open(base_path + "spiking_results.pkl", "rb") as f:
  927. spiking_results = pickle.load(f)
  928. with open(base_path + "passive_results.pkl", "rb") as f:
  929. passive_results = pickle.load(f)
  930. with open(base_path + "mEPSC_even_results.pkl", "rb") as f:
  931. mEPSC_even_results = pickle.load(f)
  932. with open(base_path + "spiking_even_results.pkl", "rb") as f:
  933. spiking_even_results = pickle.load(f)
  934. with open(base_path + "passive_even_results.pkl", "rb") as f:
  935. passive_even_results = pickle.load(f)
  936. with open(base_path + "mEPSC_rnd_results.pkl", "rb") as f:
  937. mEPSC_rnd_results = pickle.load(f)
  938. with open(base_path + "spiking_rnd_results.pkl", "rb") as f:
  939. spiking_rnd_results = pickle.load(f)
  940. with open(base_path + "passive_rnd_results.pkl", "rb") as f:
  941. passive_rnd_results = pickle.load(f)
  942. # %%
  943. mEPSC_seed_accuracies = mEPSC_results["Accuracies"]
  944. mEPSC_classification_reports = mEPSC_results["Classification Reports"]
  945. mEPSC_feature_importances = mEPSC_results["Feature Importances"]
  946. spiking_seed_accuracies = spiking_results["Accuracies"]
  947. spiking_classification_reports = spiking_results["Classification Reports"]
  948. spiking_feature_importances = spiking_results["Feature Importances"]
  949. passive_seed_accuracies = passive_results["Accuracies"]
  950. passive_classification_reports = passive_results["Classification Reports"]
  951. passive_feature_importances = passive_results["Feature Importances"]
  952. mEPSC_even_seed_accuracies = mEPSC_even_results["Accuracies"]
  953. spiking_even_seed_accuracies = spiking_even_results["Accuracies"]
  954. passive_even_seed_accuracies = passive_even_results["Accuracies"]
  955. mEPSC_rnd_seed_accuracies = mEPSC_rnd_results["Accuracies"]
  956. spiking_rnd_seed_accuracies = spiking_rnd_results["Accuracies"]
  957. passive_rnd_seed_accuracies = passive_rnd_results["Accuracies"]
  958. # %% [markdown]
  959. # ### **Feature Importances**
  960. #
  961. # %%
  962. seeds = mEPSC_feature_importances.keys()
  963. mEPSC_imp = np.vstack([mEPSC_feature_importances[seed] for seed in seeds])
  964. spiking_imp = np.vstack([spiking_feature_importances[seed] for seed in seeds])
  965. passive_imp = np.vstack([passive_feature_importances[seed] for seed in seeds])
  966. # %%
  967. mEPSC_summary = pd.DataFrame({
  968. "mean": mEPSC_imp.mean(axis=0),
  969. "std": mEPSC_imp.std(axis=0),
  970. "top5_freq": (
  971. pd.DataFrame(mEPSC_imp, columns=mEPSC_names)
  972. .rank(axis=1, ascending=False) <= 5
  973. ).mean(axis=0).values,
  974. }, index=mEPSC_names).sort_values("mean", ascending=False)
  975. passive_summary = pd.DataFrame({
  976. "mean": passive_imp.mean(axis=0),
  977. "std": passive_imp.std(axis=0),
  978. "top5_freq": (
  979. pd.DataFrame(passive_imp, columns=passive_names_pretty)
  980. .rank(axis=1, ascending=False) <= 5
  981. ).mean(axis=0).values,
  982. }, index=passive_names_pretty).sort_values("mean", ascending=False)
  983. spiking_summary = pd.DataFrame({
  984. "mean": spiking_imp.mean(axis=0),
  985. "std": spiking_imp.std(axis=0),
  986. "top5_freq": (
  987. pd.DataFrame(spiking_imp, columns=spiking_names)
  988. .rank(axis=1, ascending=False) <= 5
  989. ).mean(axis=0).values,
  990. }, index=spiking_names).sort_values("mean", ascending=False)
  991. # %%
  992. mEPSC_importances = mEPSC_summary.head(5)
  993. spiking_importances = spiking_summary.head(5)
  994. passive_importances = passive_summary.head(5)
  995. mEPSC_top5 = mEPSC_importances.index.tolist()
  996. spiking_top5 = spiking_importances.index.tolist()
  997. passive_top5 = passive_importances.index.tolist()
  998. # %%
  999. passive_top5 = ['Input Resistance',
  1000. 'Membrane Time Constant (ms)',
  1001. 'Rebound Voltage (mV)',
  1002. 'Voltage Deflection Onset (mV)',
  1003. 'RMP (mV)']
  1004. # %%
  1005. spiking_top5 = ['Burst Length (ms)',
  1006. 'Spike Frequency',
  1007. 'Average Spike Amplitude',
  1008. 'Average Spike Peaks',
  1009. 'Spikes per Burst']
  1010. # %%
  1011. mEPSC_top5 = ['Baseline',
  1012. 'Amplitude',
  1013. 'HalfWidth',
  1014. '10-90 Slope',
  1015. 'Inter-Event Interval (ms)']
  1016. # %%
  1017. print("\n" + "="*60)
  1018. print("ALL FEATURES BY IMPORTANCE")
  1019. print("="*60)
  1020. print("\nmEPSC Dataset:")
  1021. print(f" {'Rank':<6} {'Feature':<35} {'Importance'}")
  1022. print(f" {'-'*6} {'-'*35} {'-'*10}")
  1023. for rank, (name, row) in enumerate(mEPSC_summary.iterrows(), start=1):
  1024. marker = " " if rank > 5 else "* "
  1025. print(f" {marker}{rank:<6} {name:<35} {row['mean']:.4f}")
  1026. print("\nSpiking Dataset:")
  1027. print(f" {'Rank':<6} {'Feature':<35} {'Importance'}")
  1028. print(f" {'-'*6} {'-'*35} {'-'*10}")
  1029. for rank, (name, row) in enumerate(spiking_summary.iterrows(), start=1):
  1030. marker = " " if rank > 5 else "* "
  1031. print(f" {marker}{rank:<6} {name:<35} {row['mean']:.4f}")
  1032. print("\nPassive Dataset:")
  1033. print(f" {'Rank':<6} {'Feature':<35} {'Importance'}")
  1034. print(f" {'-'*6} {'-'*35} {'-'*10}")
  1035. for rank, (name, row) in enumerate(passive_summary.iterrows(), start=1):
  1036. marker = " " if rank > 5 else "* "
  1037. print(f" {marker}{rank:<6} {name:<35} {row['mean']:.4f}")
  1038. print("\n* = included in plot")
  1039. print("="*60)
  1040. # %% [markdown]
  1041. # ### **t-SNE**
  1042. # %%
  1043. palette = 'tab10'
  1044. # %%
  1045. tSNE_X_mEPSC = mEPSC_even.drop("label", axis=1)
  1046. tSNE_y_mEPSC = mEPSC_even["label"]
  1047. # %%
  1048. tSNE_X_passive = passive_even.drop("label", axis=1)
  1049. tSNE_y_passive = passive_even["label"]
  1050. # %%
  1051. tSNE_X_spiking = spiking_even.drop("label", axis=1)
  1052. tSNE_y_spiking = spiking_even["label"]
  1053. # %%
  1054. # Re-run t-SNE with 3 components for mEPSC data
  1055. tSNE_mEPSC_3d = TSNE(n_components=3, verbose=1, perplexity=25, n_iter=300, random_state=42)
  1056. mEPSC_results_3d = tSNE_mEPSC_3d.fit_transform(tSNE_X_mEPSC)
  1057. mEPSC_tSNE_comp1_3d = mEPSC_results_3d[:,0]
  1058. mEPSC_tSNE_comp2_3d = mEPSC_results_3d[:,1]
  1059. mEPSC_tSNE_comp3_3d = mEPSC_results_3d[:,2]
  1060. # %%
  1061. # Re-run t-SNE with 3 components for passive data
  1062. tSNE_passive_3d = TSNE(n_components=3, verbose=1, perplexity=25, n_iter=300, random_state=42)
  1063. passive_results_3d = tSNE_passive_3d.fit_transform(tSNE_X_passive)
  1064. passive_tSNE_comp1_3d = passive_results_3d[:,0]
  1065. passive_tSNE_comp2_3d = passive_results_3d[:,1]
  1066. passive_tSNE_comp3_3d = passive_results_3d[:,2]
  1067. # %%
  1068. # Re-run t-SNE with 3 components for spiking data
  1069. tSNE_spiking_3d = TSNE(n_components=3, verbose=1, perplexity=25, n_iter=300, random_state=42)
  1070. spiking_results_3d = tSNE_spiking_3d.fit_transform(tSNE_X_spiking)
  1071. spiking_tSNE_comp1_3d = spiking_results_3d[:,0]
  1072. spiking_tSNE_comp2_3d = spiking_results_3d[:,1]
  1073. spiking_tSNE_comp3_3d = spiking_results_3d[:,2]
  1074. # %% [markdown]
  1075. # ### **Figures**
  1076. # %% [markdown]
  1077. # #### **Model vs Accuracy Figure**
  1078. # %%
  1079. #POSTER VERSION
  1080. plot_colors = {
  1081. "Action Potentials": "#4393c3",
  1082. "Passive": "#92c5de",
  1083. "mEPSC": "#2166ac"
  1084. }
  1085. model_labels = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  1086. # Narrower width compresses points together; taller height gives more breathing room
  1087. fig, ax = plt.subplots(figsize=(10, 5), dpi=300)
  1088. plt.style.use('seaborn-v0_8-darkgrid')
  1089. # Lighter background
  1090. ax.set_facecolor('#f7f9fc')
  1091. fig.patch.set_facecolor('#f7f9fc')
  1092. datasets = {
  1093. "Action Potentials": (spiking_accuracies, spiking_std),
  1094. "Passive": (passive_accuracies, passive_std),
  1095. "mEPSC": (mEPSC_accuracies, mEPSC_std),
  1096. }
  1097. for label, (means, stds) in datasets.items():
  1098. ax.errorbar(
  1099. model_labels, means, yerr=stds,
  1100. label=label,
  1101. color=plot_colors[label],
  1102. marker='o',
  1103. linewidth=2.0,
  1104. markersize=6,
  1105. capsize=4,
  1106. capthick=1.5,
  1107. elinewidth=1.5,
  1108. )
  1109. ax.set_ylabel("Classification Accuracy", fontsize=16, fontweight='bold')
  1110. ax.set_xlabel("Model Type", fontsize=16, fontweight='bold')
  1111. ax.spines['top'].set_visible(False)
  1112. ax.spines['right'].set_visible(False)
  1113. plt.xticks(rotation=45, ha='right', fontsize=13, fontweight='bold')
  1114. plt.yticks(fontsize=13, fontweight='bold')
  1115. legend = ax.legend(
  1116. title="Feature Set",
  1117. fontsize=12,
  1118. title_fontsize=13,
  1119. frameon=False,
  1120. bbox_to_anchor=(1.02, 1.0),
  1121. loc='upper left'
  1122. )
  1123. # Bold legend title
  1124. legend.get_title().set_fontweight('bold')
  1125. plt.tight_layout()
  1126. plt.show()
  1127. # %%
  1128. #HELPER FUNCTIONS FOR PAPER VERSION
  1129. #Loading stats results for significance markers
  1130. tukey_df = pd.read_csv(base_path + 'MC for accuracies.csv', index_col=0)
  1131. tukey_df[['feat1','model1']] = tukey_df['group1'].str.split('_', n=1, expand=True)
  1132. tukey_df[['feat2','model2']] = tukey_df['group2'].str.split('_', n=1, expand=True)
  1133. model_order = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  1134. def cld_within_feature(feature_set):
  1135. """Compact letters comparing models WITHIN one feature set (for panels B/C/D)"""
  1136. sub = tukey_df[(tukey_df['feat1']==feature_set) & (tukey_df['feat2']==feature_set)]
  1137. G = nx.Graph()
  1138. G.add_nodes_from(model_order)
  1139. for _, row in sub.iterrows():
  1140. if row['reject']:
  1141. G.add_edge(row['model1'], row['model2'])
  1142. coloring = nx.coloring.greedy_color(G, strategy='largest_first')
  1143. colors_sorted = sorted(set(coloring.values()))
  1144. letter_map = {c: chr(97+i) for i, c in enumerate(colors_sorted)}
  1145. return {m: letter_map[coloring[m]] for m in model_order}
  1146. def cld_within_model(model):
  1147. """Compact letters comparing feature sets WITHIN one model (for panel A)"""
  1148. sub = tukey_df[(tukey_df['model1']==model) & (tukey_df['model2']==model)]
  1149. G = nx.Graph()
  1150. G.add_nodes_from(['AP','Passive','mEPSC'])
  1151. for _, row in sub.iterrows():
  1152. if row['reject']:
  1153. G.add_edge(row['feat1'], row['feat2'])
  1154. coloring = nx.coloring.greedy_color(G, strategy='largest_first')
  1155. colors_sorted = sorted(set(coloring.values()))
  1156. letter_map = {c: chr(97+i) for i, c in enumerate(colors_sorted)}
  1157. return {f: letter_map[coloring[f]] for f in ['AP','Passive','mEPSC']}
  1158. def tukey_row(a, b):
  1159. """Return the Tukey row for groups a,b in either order."""
  1160. r = tuk[((tuk.group1 == a) & (tuk.group2 == b)) |
  1161. ((tuk.group1 == b) & (tuk.group2 == a))]
  1162. if len(r) != 1:
  1163. raise KeyError(f"expected 1 row for {a} vs {b}, got {len(r)}")
  1164. return r.iloc[0]
  1165. SIG_LEVELS = [(1e-4, '**'), (0.05, '*')]
  1166. def stars(p):
  1167. for thresh, sym in SIG_LEVELS:
  1168. if p < thresh:
  1169. return sym
  1170. return None
  1171. def _summarise(d):
  1172. mu = [np.mean(v) * 100 for v in d.values()]
  1173. sd = [np.std(v) * 100 for v in d.values()]
  1174. return mu, sd
  1175. def compact_letters(fs):
  1176. ns = {m: set() for m in model_labels}
  1177. for m1, m2 in itertools.combinations(model_labels, 2):
  1178. if stars(tukey_row(f'{fs}_{m1}', f'{fs}_{m2}')['p-adj']) is None:
  1179. ns[m1].add(m2)
  1180. ns[m2].add(m1)
  1181. order = sorted(model_labels, key=lambda m: -ACC[fs][model_labels.index(m)])
  1182. cliques = []
  1183. for m in order:
  1184. cand = [m] + [x for x in order if x in ns[m]]
  1185. for start in range(len(cand)):
  1186. cl = {cand[start]}
  1187. for x in cand:
  1188. if x in cl:
  1189. continue
  1190. if all(y in ns[x] for y in cl):
  1191. cl.add(x)
  1192. if cl not in cliques:
  1193. cliques.append(cl)
  1194. cliques = [c for c in cliques if not any(c < d for d in cliques)]
  1195. cliques.sort(key=lambda c: -max(ACC[fs][model_labels.index(m)] for m in c))
  1196. lab = {m: '' for m in model_labels}
  1197. for letter, c in zip('abcdefghij', cliques):
  1198. for m in c:
  1199. lab[m] += letter
  1200. return lab, cliques
  1201. def draw_panel(axis, fs):
  1202. a, e = ACC[fs], STD[fs]
  1203. axis.errorbar(xpos, a, yerr=e, color=plot_colors[SET2COL[fs]], marker='o',
  1204. linewidth=0, markersize=6, capsize=4, capthick=1.2, elinewidth=1.2)
  1205. for m1, m2 in itertools.combinations(model_labels, 2):
  1206. r = tukey_row(f'{fs}_{m1}', f'{fs}_{m2}')
  1207. audit.append((SET2COL[fs], f'{fs}: {m1} vs {m2}', int(r.csv_line), r.group1,
  1208. r.group2, r.meandiff, r['p-adj'], stars(r['p-adj']) or 'ns'))
  1209. top = 105
  1210. if ANNOT_MODE == 'cld':
  1211. lab, _ = compact_letters(fs)
  1212. for i, m in enumerate(model_labels):
  1213. axis.text(i, a[i] + e[i] + 2.0, ','.join(lab[m]), ha='center',
  1214. fontsize=11, fontweight='bold')
  1215. print(f"[{SET2COL[fs]}] CLD: " + ", ".join(
  1216. f"{m}={','.join(lab[m])}" for m in sorted(model_labels,
  1217. key=lambda m: -a[model_labels.index(m)])))
  1218. elif ANNOT_MODE == 'brackets':
  1219. sig = []
  1220. for m1, m2 in itertools.combinations(model_labels, 2):
  1221. sym = stars(tukey_row(f'{fs}_{m1}', f'{fs}_{m2}')['p-adj'])
  1222. if sym:
  1223. sig.append((model_labels.index(m1), model_labels.index(m2), sym))
  1224. sig.sort(key=lambda t: abs(t[1] - t[0]))
  1225. y = max(np.array(a) + np.array(e)) + 3
  1226. for i1, i2, sym in sig:
  1227. axis.plot([i1, i1, i2, i2], [y, y + 0.9, y + 0.9, y], color='black', lw=1)
  1228. axis.text((i1 + i2) / 2, y + 1.1, sym, ha='center', fontsize=7)
  1229. y += 3.4
  1230. top = y + 3
  1231. axis.set_ylim(40, top)
  1232. #CONSTANTS AND VARIABLES
  1233. letters_A = {m: cld_within_model(m) for m in model_order}
  1234. letters_AP = cld_within_feature('AP')
  1235. letters_Passive = cld_within_feature('Passive')
  1236. letters_mEPSC = cld_within_feature('mEPSC')
  1237. TUKEY_CSV = base_path + 'MC for accuracies.csv'
  1238. OUT_TIF = base_path + 'figure3.tif'
  1239. ANNOT_MODE = 'cld'
  1240. SHOW_NS_A = False
  1241. LABEL_ALL = True
  1242. DPI = 300
  1243. FIGSIZE = (16, 10)
  1244. model_labels = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  1245. feature_sets = ['AP', 'Passive', 'mEPSC']
  1246. plot_colors = {"Action Potentials": "#4393c3", "Passive": "#92c5de", "mEPSC": "#2166ac"}
  1247. SET2COL = {'AP': "Action Potentials", 'Passive': "Passive", 'mEPSC': "mEPSC"}
  1248. ARM_START = 0.15
  1249. LANE_STEP = 0.26
  1250. ARM_LEN = 0.09
  1251. STAR_SIZE = 12
  1252. STAR_PAD = 0.03
  1253. STAR_STACK = True
  1254. # %%
  1255. tuk = pd.read_csv(TUKEY_CSV, index_col=0)
  1256. tuk['csv_line'] = tuk.index + 2 # 1-indexed line in the raw file
  1257. xpos = np.arange(len(model_labels))
  1258. XLIM = (-0.5, len(model_labels) + 0.05) # same on every panel so columns align
  1259. spiking_accuracies, spiking_std = _summarise(spiking_seed_accuracies)
  1260. passive_accuracies, passive_std = _summarise(passive_seed_accuracies)
  1261. mEPSC_accuracies, mEPSC_std = _summarise(mEPSC_seed_accuracies)
  1262. ACC = {'AP': spiking_accuracies, 'Passive': passive_accuracies, 'mEPSC': mEPSC_accuracies}
  1263. STD = {'AP': spiking_std, 'Passive': passive_std, 'mEPSC': mEPSC_std}
  1264. plt.style.use('seaborn-v0_8-darkgrid')
  1265. fig, ax = plt.subplots(nrows=2, ncols=2, figsize=FIGSIZE)
  1266. audit = []
  1267. # ============ PANEL A : within model, across feature sets ============
  1268. for fs in feature_sets:
  1269. ax[0, 0].errorbar(xpos, ACC[fs], yerr=STD[fs], label=SET2COL[fs],
  1270. color=plot_colors[SET2COL[fs]], marker='o', linewidth=0,
  1271. markersize=5, capsize=4, capthick=1.2, elinewidth=1.2)
  1272. for i, m in enumerate(model_labels):
  1273. pairs = []
  1274. for a, b in itertools.combinations(feature_sets, 2):
  1275. r = tukey_row(f'{a}_{m}', f'{b}_{m}')
  1276. sym = stars(r['p-adj'])
  1277. audit.append(('A', f'{m}: {a} vs {b}', int(r.csv_line), r.group1, r.group2,
  1278. r.meandiff, r['p-adj'], sym or 'ns'))
  1279. if sym is None and not SHOW_NS_A:
  1280. continue
  1281. pairs.append((ACC[a][i], ACC[b][i], sym or 'ns'))
  1282. pairs.sort(key=lambda t: abs(t[0] - t[1]))
  1283. for lane, (y1, y2, sym) in enumerate(pairs):
  1284. x0 = i + ARM_START + lane * LANE_STEP
  1285. x1 = x0 + ARM_LEN
  1286. lo, hi = sorted([y1, y2])
  1287. ax[0, 0].plot([x1, x1], [lo, hi], color='black', lw=1)
  1288. ax[0, 0].plot([x0, x1], [lo, lo], color='black', lw=1)
  1289. ax[0, 0].plot([x0, x1], [hi, hi], color='black', lw=1)
  1290. txt = '\n'.join(sym) if (STAR_STACK and sym != 'ns') else sym
  1291. ax[0, 0].text(x1 + STAR_PAD, (lo + hi) / 2, txt, fontsize=STAR_SIZE,
  1292. style='italic' if sym == 'ns' else 'normal',
  1293. va='center', ha='left', linespacing=0.25)
  1294. ax[0, 0].set_ylim(40, 105)
  1295. ax[0, 0].legend(title="Feature Set", fontsize=9, title_fontsize=10,
  1296. frameon=False, loc='lower left')
  1297. draw_panel(ax[0, 1], 'AP')
  1298. draw_panel(ax[1, 0], 'Passive')
  1299. draw_panel(ax[1, 1], 'mEPSC')
  1300. # ============ shared cosmetics ============
  1301. uniform_y = len({tuple(a.get_ylim()) for a in ax.ravel()}) == 1
  1302. for (row, col), axis in np.ndenumerate(ax):
  1303. axis.set_xticks(xpos)
  1304. axis.set_xlim(*XLIM)
  1305. for sp in ('top', 'right'):
  1306. axis.spines[sp].set_visible(False)
  1307. show_y = LABEL_ALL or col == 0
  1308. show_x = LABEL_ALL or row == 1
  1309. axis.set_ylabel("Classification Accuracy (%)" if show_y else "", fontsize=12)
  1310. if not show_y and uniform_y:
  1311. axis.tick_params(labelleft=False)
  1312. axis.set_xlabel("Model Type" if show_x else "", fontsize=12)
  1313. if show_x:
  1314. axis.set_xticklabels(model_labels, rotation=45)
  1315. else:
  1316. axis.set_xticklabels([])
  1317. for axis, letter in zip(ax.ravel(), ['A', 'B', 'C', 'D']):
  1318. axis.text(-0.06, 1.06, letter, transform=axis.transAxes,
  1319. fontsize=15, fontweight='bold', va='top', ha='right')
  1320. plt.tight_layout()
  1321. fig.subplots_adjust(wspace=0.10, hspace=0.18)
  1322. # TIFF export. matplotlib renders RGBA; most journals require flattened RGB,
  1323. # so re-open and drop the alpha channel over white.
  1324. fig.savefig(OUT_TIF, dpi=DPI, facecolor='white',
  1325. pil_kwargs={'compression': 'tiff_lzw'})
  1326. _im = Image.open(OUT_TIF)
  1327. if _im.mode != 'RGB':
  1328. _flat = Image.new('RGB', _im.size, 'white')
  1329. _flat.paste(_im, mask=_im.split()[-1] if _im.mode == 'RGBA' else None)
  1330. _flat.save(OUT_TIF, compression='tiff_lzw', dpi=(DPI, DPI))
  1331. _im.close()
  1332. plt.show()
  1333. # ---------------- audit table for the supplement ----------------
  1334. audit_df = pd.DataFrame(audit, columns=['panel', 'comparison', 'csv_line',
  1335. 'group1', 'group2', 'meandiff',
  1336. 'p_adj', 'symbol'])
  1337. audit_df.to_csv(base_path + 'figure3_significance_audit.csv', index=False)
  1338. # %%
  1339. #Generated by Claude to assist with organizing MCs into output format for supplementary table S1
  1340. SUPP_PREFIX = base_path + 'figure3_supp_'
  1341. PP = 100.0
  1342. def fmt_p(p):
  1343. """Adjusted p-values hit the solver's numerical floor (~1.9e-14), so anything
  1344. below the reporting threshold is shown as an inequality rather than a
  1345. spuriously precise number."""
  1346. if p < 1e-4:
  1347. return '<0.0001'
  1348. if p < 1e-3:
  1349. return f'{p:.5f}'
  1350. return f'{p:.4f}'
  1351. def oriented(g1, g2):
  1352. """Tukey rows store an arbitrary group order. Return diff and CI expressed as
  1353. (mean of g2) - (mean of g1) regardless of how the row is stored."""
  1354. r = tukey_row(g1, g2)
  1355. diff, lo, hi = r.meandiff, r.lower, r.upper
  1356. if r.group1 != g1: # row is stored the other way round
  1357. diff, lo, hi = -diff, -hi, -lo # negate and swap the bounds
  1358. return diff * PP, lo * PP, hi * PP, r['p-adj'], int(r.csv_line)
  1359. # ---------------- Table S1: group summary ----------------
  1360. letters = {fs: compact_letters(fs)[0] for fs in feature_sets}
  1361. _seed_dicts = {'AP': spiking_seed_accuracies,
  1362. 'Passive': passive_seed_accuracies,
  1363. 'mEPSC': mEPSC_seed_accuracies}
  1364. rows_s1 = []
  1365. for fs in feature_sets:
  1366. for j, m in enumerate(model_labels):
  1367. rows_s1.append({
  1368. 'Feature set': SET2COL[fs],
  1369. 'Model': m,
  1370. 'Mean accuracy (%)': round(ACC[fs][j], 2),
  1371. 'SD (%)': round(STD[fs][j], 2),
  1372. 'n (seeds)': len(list(_seed_dicts[fs].values())[j]),
  1373. 'Group': ','.join(letters[fs][m]),
  1374. })
  1375. table_s1 = pd.DataFrame(rows_s1)
  1376. # ---------------- Table S2: pairwise comparisons ----------------
  1377. rows_s2 = []
  1378. # Panel A -- within model, across feature sets
  1379. for m in model_labels:
  1380. for a, b in itertools.combinations(feature_sets, 2):
  1381. d, lo, hi, p, line = oriented(f'{a}_{m}', f'{b}_{m}')
  1382. rows_s2.append({
  1383. 'Panel': 'A', 'Stratum': m,
  1384. 'Comparison': f'{SET2COL[b]} - {SET2COL[a]}',
  1385. 'Difference (pp)': round(d, 2),
  1386. '95% CI': f'[{lo:.2f}, {hi:.2f}]',
  1387. 'p-adj': fmt_p(p),
  1388. 'Significance': stars(p) or 'n.s.',
  1389. 'Source row': line,
  1390. })
  1391. # Panels B/C/D -- within feature set, across models
  1392. for fs, panel in zip(feature_sets, ['B', 'C', 'D']):
  1393. for m1, m2 in itertools.combinations(model_labels, 2):
  1394. d, lo, hi, p, line = oriented(f'{fs}_{m1}', f'{fs}_{m2}')
  1395. rows_s2.append({
  1396. 'Panel': panel, 'Stratum': SET2COL[fs],
  1397. 'Comparison': f'{m2} - {m1}',
  1398. 'Difference (pp)': round(d, 2),
  1399. '95% CI': f'[{lo:.2f}, {hi:.2f}]',
  1400. 'p-adj': fmt_p(p),
  1401. 'Significance': stars(p) or 'n.s.',
  1402. 'Source row': line,
  1403. })
  1404. table_s2 = pd.DataFrame(rows_s2)
  1405. # ---------------- sanity checks ----------------
  1406. assert len(table_s2) == 21 + 3 * 21, f"expected 84 comparisons, got {len(table_s2)}"
  1407. _recon = []
  1408. for _, r in table_s2.iterrows():
  1409. lo, hi = [float(v) for v in r['95% CI'].strip('[]').split(',')]
  1410. _recon.append(lo <= r['Difference (pp)'] <= hi)
  1411. assert all(_recon), "a difference fell outside its own CI -- check orientation logic"
  1412. table_s1.to_csv(SUPP_PREFIX + 'tableS1_group_summary.csv', index=False)
  1413. table_s2.to_csv(SUPP_PREFIX + 'tableS2_pairwise_tukey.csv', index=False)
  1414. # %%
  1415. print("="*10, "Model Accuracies and Standard Deviations", "="*10)
  1416. print("="*15, "mEPSC", "="*15)
  1417. for model, accuracies in mEPSC_seed_accuracies.items():
  1418. mu = np.mean(accuracies)
  1419. sigma = np.std(accuracies)
  1420. print(f"{model}: {mu:.4f} ± {sigma:.4f}")
  1421. print("="*15, "AP", "="*15)
  1422. for model, accuracies in spiking_seed_accuracies.items():
  1423. mu = np.mean(accuracies)
  1424. sigma = np.std(accuracies)
  1425. print(f"{model}: {mu:.4f} ± {sigma:.4f}")
  1426. print("="*15, "Passive", "="*15)
  1427. for model, accuracies in passive_seed_accuracies.items():
  1428. mu = np.mean(accuracies)
  1429. sigma = np.std(accuracies)
  1430. print(f"{model}: {mu:.4f} ± {sigma:.4f}")
  1431. # %% [markdown]
  1432. # #### **Classification Report Figure**
  1433. # %%
  1434. passive_diestrus_recalls = []
  1435. passive_estrus_recalls = []
  1436. passive_Eproestrus_recalls = []
  1437. passive_Lproestrus_recalls = []
  1438. for model_report in passive_classification_reports.values():
  1439. passive_diestrus_recalls.append(model_report['0']['recall'])
  1440. passive_estrus_recalls.append(model_report['1']['recall'])
  1441. passive_Eproestrus_recalls.append(model_report['2']['recall'])
  1442. passive_Lproestrus_recalls.append(model_report['3']['recall'])
  1443. spiking_diestrus_recalls = []
  1444. spiking_estrus_recalls = []
  1445. spiking_Eproestrus_recalls = []
  1446. spiking_Lproestrus_recalls = []
  1447. for model_report in spiking_classification_reports.values():
  1448. spiking_diestrus_recalls.append(model_report['0']['recall'])
  1449. spiking_estrus_recalls.append(model_report['1']['recall'])
  1450. spiking_Eproestrus_recalls.append(model_report['2']['recall'])
  1451. spiking_Lproestrus_recalls.append(model_report['3']['recall'])
  1452. mEPSC_diestrus_recalls = []
  1453. mEPSC_estrus_recalls = []
  1454. mEPSC_Eproestrus_recalls = []
  1455. mEPSC_Lproestrus_recalls = []
  1456. for model_report in mEPSC_classification_reports.values():
  1457. mEPSC_diestrus_recalls.append(model_report['0']['recall'])
  1458. mEPSC_estrus_recalls.append(model_report['1']['recall'])
  1459. mEPSC_Eproestrus_recalls.append(model_report['2']['recall'])
  1460. mEPSC_Lproestrus_recalls.append(model_report['3']['recall'])
  1461. # %%
  1462. print("Passive")
  1463. print(f"Diestrus: {np.mean(passive_diestrus_recalls):.4f} +- {np.std(passive_diestrus_recalls):.4f}")
  1464. print(f"Estrus: {np.mean(passive_estrus_recalls):.4f} +- {np.std(passive_estrus_recalls):.4f}")
  1465. print(f"Early Proestrus: {np.mean(passive_Eproestrus_recalls):.4f} +- {np.std(passive_Eproestrus_recalls):.4f}")
  1466. print(f"Late Proestrus: {np.mean(passive_Lproestrus_recalls):.4f} +- {np.std(passive_Lproestrus_recalls):.4f}")
  1467. print("AP")
  1468. print(f"Diestrus: {np.mean(spiking_diestrus_recalls):.4f} +- {np.std(spiking_diestrus_recalls):.4f}")
  1469. print(f"Estrus: {np.mean(spiking_estrus_recalls):.4f} +- {np.std(spiking_estrus_recalls):.4f}")
  1470. print(f"Early Proestrus: {np.mean(spiking_Eproestrus_recalls):.4f} +- {np.std(spiking_Eproestrus_recalls):.4f}")
  1471. print(f"Late Proestrus: {np.mean(spiking_Lproestrus_recalls):.4f} +- {np.std(spiking_Lproestrus_recalls):.4f}")
  1472. print("mEPSC")
  1473. print(f"Diestrus: {np.mean(mEPSC_diestrus_recalls):.4f} +- {np.std(mEPSC_diestrus_recalls):.4f}")
  1474. print(f"Estrus: {np.mean(mEPSC_estrus_recalls):.4f} +- {np.std(mEPSC_estrus_recalls):.4f}")
  1475. print(f"Early Proestrus: {np.mean(mEPSC_Eproestrus_recalls):.4f} +- {np.std(mEPSC_Eproestrus_recalls):.4f}")
  1476. print(f"Late Proestrus: {np.mean(mEPSC_Lproestrus_recalls):.4f} +- {np.std(mEPSC_Lproestrus_recalls):.4f}")
  1477. # %%
  1478. plt.style.use('seaborn-v0_8-white')
  1479. phase_labels = ['Diestrus', 'Estrus', 'Early Proestrus', 'Late Proestrus']
  1480. def make_long_df(diestrus, estrus, eproestrus, lproestrus, dataset_name):
  1481. model_names = list(passive_classification_reports.keys())
  1482. rows = []
  1483. for recalls, phase in zip([diestrus, estrus, eproestrus, lproestrus], phase_labels):
  1484. for model, val in zip(model_names, recalls):
  1485. rows.append({'phase': phase, 'Recall': val*100, 'model': model, 'Dataset': dataset_name})
  1486. return pd.DataFrame(rows)
  1487. passive_df = make_long_df(passive_diestrus_recalls, passive_estrus_recalls,
  1488. passive_Eproestrus_recalls, passive_Lproestrus_recalls, 'Passive')
  1489. spiking_df = make_long_df(spiking_diestrus_recalls, spiking_estrus_recalls,
  1490. spiking_Eproestrus_recalls, spiking_Lproestrus_recalls, 'AP')
  1491. mEPSC_df = make_long_df(mEPSC_diestrus_recalls, mEPSC_estrus_recalls,
  1492. mEPSC_Eproestrus_recalls, mEPSC_Lproestrus_recalls, 'mEPSC')
  1493. combined_df = pd.concat([spiking_df, passive_df, mEPSC_df], ignore_index=True)
  1494. class_colors = ['#E63946', '#457B9D', '#2A9D8F', '#E76F51']
  1495. fig, ax = plt.subplots(figsize=(8, 6), dpi=300)
  1496. fig.patch.set_facecolor('white')
  1497. ax.set_facecolor('white')
  1498. sns.stripplot(data=combined_df, x='Dataset', y='Recall', hue='phase',
  1499. hue_order=phase_labels, dodge=True, size=8, jitter=0.1,
  1500. alpha=0.8, palette=class_colors, ax=ax)
  1501. sns.pointplot(data=combined_df, x='Dataset', y='Recall', hue='phase',
  1502. hue_order=phase_labels, dodge=0.8 - 0.8/len(phase_labels),
  1503. estimator='mean', errorbar=None,
  1504. color='black', linestyle='none',
  1505. markers='_', markersize=20, markeredgewidth=3, ax=ax)
  1506. handles, labels = ax.get_legend_handles_labels()
  1507. ax.legend(handles[:4], labels[:4], title='Phase', bbox_to_anchor=(1.05, 1),
  1508. loc='upper left', frameon=True, edgecolor='black', facecolor='white',
  1509. borderaxespad=0.)
  1510. sig_pairs = {
  1511. 'AP': [((0, 1), '**'), ((2, 3), '***'), ((1, 3), '****')],
  1512. 'Passive': [((0, 1), '*'), ((1, 3), '***')],
  1513. 'mEPSC': [((2, 3), '*'), ((1, 3), '***'), ((0, 3), '*')],
  1514. }
  1515. dataset_order = ['AP', 'Passive', 'mEPSC']
  1516. n_hue = len(phase_labels)
  1517. hue_offset = {i: -0.4 + (0.8 / n_hue) * (i + 0.5) for i in range(n_hue)} # -0.3,-0.1,0.1,0.3
  1518. tick_h = 2.5
  1519. pad = 6.0
  1520. step = 9.0
  1521. for ds, pairs in sig_pairs.items():
  1522. xc = dataset_order.index(ds)
  1523. y_base = combined_df.loc[combined_df['Dataset'] == ds, 'Recall'].max() + pad
  1524. for k, ((a, b), stars) in enumerate(pairs):
  1525. y = y_base + k * step
  1526. x1 = xc + hue_offset[a]
  1527. x2 = xc + hue_offset[b]
  1528. ax.plot([x1, x1, x2, x2], [y, y + tick_h, y + tick_h, y],
  1529. lw=1.5, c='black', solid_capstyle='projecting', clip_on=False)
  1530. ax.text((x1 + x2) / 2, y + tick_h, stars, ha='center', va='bottom',
  1531. color='black', fontsize=13)
  1532. ax.set_ylim(0, 130)
  1533. ax.set_ylabel("Recall (%)", fontsize=14)
  1534. ax.set_xlabel("Dataset", fontsize=14)
  1535. fig.subplots_adjust(right=0.75)
  1536. plt.tight_layout()
  1537. # Explicitly set y-ticks from 0 to 100 even though the graph extends above that for significance bars
  1538. custom_yticks = np.arange(0, 101, 20)
  1539. ax.set_yticks(custom_yticks)
  1540. ax.set_yticklabels([str(int(t)) for t in custom_yticks])
  1541. plt.savefig(base_path + 'figure6.tif', bbox_inches='tight')
  1542. plt.show()
  1543. # %% [markdown]
  1544. # #### **Feature Importances Figure**
  1545. # %%
  1546. # Set the plot style and color
  1547. plt.style.use('seaborn-v0_8-darkgrid')
  1548. plot_color = "#2166ac"
  1549. # Create figure with 3 subplots
  1550. fig, axes = plt.subplots(3, 1, figsize=(4, 8), dpi=300)
  1551. # ============== Panel A: Spiking Dataset ==============
  1552. # Top 5 features
  1553. axes[0].bar(
  1554. range(5),
  1555. spiking_importances["mean"]*100,
  1556. color=plot_color,
  1557. linewidth=0
  1558. )
  1559. axes[0].set_ylabel("Feature importance (%)", fontsize=7)
  1560. axes[0].set_title("A", fontsize=10, fontweight='bold', loc='left', x=-0.125, y=1)
  1561. axes[0].set_xticks(range(5))
  1562. axes[0].set_xticklabels(spiking_top5, rotation=45, ha='right', fontsize=7)
  1563. axes[0].tick_params(axis='y', labelsize=7)
  1564. axes[0].spines['top'].set_visible(False)
  1565. axes[0].spines['right'].set_visible(False)
  1566. # ============== Panel B: Passive Dataset ==============
  1567. # Top 5 features
  1568. axes[1].bar(
  1569. range(5),
  1570. passive_importances["mean"]*100,
  1571. color=plot_color,
  1572. linewidth=0
  1573. )
  1574. axes[1].set_ylabel("Feature importance (%)", fontsize=7)
  1575. axes[1].set_title("B", fontsize=10, fontweight='bold', loc='left', x=-0.125, y=1)
  1576. axes[1].set_xticks(range(5))
  1577. axes[1].set_xticklabels(passive_top5, rotation=45, ha='right', fontsize=7)
  1578. axes[1].tick_params(axis='y', labelsize=7)
  1579. axes[1].spines['top'].set_visible(False)
  1580. axes[1].spines['right'].set_visible(False)
  1581. # ============== Panel C: mEPSC Dataset ==============
  1582. # Top 5 features
  1583. axes[2].bar(
  1584. range(5),
  1585. mEPSC_importances["mean"]*100,
  1586. color=plot_color,
  1587. linewidth=0
  1588. )
  1589. axes[2].set_ylabel("Feature importance (%)", fontsize=7)
  1590. axes[2].set_title("C", fontsize=10, fontweight='bold', loc='left', x=-0.125, y=1)
  1591. axes[2].set_xticks(range(5))
  1592. axes[2].set_xticklabels(mEPSC_top5, rotation=45, ha='right', fontsize=7)
  1593. axes[2].tick_params(axis='y', labelsize=7)
  1594. axes[2].spines['top'].set_visible(False)
  1595. axes[2].spines['right'].set_visible(False)
  1596. # Set the same y-axis limits for all three panels
  1597. for ax in axes:
  1598. ax.set_ylim(0, 35) # Add 5% padding at the top
  1599. plt.tight_layout()
  1600. plt.savefig(base_path + 'figure4.svg', bbox_inches='tight')
  1601. plt.show()
  1602. # %% [markdown]
  1603. # #### **Data Types Figure**
  1604. # %%
  1605. fig2_X_mEPSC_diestrus = read_abf(raw_mEPSC_diestrus[0])
  1606. fig2_X_mEPSC_estrus = read_abf(raw_mEPSC_estrus[0])
  1607. fig2_X_mEPSC_Eproestrus = read_abf(raw_mEPSC_Eproestrus[0])
  1608. fig2_X_mEPSC_Lproestrus = read_abf(raw_mEPSC_Lproestrus[0])
  1609. fig2_X_spike_diestrus = read_abf(raw_spike_diestrus[0])
  1610. fig2_X_spike_estrus = read_abf(raw_spike_estrus[0])
  1611. fig2_X_spike_Eproestrus = read_abf(raw_spike_Eproestrus[0])
  1612. fig2_X_spike_Lproestrus = read_abf(raw_spike_Lproestrus[0])
  1613. fig2_X_passive_diestrus = read_abf(raw_passive_diestrus[0])
  1614. fig2_X_passive_estrus = read_abf(raw_passive_estrus[0])
  1615. fig2_X_passive_Eproestrus = read_abf(raw_passive_Eproestrus[0])
  1616. fig2_X_passive_Lproestrus = read_abf(raw_passive_Lproestrus[0])
  1617. # %%
  1618. abf3 = pyabf.ABF(raw_spike_diestrus[0])
  1619. abf3.dataPointsPerMs*1000
  1620. # %%
  1621. abf = pyabf.ABF(raw_mEPSC_diestrus[0])
  1622. abf.dataPointsPerMs*1000
  1623. # %%
  1624. def filter_signal(current, freq=20000):
  1625. raw_signal = current
  1626. adjusted_signal = raw_signal - np.median(raw_signal)
  1627. b_lowpass, a_lowpass = signal.bessel(4, 800, 'low', analog=False, norm='phase', fs=freq)
  1628. b_notch, a_notch = signal.iirnotch(60.0, 30.0, fs=freq)
  1629. b_multiband = signal.convolve(b_lowpass, b_notch)
  1630. a_multiband = signal.convolve(a_lowpass, a_notch)
  1631. filtered_signal = signal.filtfilt(b_multiband, a_multiband, adjusted_signal)
  1632. return filtered_signal
  1633. # %%
  1634. filtered = filter_signal(fig2_X_mEPSC_diestrus)
  1635. # %%
  1636. abf4 = pyabf.ABF(raw_spike_diestrus[0])
  1637. x = len(fig2_X_mEPSC_diestrus)/abf4.sweepCount
  1638. # %%
  1639. sig = fig2_X_spike_diestrus[800000:835000]
  1640. # %%
  1641. abf2 = pyabf.ABF(raw_passive_diestrus[0])
  1642. # %%
  1643. len(abf2.sweepX)
  1644. # %%
  1645. # Data figure - 3 side by side plots
  1646. fig, ax = plt.subplots(nrows=2, ncols=2, figsize=(16, 12), dpi=300) # Changed to 2x2 grid and adjusted figsize
  1647. plt.style.use('seaborn-v0_8-darkgrid')
  1648. plot_color = "#2166ac"
  1649. # Panel A: Spiking plot
  1650. spike_data = fig2_X_spike_diestrus[800000:835000]
  1651. time_spike = np.arange(len(spike_data)) / abf3.dataPointsPerMs
  1652. ax[0,0].plot(time_spike, spike_data, color=plot_color, linewidth=1.0)
  1653. ax[0,0].set_ylabel("Voltage (mV)", fontsize=13)
  1654. ax[0,0].set_xlabel("Time (ms)", fontsize=13)
  1655. ax[0,0].spines['top'].set_visible(False)
  1656. ax[0,0].spines['right'].set_visible(False)
  1657. # Panel B: Sweep data
  1658. for sweep in abf2.sweepList:
  1659. abf2.setSweep(sweep)
  1660. time_sweep = abf2.sweepX[:40000] * 1000
  1661. ax[0,1].plot(time_sweep, abf2.sweepY[:40000], color=plot_color,
  1662. linewidth=0.5, alpha=0.7)
  1663. ax[0,1].set_ylabel("Voltage (mV)", fontsize=13)
  1664. ax[0,1].set_xlabel("Time (ms)", fontsize=13)
  1665. ax[0,1].spines['top'].set_visible(False)
  1666. ax[0,1].spines['right'].set_visible(False)
  1667. # Panel C: Filtered current
  1668. current_data = filtered[235000:255000]
  1669. time_current = np.arange(len(current_data)) / abf.dataPointsPerMs
  1670. ax[1,0].plot(time_current, current_data, color=plot_color, linewidth=1.0)
  1671. ax[1,0].set_ylabel("Current (pA)", fontsize=13)
  1672. ax[1,0].set_xlabel("Time (ms)", fontsize=13)
  1673. ax[1,0].spines['top'].set_visible(False)
  1674. ax[1,0].spines['right'].set_visible(False)
  1675. #Panel D: Burst metrics
  1676. last_burst = spiking_diestrus.iloc[1]
  1677. spike_data = fig2_X_spike_diestrus[800000:835000]
  1678. time_spike = np.arange(len(spike_data)) / abf3.dataPointsPerMs
  1679. ax[1,1].plot(time_spike, spike_data, color=plot_color, linewidth=1.0)
  1680. ax[1,1].set_ylabel("Voltage (mV)", fontsize=13)
  1681. ax[1,1].set_xlabel("Time (ms)", fontsize=13)
  1682. ax[1,1].spines['top'].set_visible(False)
  1683. ax[1,1].spines['right'].set_visible(False)
  1684. # --- Panel D annotations ---
  1685. burst_start_ms = last_burst['burst_start'] * 1000 - (800000 / abf3.dataPointsPerMs)
  1686. burst_end_ms = last_burst['burst_end'] * 1000 - (800000 / abf3.dataPointsPerMs)
  1687. peak_val = last_burst['avg_spike_peaks'] # 29.38 mV
  1688. amp_val = last_burst['avg_spike_amplitude'] # 81.51 mV
  1689. base_val = peak_val - amp_val # true amplitude foot (~ -52 mV)
  1690. # --- 1. Burst duration ---
  1691. plateau_val = np.percentile(spike_data, 65)
  1692. y_duration = plateau_val
  1693. ax[1,1].annotate('', xy=(burst_end_ms + 30, y_duration), xytext=(burst_start_ms - 20, y_duration),
  1694. arrowprops=dict(arrowstyle='<->', color='black', lw=1.2))
  1695. ax[1,1].text((burst_start_ms + burst_end_ms) / 2, y_duration - 3,
  1696. f"Burst duration\n{last_burst['burst_length']:.2f} s",
  1697. ha='center', va='top', fontsize=12)
  1698. # --- 2. Avg spike amplitude (bracket now spans peak -> peak - amplitude) ---
  1699. bracket_x = burst_start_ms - 200 # sits left of the rise, over flat baseline
  1700. tick_width = 8
  1701. ax[1,1].plot([bracket_x, bracket_x], [base_val + 5, peak_val], color='black', lw=1.2)
  1702. ax[1,1].plot([bracket_x, bracket_x + tick_width], [peak_val, peak_val], color='black', lw=1.2)
  1703. ax[1,1].plot([bracket_x, bracket_x + tick_width], [base_val + 5, base_val + 5], color='black', lw=1.2)
  1704. ax[1,1].text(bracket_x - 12, (peak_val + base_val) / 2,
  1705. f"Average spike\namplitude\n{amp_val:.2f} mV",
  1706. color='black', fontsize=12, va='center', ha='right')
  1707. # --- 3. Avg spike peak value (dashed line only across the spike train) ---
  1708. pad = 25
  1709. ax[1,1].plot([burst_start_ms - pad, burst_end_ms + pad], [peak_val, peak_val],
  1710. color='black', linestyle='--', linewidth=0.8)
  1711. ax[1,1].text(burst_end_ms + pad + 20, peak_val,
  1712. f"Average spike\npeak value\n{peak_val:.2f} mV",
  1713. ha='left', va='center', fontsize=12, color='black')
  1714. # --- 4. Spikes per burst ---
  1715. ax[1,1].text(burst_end_ms + pad + 20, -5,
  1716. f"Spikes per burst = 5",
  1717. ha='left', va='center', fontsize=12, color='black')
  1718. # Add panel labels
  1719. for i, label in enumerate(['A', 'B', 'C', 'D']):
  1720. row = i // 2
  1721. col = i % 2
  1722. ax[row, col].text(-0.05, 1.05, label, transform=ax[row, col].transAxes,
  1723. fontsize=16, fontweight='bold', va='top', ha='right')
  1724. plt.tight_layout()
  1725. plt.savefig(base_path + 'figure1.tif', bbox_inches='tight')
  1726. plt.show()
  1727. # %% [markdown]
  1728. # #### **t-SNE Figure**
  1729. # %%
  1730. plt.style.use('default')
  1731. views = [(30, 45), (15, 180), (45, 270)]
  1732. zoom = 1.2
  1733. data_sets = [
  1734. (spiking_tSNE_comp1_3d, spiking_tSNE_comp2_3d, spiking_tSNE_comp3_3d, tSNE_y_spiking),
  1735. (passive_tSNE_comp1_3d, passive_tSNE_comp2_3d, passive_tSNE_comp3_3d, tSNE_y_passive),
  1736. (mEPSC_tSNE_comp1_3d, mEPSC_tSNE_comp2_3d, mEPSC_tSNE_comp3_3d, tSNE_y_mEPSC)
  1737. ]
  1738. row_titles = ['A', 'B', 'C']
  1739. label_style = {
  1740. (30, 45): {'x': {'pad': -6, 'size': 9}, 'y': {'pad': -6, 'size': 9}, 'z': {'pad': -10, 'size': 9}},
  1741. (15, 180): {'x': {'pad': -8, 'size': 7}, 'y': {'pad': -2, 'size': 10}, 'z': {'pad': -10, 'size': 9}},
  1742. (45, 270): {'x': {'pad': -6, 'size': 9}, 'y': {'pad': -6, 'size': 9}, 'z': {'pad': -6, 'size': 9}},
  1743. }
  1744. fig = plt.figure(figsize=(22, 18), dpi=300)
  1745. fig.patch.set_facecolor('white')
  1746. plt.subplots_adjust(wspace=0.15, hspace=0.15)
  1747. scatter = None
  1748. class_colors = ['#E63946', '#457B9D', '#2A9D8F', '#E76F51']
  1749. cmap = mcolors.ListedColormap(class_colors)
  1750. for row, (comp1, comp2, comp3, labels) in enumerate(data_sets):
  1751. for col, (elev, azim) in enumerate(views):
  1752. ax = fig.add_subplot(3, 3, row * 3 + col + 1, projection='3d')
  1753. ax.set_facecolor('white')
  1754. ax.set_box_aspect(None, zoom=zoom)
  1755. scatter = ax.scatter(
  1756. comp1, comp2, comp3,
  1757. c=labels,
  1758. cmap=cmap,
  1759. alpha=1,
  1760. s=6,
  1761. vmin=0, vmax=3,
  1762. edgecolor='w',
  1763. linewidth=0.25
  1764. )
  1765. ax.set_xticklabels([])
  1766. ax.set_yticklabels([])
  1767. ax.set_zticklabels([])
  1768. ax.tick_params(axis='both', which='both', length=0)
  1769. ax.view_init(elev=elev, azim=azim)
  1770. ax.xaxis.pane.fill = False
  1771. ax.yaxis.pane.fill = False
  1772. ax.zaxis.pane.fill = False
  1773. ax.xaxis.pane.set_edgecolor('lightgrey')
  1774. ax.yaxis.pane.set_edgecolor('lightgrey')
  1775. ax.zaxis.pane.set_edgecolor('lightgrey')
  1776. ax.grid(True, color='lightgrey', linewidth=0.5)
  1777. if col == 0:
  1778. ax.set_title(row_titles[row], fontsize=20, fontweight='bold', loc='left', pad=2)
  1779. legend_labels = ['Diestrus', 'Estrus', 'Early proestrus', 'Late proestrus']
  1780. legend_handles = [
  1781. plt.Line2D([0], [0], marker='o', color='w', markerfacecolor=c,
  1782. markersize=9, label=l)
  1783. for c, l in zip(class_colors, legend_labels)
  1784. ]
  1785. fig.legend(
  1786. handles=legend_handles,
  1787. labels=legend_labels,
  1788. fontsize=20,
  1789. title_fontsize=0,
  1790. frameon=False,
  1791. loc='center right',
  1792. bbox_to_anchor=(1.02, 0.5),
  1793. handletextpad=0.5,
  1794. labelspacing=0.8,
  1795. )
  1796. plt.savefig(base_path + 'figure5.svg', bbox_inches='tight')
  1797. plt.show()
  1798. # %% [markdown]
  1799. # ### **Stats**
  1800. # %% [markdown]
  1801. # #### **AP vs. passive vs. mEPSC accuracies**
  1802. # %%
  1803. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  1804. all_accuracies_data = []
  1805. # Populate with mEPSC seed accuracies (individual values from each seed)
  1806. for name in model_names:
  1807. for acc in mEPSC_seed_accuracies[name]:
  1808. all_accuracies_data.append({
  1809. "Accuracy": acc,
  1810. "Feature_Set": "mEPSC",
  1811. "Model": name
  1812. })
  1813. # Populate with spiking seed accuracies
  1814. for name in model_names:
  1815. for acc in spiking_seed_accuracies[name]:
  1816. all_accuracies_data.append({
  1817. "Accuracy": acc,
  1818. "Feature_Set": "AP",
  1819. "Model": name
  1820. })
  1821. # Populate with passive seed accuracies
  1822. for name in model_names:
  1823. for acc in passive_seed_accuracies[name]:
  1824. all_accuracies_data.append({
  1825. "Accuracy": acc,
  1826. "Feature_Set": "Passive",
  1827. "Model": name
  1828. })
  1829. accuracies_df = pd.DataFrame(all_accuracies_data)
  1830. # Drop rows with NaN or inf values in the 'Accuracy' column just in case
  1831. accuracies_df = accuracies_df.replace([np.inf, -np.inf], np.nan).dropna(subset=['Accuracy'])
  1832. formula = "Accuracy ~ C(Feature_Set) * C(Model)"
  1833. model = smf.ols(formula, data=accuracies_df).fit()
  1834. table = sm.stats.anova_lm(model, typ=2)
  1835. display(table)
  1836. # %%
  1837. df_grouped = accuracies_df.copy()
  1838. df_grouped['group'] = df_grouped['Feature_Set'] + '_' + df_grouped['Model']
  1839. tukey = pairwise_tukeyhsd(df_grouped['Accuracy'], df_grouped['group'])
  1840. tukey_df = pd.DataFrame(data=tukey._results_table.data[1:],
  1841. columns=tukey._results_table.data[0])
  1842. tukey_df['p-adj'] = tukey.pvalues #To get precise values
  1843. pd.set_option('display.float_format', lambda x: f'{x:.12f}')
  1844. tukey_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/MC for accuracies.csv")
  1845. # %% [markdown]
  1846. # #### **Across full vs. subsetted dataset**
  1847. # %%
  1848. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  1849. SEEDS = [42, 7, 13, 21, 99]
  1850. # Prepare data for ANOVA
  1851. all_accuracies_data = []
  1852. # Full dataset accuracies
  1853. for model_name in model_names:
  1854. for seed_acc in mEPSC_seed_accuracies[model_name]:
  1855. all_accuracies_data.append({
  1856. 'Accuracy': seed_acc,
  1857. 'Dataset_Type': 'Full',
  1858. 'Feature_Set': 'mEPSC',
  1859. 'Model': model_name
  1860. })
  1861. for seed_acc in spiking_seed_accuracies[model_name]:
  1862. all_accuracies_data.append({
  1863. 'Accuracy': seed_acc,
  1864. 'Dataset_Type': 'Full',
  1865. 'Feature_Set': 'Spiking',
  1866. 'Model': model_name
  1867. })
  1868. for seed_acc in passive_seed_accuracies[model_name]:
  1869. all_accuracies_data.append({
  1870. 'Accuracy': seed_acc,
  1871. 'Dataset_Type': 'Full',
  1872. 'Feature_Set': 'Passive',
  1873. 'Model': model_name
  1874. })
  1875. # Even label count dataset accuracies
  1876. for model_name in model_names:
  1877. for seed_acc in mEPSC_even_seed_accuracies[model_name]:
  1878. all_accuracies_data.append({
  1879. 'Accuracy': seed_acc,
  1880. 'Dataset_Type': 'Even_Label',
  1881. 'Feature_Set': 'mEPSC',
  1882. 'Model': model_name
  1883. })
  1884. for seed_acc in spiking_even_seed_accuracies[model_name]:
  1885. all_accuracies_data.append({
  1886. 'Accuracy': seed_acc,
  1887. 'Dataset_Type': 'Even_Label',
  1888. 'Feature_Set': 'Spiking',
  1889. 'Model': model_name
  1890. })
  1891. for seed_acc in passive_even_seed_accuracies[model_name]:
  1892. all_accuracies_data.append({
  1893. 'Accuracy': seed_acc,
  1894. 'Dataset_Type': 'Even_Label',
  1895. 'Feature_Set': 'Passive',
  1896. 'Model': model_name
  1897. })
  1898. accuracies_df = pd.DataFrame(all_accuracies_data)
  1899. formula = 'Accuracy ~ C(Dataset_Type) * C(Feature_Set) * C(Model)'
  1900. model = smf.ols(formula, data=accuracies_df).fit()
  1901. anova_table = sm.stats.anova_lm(model, typ=2) # Type 2 ANOVA for unbalanced designs
  1902. print("\nANOVA Results:")
  1903. display(anova_table)
  1904. # %%
  1905. df_grouped = accuracies_df.copy()
  1906. df_grouped['group'] = df_grouped['Dataset_Type'] + '_' + df_grouped['Feature_Set']
  1907. tukey = pairwise_tukeyhsd(df_grouped['Accuracy'], df_grouped['group'])
  1908. print(tukey.summary())
  1909. # %% [markdown]
  1910. # #### **Across randomized vs. non-randomized**
  1911. # %%
  1912. model_names = ['RFC', 'GBC', 'LR', 'KNN', 'MLP', 'SVC', 'DTC']
  1913. SEEDS = [42, 7, 13, 21, 99]
  1914. # Prepare data for ANOVA
  1915. all_accuracies_rnd_data = []
  1916. # Full dataset accuracies — real labels
  1917. for model_name in model_names:
  1918. for seed_acc in mEPSC_seed_accuracies[model_name]:
  1919. all_accuracies_rnd_data.append({
  1920. 'Accuracy': seed_acc,
  1921. 'Label_Type': 'Real',
  1922. 'Feature_Set': 'mEPSC',
  1923. 'Model': model_name
  1924. })
  1925. for seed_acc in spiking_seed_accuracies[model_name]:
  1926. all_accuracies_rnd_data.append({
  1927. 'Accuracy': seed_acc,
  1928. 'Label_Type': 'Real',
  1929. 'Feature_Set': 'Spiking',
  1930. 'Model': model_name
  1931. })
  1932. for seed_acc in passive_seed_accuracies[model_name]:
  1933. all_accuracies_rnd_data.append({
  1934. 'Accuracy': seed_acc,
  1935. 'Label_Type': 'Real',
  1936. 'Feature_Set': 'Passive',
  1937. 'Model': model_name
  1938. })
  1939. # Randomized label control accuracies
  1940. for model_name in model_names:
  1941. for seed_acc in mEPSC_rnd_seed_accuracies[model_name]:
  1942. all_accuracies_rnd_data.append({
  1943. 'Accuracy': seed_acc,
  1944. 'Label_Type': 'Randomized',
  1945. 'Feature_Set': 'mEPSC',
  1946. 'Model': model_name
  1947. })
  1948. for seed_acc in spiking_rnd_seed_accuracies[model_name]:
  1949. all_accuracies_rnd_data.append({
  1950. 'Accuracy': seed_acc,
  1951. 'Label_Type': 'Randomized',
  1952. 'Feature_Set': 'Spiking',
  1953. 'Model': model_name
  1954. })
  1955. for seed_acc in passive_rnd_seed_accuracies[model_name]:
  1956. all_accuracies_rnd_data.append({
  1957. 'Accuracy': seed_acc,
  1958. 'Label_Type': 'Randomized',
  1959. 'Feature_Set': 'Passive',
  1960. 'Model': model_name
  1961. })
  1962. accuracies_rnd_df = pd.DataFrame(all_accuracies_rnd_data)
  1963. formula = 'Accuracy ~ C(Label_Type) * C(Feature_Set) * C(Model)'
  1964. rnd_anova_model = smf.ols(formula, data=accuracies_rnd_df).fit()
  1965. anova_rnd_table = sm.stats.anova_lm(rnd_anova_model, typ=2) # Type 2 for unbalanced designs
  1966. display(anova_rnd_table)
  1967. # %% [markdown]
  1968. # #### **Across Phase and Feature Set in Classification Reports**
  1969. # %%
  1970. all_accuracies_data = []
  1971. for name in model_names:
  1972. rep = mEPSC_classification_reports[name]
  1973. iter = 0
  1974. for key, value in rep.items():
  1975. if iter > 3:
  1976. break
  1977. all_accuracies_data.append({
  1978. "Recall": value['recall'], #recall
  1979. "Model": name,
  1980. "Phase": key,
  1981. "Feature_Set": "mEPSC"
  1982. })
  1983. iter += 1
  1984. for name in model_names:
  1985. rep = spiking_classification_reports[name]
  1986. iter = 0
  1987. for key, value in rep.items():
  1988. if iter > 3:
  1989. break
  1990. all_accuracies_data.append({
  1991. "Recall": value['recall'], #recall
  1992. "Model": name,
  1993. "Phase": key,
  1994. "Feature_Set": "AP"
  1995. })
  1996. iter += 1
  1997. for name in model_names:
  1998. rep = passive_classification_reports[name]
  1999. iter = 0
  2000. for key, value in rep.items():
  2001. if iter > 3:
  2002. break
  2003. all_accuracies_data.append({
  2004. "Recall": value['recall'], #recall
  2005. "Model": name,
  2006. "Phase": key,
  2007. "Feature_Set": "Passive"
  2008. })
  2009. iter += 1
  2010. accuracies_df = pd.DataFrame(all_accuracies_data)
  2011. accuracies_df = accuracies_df.replace([np.inf, -np.inf], np.nan).dropna(subset=['Recall'])
  2012. formula = "Recall ~ C(Phase) * C(Feature_Set) + C(Phase) * C(Model) + C(Feature_Set) * C(Model)"
  2013. model = smf.ols(formula, data=accuracies_df).fit()
  2014. table = sm.stats.anova_lm(model, typ=2)
  2015. display(table)
  2016. # %%
  2017. df_grouped = accuracies_df.copy()
  2018. df_grouped['group'] = df_grouped['Phase'] + '_' + df_grouped['Feature_Set']
  2019. tukey = pairwise_tukeyhsd(df_grouped['Recall'], df_grouped['group'])
  2020. tukey_df = pd.DataFrame(data=tukey._results_table.data[1:],
  2021. columns=tukey._results_table.data[0])
  2022. tukey_df['p-adj'] = tukey.pvalues
  2023. pd.set_option('display.float_format', lambda x: f'{x:.12f}')
  2024. tukey_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/MC for recall.csv")
  2025. print(tukey.summary())
  2026. # %%
  2027. df_grouped = accuracies_df.copy()
  2028. df_grouped['group'] = df_grouped['Feature_Set'] + '_' + df_grouped['Model']
  2029. tukey = pairwise_tukeyhsd(df_grouped['Accuracy'], df_grouped['group'])
  2030. tukey_df = pd.DataFrame(data=tukey._results_table.data[1:],
  2031. columns=tukey._results_table.data[0])
  2032. tukey_df['p-adj'] = tukey.pvalues
  2033. pd.set_option('display.float_format', lambda x: f'{x:.12f}')
  2034. tukey_df.to_csv("/content/drive/MyDrive/Estrous Classification Project/MC for accuracies.csv")

Estrous_Classification.ipynb at commit 8495b02, under MIT · at the source

Overview

Authors: Armaan Raina1, John Meitzen1,2,3
ORCID iDs: John Meitzen
  1. Department of Biological Sciences, North Carolina State University, Raleigh, NC USA
  2. Center for Human Health and the Environment, North Carolina State University, Raleigh, NC USA
  3. Dept. of Biological Sciences, NC State University, 144 David Clark Labs, Campus Box 7617, Raleigh, NC 27695-7617 USA
Institutions: North Carolina State University (United States)
Journal: Neuroinformatics, volume 24, issue 3, article 57
Dates: received 8 July 2026; accepted 25 August 2026; published online 1 September 2026; in print 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1007/s12021-026-09814-0 · PMID 42678460 · PMCID PMC13534189 · OpenAlex W7204910555
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: intracellular / patch clamp (modality), rat (organism), methods / tools (subfield)
Methods: Smoothing, state filtering, decompositions, Statistics, Machine learning, Preprocessing, Evoked potentials, Single-unit activity, calcium imaging
Keywords: machine learning, estrous cycle, nucleus accumbens, electrophysiology, neuroinformatics
MeSH: Estrous Cycle*, Machine Learning*, Models, Neurological*, Neurons*, Nucleus Accumbens*, Action Potentials, Animals, Classification Algorithms, Excitatory Postsynaptic Potentials, Female, Medium Spiny Neurons, Predictive Learning Models, Rats (* major topic)
Topic: Neural dynamics and brain function (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: NIEHS NIH HHS (P30ES025128, P30 ES025128); North Carolina State University (Provost Professional Experience Award); National Institute of Environmental Health Sciences (P30ES025128)
Citations: not cited yet (Europe PMC); 73 references in the paper

Abstract

Machine learning models (MLMs) have been used to classify neuron subtypes based upon electrophysiological attributes. What has been less explored is using MLMs to assess differences within a neuron type, especially in the context of subtle changes induced by neuromodulation such as the rodent estrous cycle. Previous research found that the estrous cycle shifts rat nucleus accumbens (NAc) medium spiny neuron (MSN) electrophysiology, including action potential (AP), passive, and miniature excitatory post-synaptic current (mEPSC) properties. This plasticity provides a model system to investigate this question. We hypothesized that MLM classification accuracy would differ when trained across these data types, reflecting the information each encodes regarding estrous phase identity. To test this hypothesis, we extracted electrophysiological features across four estrous cycle phases from a publicly available dataset, and employed this data to train MLMs to classify estrous cycle phase origin. We found: MLMs identified estrous phase origin with up to 94% accuracy (Random Forest); MLM accuracy differed by model and feature set; passive, AP and mEPSC features yielded sequential performances with feature importance analysis revealing input resistance, AP burst length, and mEPSC baseline current as most discriminative for their respective feature sets; the late proestrus phase demonstrated the most distinct electrophysiological profile. These findings demonstrate that MLMs can detect changes associated with the estrous cycle and are useful for assessing electrophysiological differences within a neuron type. They further represent a broader application of MLMs as computational tools useful for understanding the sensitivity of neural properties to the influences of neuromodulatory cycles.

Supplementary Information: The online version contains supplementary material available at https://doi.org/10.1007/s12021-026-09814-0.

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

doi:10.5061/dryad.k0p2ngfgm

License: none: the authors keep all their rights
State: the link answers, verified on 26 September 2026
Evidence: the link answers
Software Heritage: not checked
Found in: “Data Availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 26 September 2026: the link answers (HTTP 200)
  • 26 September 2026: the link answers (HTTP 200)

Armaan-Raina/Estrous-Phase-Classification

License: MIT
State: the link answers, verified on 26 September 2026
Evidence: files inventoried
Commit: 8495b02c0cd77faeb25c2218dff6c0b2a9311508, 25 August 2026
Languages: Jupyter (1)
Size: 3 files, 1 script
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, license file, 1 notebook
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (1 file), NetworkX (1 file), NumPy (1 file), pandas (1 file), Pillow (1 file), pyABF (1 file), scikit-learn (1 file), SciPy (1 file), seaborn (1 file), statsmodels (1 file)
Availability: 1 check, the latest on 26 September 2026: the link answers
  • 26 September 2026: the link answers
3 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;
  • 1 script, each with its path and the digest of its content;
  • 4 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

The parent electrophysiological dataset used in this study is available at the Dryad Data Repository (https://datadryad.org/dataset/doi:10.5061/dryad.k0p2ngfgm). The code generated during this study to analyze the parent dataset is available in the GitHub repository (https://github.com/Armaan-Raina/Estrous-Phase-Classification).

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 3, 28 September 2026

  • Publisher: n/a → Springer Science+Business Media

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 2 authors, 5 keywords, 13 MeSH terms, 3 funders, 69 references.

Cite

This paper

Raina, A., & Meitzen, J. (2026). Application of Machine Learning Models to Identify Differences in Neural Electrophysiological Properties Across Estrous Cycle Phases. Neuroinformatics, 24(3), 57. https://doi.org/10.1007/s12021-026-09814-0

BibTeX

@article{raina2026application,
author = {Raina, Armaan and Meitzen, John},
title = {{Application of Machine Learning Models to Identify Differences in Neural Electrophysiological Properties Across Estrous Cycle Phases}},
journal = {Neuroinformatics},
year = {2026},
month = sep,
volume = {24},
number = {3},
pages = {57},
publisher = {Springer Science+Business Media},
issn = {1539-2791},
doi = {10.1007/s12021-026-09814-0},
url = {https://doi.org/10.1007/s12021-026-09814-0},
pmid = {42678460},
pmcid = {PMC13534189}
}

RIS

TY - JOUR
AU - Raina, Armaan
AU - Meitzen, John
TI - Application of Machine Learning Models to Identify Differences in Neural Electrophysiological Properties Across Estrous Cycle Phases
T2 - Neuroinformatics
J2 - Neuroinformatics
PY - 2026
DA - 2026/09/01
VL - 24
IS - 3
SP - 57
SN - 1539-2791
PB - Springer Science+Business Media
DO - 10.1007/s12021-026-09814-0
UR - https://doi.org/10.1007/s12021-026-09814-0
LA - en
ER -

CSL-JSON

{
"id": "10.1007/s12021-026-09814-0",
"type": "article-journal",
"title": "Application of Machine Learning Models to Identify Differences in Neural Electrophysiological Properties Across Estrous Cycle Phases",
"container-title": "Neuroinformatics",
"author": [
{
"family": "Raina",
"given": "Armaan"
},
{
"family": "Meitzen",
"given": "John"
}
],
"container-title-short": "Neuroinformatics",
"volume": "24",
"issue": "3",
"page": "57",
"DOI": "10.1007/s12021-026-09814-0",
"PMID": "42678460",
"PMCID": "PMC13534189",
"ISSN": "1539-2791",
"publisher": "Springer Science+Business Media",
"URL": "https://doi.org/10.1007/s12021-026-09814-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
9,
1
]
]
}
}

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-76939-w [code]
HIPPIE: a generative model for electrophysiological analysis across species, technologies, and modalities.
Journal: Nature communications
In common: Pillow, statsmodels, seaborn, 5 other tools, rat, methods / tools, 1 reference
[2] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: NetworkX, Pillow, statsmodels, 6 other tools
[3] doi:10.3390/ijms27167275 [code]
Integrative Multi-Omics Analysis of Multiple Sclerosis Reveals Cell-Type-Specific Regulatory Landscapes and Discordant Methylation-Expression Coupling.
Journal: International journal of molecular sciences
In common: NetworkX, Pillow, statsmodels, 6 other tools
[4] doi:10.1371/journal.pcbi.1014571 [code]
SynAPSeg: A novel dataset and image analysis framework for deep learning-based synapse detection and quantification.
Journal: PLoS computational biology
In common: NetworkX, Pillow, statsmodels, 6 other tools
[5] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: NetworkX, Pillow, statsmodels, 6 other tools
[6] doi:10.1038/s41467-026-74357-6 [code]
Hippocampo-neocortical interaction as compressive retrieval-augmented generation.
Journal: Nature communications
In common: NetworkX, Pillow, statsmodels, 6 other tools
[7] doi:10.1016/j.isci.2026.116206 [code]
Gut distension evokes rapid neural dynamics in vagal and hindbrain populations of larval zebrafish.
Journal: iScience
In common: NetworkX, Pillow, statsmodels, 6 other tools
[8] doi:10.3389/fncom.2026.1786996 [code]
Schumann-anchored golden ratio organization of human neural oscillations.
Journal: Frontiers in computational neuroscience
In common: NetworkX, Pillow, statsmodels, 6 other tools
[9] doi:10.1016/j.celrep.2026.117420 [code]
Neural population dynamics of direct electrical stimulation of neocortex.
Journal: Cell reports
In common: NetworkX, Pillow, statsmodels, 6 other tools
[10] doi: [code]
Naturalistic behavior and self-generated neural activity predictive of self-correction
Journal: bioRxiv : the preprint server for biology
In common: NetworkX, Pillow, statsmodels, 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.