-
Notifications
You must be signed in to change notification settings - Fork 12
Expand file tree
/
Copy pathParallelRAMSimulationsEFR.py
More file actions
305 lines (248 loc) · 10.5 KB
/
Copy pathParallelRAMSimulationsEFR.py
File metadata and controls
305 lines (248 loc) · 10.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
#!/usr/bin/env python3
"""
RAM EFR Analysis Script - Parallel Processing Pipeline
This script provides a parallel processing pipeline for RAM EFR analysis:
1. Processes multiple subjects with custom audiogram data
2. Converts audiogram data to Poles using OHC_ind function for each subject
3. Runs simulations with user-specified ANF distributions (HSR, MSR, LSR)
4. Calculates EFR values for each simulation
5. Saves results to CSV file with subject names and EFR values
Usage:
Configure your settings in the main block and run the script.
The script will process all subjects in parallel using all available CPU cores.
Created by: Brent Nissens
Date: October 22, 2025
"""
import numpy as np
import os
import sys
import pandas as pd
import multiprocessing as mp
import time
from tqdm import tqdm
from concurrent.futures import ProcessPoolExecutor, as_completed
from get_RAM_stims import get_RAM_stims
from model2018 import model2018
import OHC_ind
def load_shera_poles_profile(sheraP_folder):
"""
Load Shera poles profile from a specific path.
"""
try:
sheraP_data = np.loadtxt(sheraP_folder + '/StartingPoles.dat')
# Take only the first row if there are multiple rows
if sheraP_data.ndim > 1:
sheraP = sheraP_data[0, :]
else:
sheraP = sheraP_data
return sheraP
except Exception as e:
print(f"Error loading Shera poles profile from {sheraP_folder}: {e}")
return None
def run_simulation(stim, fs, sheraP, HSR, MSR, LSR):
"""
Run a simulation with the given parameters.
"""
try:
output = model2018(stim, fs, 'abr', 1, 'evihmlbw', 1, sheraP, 0.05, 'vel', HSR, MSR, LSR, 1, os.getcwd())
return output
except Exception as e:
print(f"Error running simulation: {e}")
return None
def calculate_EFR(output):
"""
Calculate the EFR from the output.
"""
# This function computes the EFR (Envelope Following Response) harmonics sum from simulation output.
# It mirrors get_RAM_EFRS1, but omits plotting and converts to microV.
try:
# Handle the case where output is a list containing a dictionary
if isinstance(output, list) and len(output) > 0:
output = output[0] # Get the first (and likely only) element
# Extract relevant waveforms and sampling frequency from the output object/struct.
fs = float(output.fs_an)
w1 = output.w1.flatten()
w3 = output.w3.flatten()
w5 = output.w5.flatten()
EFR = w1 + w3 + w5
# Fourier Transform
L = len(EFR)
Y = np.fft.fft(EFR)
P2 = np.abs(Y / L)
P1 = P2[:L//2 + 1]
P1[1:-1] = 2 * P1[1:-1]
f = fs * np.arange(L//2 + 1) / L
# Find indices for 4 harmonics of 110 Hz
fundamental = 110 # Hz
num_harmonics = 4
harmonics = np.arange(1, num_harmonics + 1) * fundamental
idx = []
for harmonic in harmonics:
idx.append(np.argmin(np.abs(f - harmonic)))
# Calculate sum and convert to microV
harmonic_sum = np.sum(P1[idx]) * 1e6
return harmonic_sum
except Exception as e:
print(f"Error calculating EFR: {e}")
return np.nan
def process_single_subject(args):
"""
Process a single subject's simulation.
Parameters:
-----------
args : tuple
(subject_data, HSR, MSR, LSR, stim, fs, poles_output_dir)
where subject_data is (idx, subject_name, hl_freqs_hz, hl_db)
HSR, MSR, LSR can be scalars or arrays (frequency-dependent ANF distributions)
Returns:
--------
dict
Result row with subject name and EFR value
"""
subject_data, HSR, MSR, LSR, stim, fs, poles_output_dir = args
idx, subject_name, hl_freqs_hz, hl_db = subject_data
try:
# Create poles using OHC_ind (without showing figures)
OHC_ind.ohc_ind(
name=subject_name,
hl_freqs_hz=hl_freqs_hz,
hl_db=hl_db,
base_dir=poles_output_dir,
show_figs=False
)
# Load the poles profile
subject_poles_path = os.path.join(poles_output_dir, 'Poles', subject_name)
sheraP = load_shera_poles_profile(subject_poles_path)
if sheraP is None:
print(f" Could not load poles for {subject_name}, skipping...")
return None
# Run simulation with ANF distributions
# Note: HSR, MSR, LSR can be scalars (constant across frequency) or
# arrays (frequency-dependent ANF distributions)
output = run_simulation(stim, fs, sheraP, HSR, MSR, LSR)
if output is None:
return None
# Calculate EFR
efr = calculate_EFR(output)
# Store results
result_row = {
'Name': subject_name,
'EFR': efr
}
return result_row
except Exception as e:
print(f"Error processing subject {subject_name}: {e}")
return None
if __name__ == "__main__":
# ============================================================================
# CONFIGURATION - Modify these settings for your use case
# ============================================================================
# Path to Excel file containing audiogram data
# Expected format: Excel file with columns 'ID' and 'Audio_XXXHz' (e.g., 'Audio_125Hz', 'Audio_250Hz')
excel_path = './data/audiograms.xlsx'
# ANF distribution settings
# HSR, MSR, LSR can be scalars (constant across frequency) or arrays (frequency-dependent)
# For frequency-dependent distributions, provide arrays with values for each frequency channel
# Default values: HSR=13, MSR=3, LSR=3 (constant across all frequencies)
HSR = 13 # High Spontaneous Rate fibers
MSR = 3 # Medium Spontaneous Rate fibers
LSR = 3 # Low Spontaneous Rate fibers
# Poles output directory (where OHC_ind will save generated poles)
poles_output_dir = '.'
# Stimulus parameters
fs = 1e5 # Sampling frequency in Hz
fRAM = np.array([4000]) # RAM frequency in Hz
# Output settings
output_csv = 'EFR_results.csv' # Output CSV filename
# Parallel processing settings
num_workers = mp.cpu_count() # Number of parallel workers (default: all CPU cores)
# ============================================================================
# MAIN PROCESSING PIPELINE
# ============================================================================
print("="*60)
print("Starting RAM EFR analysis pipeline...")
print("="*60 + "\n")
# Load Excel data
print("="*60)
print(f"Loading data from {excel_path}")
print("="*60 + "\n")
try:
df = pd.read_excel(excel_path)
except Exception as e:
print(f"Error loading Excel file: {e}")
sys.exit(1)
# Get audiogram column names (assuming they are Audio_XXXHz)
audio_columns = [col for col in df.columns if 'Audio_' in col]
if not audio_columns:
print("Error: No audiogram columns found (expected format: 'Audio_XXXHz')")
sys.exit(1)
# Initialize stimulus
stim = get_RAM_stims(fs, fRAM)
print("="*60)
print(f"Generated RAM stimulus with fs={fs} Hz, fRAM={fRAM} Hz")
print("="*60 + "\n")
# Prepare all subject data for parallel processing
subject_data_list = []
for idx, row in df.iterrows():
subject_name = str(row.get('ID', f'Subject_{idx}'))
# Extract audiogram frequencies and values
audiogram_data = []
for col in audio_columns:
freq_hz = int(col.replace('Audio_', '').replace('Hz', ''))
hl_db = row[col]
audiogram_data.append((freq_hz, hl_db))
# Sort by frequency
audiogram_data.sort(key=lambda x: x[0])
hl_freqs_hz = [x[0] for x in audiogram_data]
hl_db = [x[1] for x in audiogram_data]
subject_data_list.append((idx, subject_name, hl_freqs_hz, hl_db))
print("="*60)
print(f"Processing {len(subject_data_list)} subjects using {num_workers} workers")
print(f"ANF distributions: HSR={HSR}, MSR={MSR}, LSR={LSR}")
print("="*60 + "\n")
# Prepare arguments for each worker
worker_args = [(sd, HSR, MSR, LSR, stim, fs, poles_output_dir) for sd in subject_data_list]
# Process subjects in parallel
start_time = time.time()
completed_subjects = 0
results = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {executor.submit(process_single_subject, args): idx for idx, args in enumerate(worker_args)}
# Collect results with progress bar
pbar = tqdm(total=len(futures), desc="Processing subjects")
for future in as_completed(futures):
result = future.result()
if result is not None:
results.append(result)
# Update progress and time estimates
completed_subjects += 1
elapsed = time.time() - start_time
avg_time_per_subject = elapsed / completed_subjects if completed_subjects > 0 else 0
remaining_subjects = len(futures) - completed_subjects
estimated_time_left = avg_time_per_subject * remaining_subjects / num_workers if num_workers > 0 else 0
# Format time estimates
elapsed_str = f"{int(elapsed // 3600)}h {int((elapsed % 3600) // 60)}m {int(elapsed % 60)}s"
eta_str = f"{int(estimated_time_left // 3600)}h {int((estimated_time_left % 3600) // 60)}m {int(estimated_time_left % 60)}s"
pbar.set_postfix({
'Elapsed': elapsed_str,
'ETA': eta_str,
'Avg/Subj': f"{avg_time_per_subject:.1f}s"
})
pbar.update(1)
pbar.close()
# Calculate total processing time
total_time = time.time() - start_time
hours = int(total_time // 3600)
minutes = int((total_time % 3600) // 60)
seconds = int(total_time % 60)
# Convert results to DataFrame and save to CSV
if results:
results_df = pd.DataFrame(results)
results_df.to_csv(output_csv, index=False)
print(f"\nResults saved to {output_csv}")
print(f"\nTotal subjects processed: {len(results)}")
print(f"Total processing time: {hours}h {minutes}m {seconds}s")
print("\nResults summary:")
print(results_df)
else:
print("\nWarning: No results were generated. Check your configuration and input data.")