prepare for new version

This commit is contained in:
2026-08-31 22:00:35 -07:00
parent a9ea0efcb4
commit f0b8266cf9
8 changed files with 613 additions and 32 deletions
+525 -21
View File
@@ -163,6 +163,42 @@ QC_METRIC_LABELS = {
}
ROI_MAP = {
'ROI_OccipitoParietal_BA18_19_7': [
'S1_D1', 'S1_D2', 'S4_D3'
],
'ROI_SuperiorParietal_BA7_5': [
'S2_D1', 'S3_D1', 'S2_D2', 'S2_D4', 'S4_D4'
],
'ROI_TPJ_AngularGyrus_BA39_19_40': [
'S4_D2', 'S5_D3', 'S4_D5', 'S5_D5'
],
# --- SPLIT BA40 REGIONS ---
'ROI_Posterior_BA40_Parietal': [
'S6_D4', 'S6_D5'
],
'ROI_Anterior_BA40_Sensorimotor': [
'S10_D5', 'S8_D5'
],
# --------------------------
'ROI_Supramarginal_Inferior_BA40_6': [
'S6_D6'
],
'ROI_VentralSomatosensory_BA1_2_3_43_48': [
'S8_D6', 'S9_D6'
],
'ROI_Sensorimotor_BA1_2_3_4_6': [
'S6_D8', 'S10_D8'
],
'ROI_DLPFC_FEF_BA6_8_9': [
'S7_D7', 'S7_D8'
],
'ROI_Broca_VLPFC_BA6_44_45_4': [
'S9_D7', 'S10_D7'
]
}
DOWNSAMPLE: bool
DOWNSAMPLE_FREQUENCY: int
@@ -303,6 +339,7 @@ GROUP: str = "Default"
FOLDING_BYP: bool = False
FEATURE_1: bool = False
FEATURE_2: bool = False
# Ensure that we are working in the directory of this file
script_dir = os.path.dirname(os.path.abspath(__file__))
@@ -610,10 +647,19 @@ def process_multiple_participants(file_paths, file_params, file_metadata,
process_multiple_participants._success_count = success_count["value"]
process_multiple_participants._failed_stages = failed_stages["value"]
successes = [r for r in qc_rows if r.get("status") == "success"]
population_flags = flag_population_outliers(successes)
for path, flags in population_flags.items():
metric_summary = ", ".join(f"{f['metric']}={f['value']} (median={f['median']:.1f}, z={f['z']})" for f in flags)
n_metrics_flagged = len(flags)
severity = "STRONG" if n_metrics_flagged >= 3 else "possible"
logger.warning(f"{severity} outlier: {path} - flagged on {n_metrics_flagged} metric(s): {metric_summary}")
if qc_summary_path and qc_rows:
if FEATURE_1:
try:
write_qc_excel_summary(qc_rows, qc_summary_path)
write_qc_excel_summary(qc_rows, qc_summary_path, population_flags=population_flags)
logger.info(f"QC summary written to {qc_summary_path} ({len(qc_rows)} participant(s))")
except Exception as e:
logger.error(f"Failed to write QC summary: {e}")
@@ -622,6 +668,86 @@ def process_multiple_participants(file_paths, file_params, file_metadata,
def flag_population_outliers(
qc_rows: list[dict],
metrics: list[str] | None = None,
z_threshold: float = 2.5,
) -> dict[str, list[dict]]:
"""
Given QC metrics collected across an entire batch (each participant
processed with identical, fixed thresholds - nothing adaptive changes
per-participant), flags participants whose value on a given metric is
a statistical outlier relative to the rest of the batch.
This does NOT change what gets marked as a bad channel, and does NOT
imply anything is wrong with the metric thresholds themselves - it
answers a different question: "given everyone went through the same
pipeline with the same settings, does this participant's result look
unusual compared to everyone else who did." A real, uniformly bad
dataset (e.g. wrong population for the config) would show up as
outliers across MULTIPLE metrics; a participant with one borderline
metric and nothing else is much weaker evidence of a real problem.
Uses a modified z-score (median + MAD-based, not mean + std) since
QC metric distributions across a real population are often skewed by
a few genuinely bad files - MAD-based scoring is robust to those
outliers dominating the very estimate used to detect them, unlike a
standard mean/std z-score.
Parameters
----------
qc_rows : list of per-participant QC dicts (successful runs only -
filter out FAILED entries before calling this).
metrics : list of QC dict keys to check. Defaults to the standard
bad-channel-count metrics if not specified.
z_threshold : float, default 2.5
Modified z-score magnitude above which a participant is flagged
for that metric. 2.5 is a commonly used, moderately conservative
starting point for outlier flagging - not validated against your
specific data, worth adjusting based on what you see in practice.
Returns
-------
dict[str, list[dict]]
Keyed by file_path, listing which metric(s) that participant was
flagged as an outlier on and by how much, e.g.
{"sub-07.snirf": [{"metric": "n_bad_sci", "z": 3.1, "value": 22, "median": 3}]}
"""
if metrics is None:
metrics = [
"n_bad_sci", "n_bad_snr", "n_bad_psp", "n_bad_coeff_var",
"n_bad_mad", "n_bad_psd_noise", "n_bad_dropout",
"n_bad_channels_total", "pct_bad_channels",
]
flagged: dict[str, list[dict]] = {}
for metric in metrics:
values = np.array([row.get(metric) for row in qc_rows if row.get(metric) is not None], dtype=float)
if len(values) < 4:
logger.info(f"Skipping outlier check for '{metric}' - too few participants ({len(values)}) for a meaningful comparison.")
continue
median = np.median(values)
mad = np.median(np.abs(values - median))
if mad == 0:
continue # every participant identical on this metric - nothing to flag
for row in qc_rows:
val = row.get(metric)
if val is None:
continue
modified_z = 0.6745 * (val - median) / mad
if abs(modified_z) >= z_threshold:
path = row.get("file_path", "unknown")
flagged.setdefault(path, []).append({
"metric": metric, "z": round(float(modified_z), 2),
"value": val, "median": float(median),
})
return flagged
def markbad(data, ax, ch_names: list[str]) -> None:
"""
Add a strikethrough to a plot for channels marked as bad.
@@ -1239,16 +1365,13 @@ def calculate_peak_power(data: BaseRaw, time_window: int = 3, threshold: float =
psp = scores.mean(axis=1)
bad_channels = list(compress(cast(list[str], data.ch_names), psp < threshold))
plot_data = data.copy()
plot_data.info["bads"] = bad_channels
existing_bads = set(data.info.get("bads", []))
data.info["bads"] = list(existing_bads | set(bad_channels))
# Determine the colors based on the threshold, and create the figures
color_stops = ([0.0, threshold, threshold+0.1, threshold+0.2, 1.0], [0.0, threshold, threshold, 1.0])
psp1, psp2 = plot_timechannel_quality_metrics(plot_data, scores, times, color_stops, threshold, "Peak Spectral Power")
print(f"thresh: {threshold}")
psp1, psp2 = plot_timechannel_quality_metrics(data, scores, times, color_stops, threshold, "Peak Spectral Power")
return list(compress(cast(list[str], getattr(data, "ch_names")), psp < threshold)), psp1, psp2
@@ -5128,6 +5251,89 @@ def find_bad_channels_coeff_var(raw, coeff_var_threshold=25.0):
def detect_hbo_hbr_anticorrelation(raw_haemo, threshold: float = -0.2):
"""
Flags channels where HbO and HbR are NOT showing the expected
physiologically anti-correlated relationship. Real hemodynamic response
typically shows HbO rising while HbR falls (and vice versa) - a channel
pair with weak or positive correlation is a common signature of motion
artifact or poor optode coupling that other QC metrics can miss, since
it's checking signal SHAPE/relationship rather than amplitude, variance,
or noise level.
IMPORTANT: must be run on raw_haemo BEFORE enhance_negative_correlation
(or any similar correction step) - that step actively forces HbO/HbR
anti-correlation, so measuring this diagnostic afterward would just be
checking whether the correction worked, not the underlying data quality.
Parameters
----------
raw_haemo : BaseRaw
Haemoglobin-concentration data (post Beer-Lambert, pre-correction).
threshold : float, default -0.2
Channels with HbO/HbR correlation ABOVE this value are flagged as
bad (i.e. not sufficiently anti-correlated). -0.2 is a permissive
starting point - real channels often land well below this (-0.5 to
-0.9), but very low SNR or task designs can naturally weaken the
correlation without indicating artifact, so this shouldn't be set
aggressively without checking against your own known-good data.
Returns
-------
tuple[list[str], Figure]
- list[str]: channel names (both hbo AND hbr for each flagged pair)
below the anti-correlation threshold.
- Figure: bar chart of correlation per channel pair, matching the
visual style of detect_sensor_dropout.
"""
ch_names = raw_haemo.ch_names
data = raw_haemo.get_data()
base_names = sorted({ch.split()[0] for ch in ch_names})
correlations = {}
bad_bases = []
for base in base_names:
try:
hbo_idx = ch_names.index(f"{base} hbo")
hbr_idx = ch_names.index(f"{base} hbr")
except ValueError:
continue # channel doesn't have both chromophores present
hbo_signal = data[hbo_idx]
hbr_signal = data[hbr_idx]
if np.std(hbo_signal) == 0 or np.std(hbr_signal) == 0:
corr = 0.0 # flat channel - can't meaningfully correlate
else:
corr = float(np.corrcoef(hbo_signal, hbr_signal)[0, 1])
correlations[base] = corr
if corr > threshold:
bad_bases.append(base)
print(f"Flagged {base}: HbO/HbR correlation = {corr:.3f} (expected below {threshold})")
bad_names = [ch for ch in ch_names if ch.split()[0] in bad_bases]
fig, ax = plt.subplots(figsize=(10, 5), constrained_layout=True)
bases_sorted = list(correlations.keys())
corr_values = [correlations[b] for b in bases_sorted]
colors = ['coral' if c > threshold else 'skyblue' for c in corr_values]
ax.bar(range(len(corr_values)), corr_values, color=colors)
ax.axhline(threshold, color='red', linestyle='--', label=f'Threshold ({threshold})')
ax.axhline(0, color='black', linewidth=0.8)
ax.set_title("HbO/HbR Anti-Correlation Check")
ax.set_ylabel("Pearson r (HbO vs HbR)")
ax.set_xlabel("Channel Pair Index")
ax.legend()
plt.close(fig)
print(f"Anti-correlation check: flagged {len(bad_bases)} optode pair(s).")
return bad_names, fig
def hr_calc(raw, seconds_to_strip_hr, l_freq, h_freq, search_min, search_max, max_low_hr, max_high_hr, smoothing_window_hr, hr_window, short_channels, short_channels_threshold, verbosity, psd_confidence_threshold: float = 3.0, band_halfwidth_hz: float = 0.3):
if short_channels:
short_chans = get_short_channels(raw, max_dist=short_channels_threshold)
@@ -5745,6 +5951,12 @@ def process_participant(file_path, file_start, progress_callback=None):
logger.info("Step 17 Completed.")
step_start = lap(step_start, timings, "Step 17")
bad_anticorr = []
if FEATURE_2 and not FOLDING_BYP:
bad_anticorr, fig_anticorr = detect_hbo_hbr_anticorrelation(raw_haemo, threshold=-0.2)
_enqueue("HbO-HbR Anti-Correlation", fig_anticorr, png_queue)
qc["n_bad_anticorrelation"] = len(bad_anticorr)
# Step 18: Enhance Negative Correlation
if ENHANCE_NEGATIVE_CORRELATION and not FOLDING_BYP:
raw_haemo = enhance_negative_correlation(raw_haemo)
@@ -5877,6 +6089,27 @@ def process_participant(file_path, file_start, progress_callback=None):
if progress_callback: progress_callback(27)
logger.info("27")
step_start = lap(step_start, timings, "Step 27")
# Step 27.5: Extract FIR Waveform Features & Enqueue Metric Plots
fir_feature_dict = {'features': np.array([]), 'feature_names': [], 'feature_channels': []}
# if HRF_MODEL.lower() == "fir":
# try:
# fir_feature_dict = extract_fir_features_real_data(
# raw=raw_haemo,
# target_condition=None, # e.g., 'reach'
# fir_delays=FIR_DELAYS, # e.g., np.arange(0, 15)
# selected_metrics=tuple(METRIC_REGISTRY.keys()), # e.g., ('Peak_Amp', 'TTP', 'AUC')
# roi_map=ROI_MAP,
# glm_est=glm_est,
# df_design_matrix=df_design_matrix,
# png_queue=png_queue
# )
# logger.info("Step 27.5: FIR features successfully extracted and metric images enqueued.")
# except Exception as e:
# logger.warning(f"Step 27.5 Failed to extract FIR features: {e}")
# Step 28: Finishing Up
png_queue.put(None) # sentinel
@@ -5895,7 +6128,7 @@ def process_participant(file_path, file_start, progress_callback=None):
logger.info(f" {name:<25} {elapsed:7.3f}s")
logger.info(f"Total processing time: {sum(timings.values()):.3f}s")
return raw_haemo, epochs, df_cha, df_roi, df_design_matrix, config_dict, fig_bytes_dict, contrast_results_dict, roi_channel_map, qc, True
return raw_haemo, epochs, df_cha, df_roi, df_design_matrix, config_dict, fig_bytes_dict, contrast_results_dict, roi_channel_map, fir_feature_dict, qc, True
@@ -6769,7 +7002,11 @@ def peak_power_fast(
return raw, scores, times
def write_qc_excel_summary(qc_rows: list[dict], output_path: str) -> None:
def write_qc_excel_summary(
qc_rows: list[dict],
output_path: str,
population_flags: dict[str, list[dict]] | None = None,
) -> None:
"""
Writes one Excel workbook summarizing QC metrics across all participants
in a batch run - participants as columns, metrics as rows, each row
@@ -6838,18 +7075,23 @@ def write_qc_excel_summary(qc_rows: list[dict], output_path: str) -> None:
header_range = f"{first_col_letter}1:{last_col_letter}1"
if key in QC_METRIC_DEVIATION_BASED:
# "worst" = furthest from the row's own median
worst_formula = (
f"=INDEX({header_range},MATCH(MAX(ABS({data_range}-MEDIAN({data_range}))),"
f"ABS({data_range}-MEDIAN({data_range})),0))"
)
elif QC_METRIC_DIRECTIONS[key]: # lower is better -> worst = max
worst_formula = f"=INDEX({header_range},MATCH(MAX({data_range}),{data_range},0))"
else: # higher is better -> worst = min
worst_formula = f"=INDEX({header_range},MATCH(MIN({data_range}),{data_range},0))"
ws.cell(row=row_num, column=worst_col, value=worst_formula).font = body_font
row_values = [row.get(key) for row in successes if row.get(key) is not None]
row_paths = [row.get("file_path") for row in successes if row.get(key) is not None]
if row_values:
med = float(np.median(row_values))
deviations = [abs(v - med) for v in row_values]
worst_idx = int(np.argmax(deviations))
worst_value = row_paths[worst_idx]
else:
worst_value = ""
ws.cell(row=row_num, column=worst_col, value=worst_value).font = body_font
else:
# existing live-formula path for direction-based metrics, unchanged
if QC_METRIC_DIRECTIONS[key]:
worst_formula = f"=INDEX({header_range},MATCH(MAX({data_range}),{data_range},0))"
else:
worst_formula = f"=INDEX({header_range},MATCH(MIN({data_range}),{data_range},0))"
ws.cell(row=row_num, column=worst_col, value=worst_formula).font = body_font
# --- Color scale, direction-aware ---
if key in QC_METRIC_DEVIATION_BASED:
# color by |value - row median| via a helper column pattern isn't
@@ -6876,9 +7118,32 @@ def write_qc_excel_summary(qc_rows: list[dict], output_path: str) -> None:
)
ws.conditional_formatting.add(data_range, rule)
if population_flags:
outlier_row = len(metric_keys) + 2
ws.cell(row=outlier_row, column=1, value="Population Outlier Flags").font = label_font
any_flags = bool(population_flags)
for i, row in enumerate(successes):
col = i + 2
path = row.get("file_path", "")
flags = population_flags.get(path, []) if population_flags else []
if flags:
text = "; ".join(f"{f['metric']} (z={f['z']})" for f in flags)
c = ws.cell(row=outlier_row, column=col, value=text)
c.fill = PatternFill(
start_color="FFEB84" if len(flags) < 3 else "F8696B",
end_color="FFEB84" if len(flags) < 3 else "F8696B",
fill_type="solid",
)
else:
c = ws.cell(row=outlier_row, column=col, value="No outliers detected")
c.fill = PatternFill(start_color="C6EFCE", end_color="C6EFCE", fill_type="solid")
c.font = body_font
c.alignment = Alignment(horizontal="center")
# --- Failed participants block, separate and clearly marked ---
if failures:
fail_row_start = len(metric_keys) + 4
fail_row_start = len(metric_keys) + 5
ws.cell(row=fail_row_start, column=1, value="FAILED PARTICIPANTS").font = Font(name="Arial", bold=True, size=12, color="CC0000")
for i, row in enumerate(failures):
r = fail_row_start + 1 + i
@@ -6888,6 +7153,11 @@ def write_qc_excel_summary(qc_rows: list[dict], output_path: str) -> None:
err_cell = ws.cell(row=r, column=2, value=row.get("error", "unknown error"))
err_cell.font = body_font
err_cell.fill = fail_fill
else:
fail_row_start = len(metric_keys) + 5
c = ws.cell(row=fail_row_start, column=1, value="All participants processed successfully - no failures.")
c.font = Font(name="Arial", bold=True, size=11, color="006100")
c.fill = PatternFill(start_color="C6EFCE", end_color="C6EFCE", fill_type="solid")
ws.column_dimensions['A'].width = 32
for i in range(n_participants):
@@ -6897,6 +7167,240 @@ def write_qc_excel_summary(qc_rows: list[dict], output_path: str) -> None:
wb.save(output_path)
METRIC_REGISTRY = {
'Peak_Amp': 'Peak_Amp',
'TTP': 'Time_to_Peak',
'AUC': 'AUC',
'Rising_Slope': 'Rising_Slope',
'Recovery_Slope': 'Recovery_Slope',
'FWHM': 'FWHM',
'Onset_Latency': 'Onset_Latency',
'P2P_Amp': 'Peak_to_Peak_Amp',
'Signal_Std': 'Signal_Std',
'RMS': 'RMS'
}
def plot_and_enqueue_waveform_metrics(
roi_curves,
fir_delays,
selected_metrics,
target_condition='reach',
png_queue=None
):
"""
Generates and enqueues visual plots for every calculated waveform metric
across all Regions of Interest (ROIs).
"""
# Structure metric values per ROI
metric_data = {METRIC_REGISTRY[m]: {} for m in selected_metrics}
for (chromo, roi_name), roi_fir_curve in roi_curves.items():
metrics = compute_waveform_metrics(roi_fir_curve, fir_delays=fir_delays, selected_metrics=selected_metrics)
for m_key, val in zip(selected_metrics, metrics):
label = METRIC_REGISTRY[m_key]
metric_data[label][f"{roi_name} ({chromo.upper()})"] = val
# Generate an image for each metric across ROIs
for metric_label, roi_dict in metric_data.items():
if not roi_dict:
continue
fig, ax = plt.subplots(figsize=(8, 4.5))
rois = list(roi_dict.keys())
values = list(roi_dict.values())
colors = ['#2b5c8f' if v >= 0 else '#d9534f' for v in values]
ax.bar(rois, values, color=colors, alpha=0.85, edgecolor='black')
ax.axhline(0, color='black', linewidth=0.8, linestyle='--')
ax.set_title(f"FIR Waveform Metric: {metric_label} [{target_condition}]", fontsize=12, fontweight='bold')
ax.set_xlabel("Region of Interest (ROI)", fontsize=10)
ax.set_ylabel(metric_label, fontsize=10)
plt.xticks(rotation=35, ha='right')
plt.grid(axis='y', linestyle=':', alpha=0.6)
plt.tight_layout()
# Enqueue figure or close
if png_queue is not None:
_enqueue(f"FIR Waveform Metric - {metric_label}", fig, png_queue)
else:
plt.close(fig)
def _compute_roi_fir_curves(
raw=None,
target_condition='reach',
fir_delays=np.arange(0, 15),
roi_map=ROI_MAP,
chromophores=('hbr',),
glm_est=None,
df_design_matrix=None
):
"""
Shared FIR-GLM curve extraction. If `glm_est` and `df_design_matrix` are supplied,
it reuses pre-calculated GLM results directly to avoid duplicate processing.
"""
roi_curves = {}
# --- SHORT-CIRCUIT: Reuse pre-calculated GLM estimation if available ---
if glm_est is not None and df_design_matrix is not None:
if hasattr(glm_est, 'to_dataframe'):
glm_df = glm_est.to_dataframe().reset_index()
elif isinstance(glm_est, pd.DataFrame):
glm_df = glm_est.copy()
else:
raise ValueError("Unsupported format for precalculated glm_est.")
glm_df.columns = [str(col).lower() for col in glm_df.columns]
cond_col = 'condition' if 'condition' in glm_df.columns else 'regressor'
ch_col = 'ch_name' if 'ch_name' in glm_df.columns else ('source' if 'source' in glm_df.columns else 'channel')
print("Available conditions in GLM:", glm_df[cond_col].unique())
fir_df = glm_df[glm_df[cond_col].astype(str).str.lower().str.contains(target_condition.lower())].copy()
print(f"Matched rows for '{target_condition}': {len(fir_df)}")
if fir_df.empty:
logger.warning(f"Condition '{target_condition}' not found in precalculated GLM estimates.")
return roi_curves
for chromo in chromophores:
chromo_df = fir_df[fir_df[ch_col].str.lower().str.contains(chromo.lower())] if ch_col in fir_df.columns else fir_df
ch_curves = {}
for ch_name, ch_group in chromo_df.groupby(ch_col):
pair = ch_name.split(' ')[0]
ch_curves[pair] = ch_group['theta'].values if 'theta' in ch_group.columns else ch_group['beta'].values
print("Extracted channel keys:", list(ch_curves.keys())[:5])
for roi_name, channels in roi_map.items():
matching_curves = [ch_curves[ch] for ch in channels if ch in ch_curves]
if matching_curves:
roi_curves[(chromo, roi_name)] = np.mean(matching_curves, axis=0)
return roi_curves
def extract_fir_features_real_data(
raw=None,
target_condition=None,
fir_delays=np.arange(0, 15),
selected_metrics=('Peak_Amp',),
roi_map=ROI_MAP,
glm_est=None,
df_design_matrix=None,
png_queue=None
):
"""
Extracts FIR scalar waveform metrics, plots and enqueues figures for each
waveform metric, and returns all outputs collapsed into a single dictionary variable.
"""
raw_cols = [
col for col in df_design_matrix.columns
if not any(k in col.lower() for k in ['drift', 'constant', 'short', 'nuisance'])
]
# Strip '_delay_0', '_delay_1', etc. to get base condition names
target_conditions = list(dict.fromkeys(
col.split('_delay_')[0] if '_delay_' in col else col
for col in raw_cols
))
collapsed_features = []
feature_names = []
feature_channels = []
# 2. Iterate over EVERY condition
for cond in target_conditions:
try:
roi_curves = _compute_roi_fir_curves(
raw=raw,
target_condition=cond,
fir_delays=fir_delays,
roi_map=roi_map,
chromophores=('hbr',),
glm_est=glm_est,
df_design_matrix=df_design_matrix
)
if not roi_curves:
print("999999999999999999999999999999999999999999999999999999999999999999999999999999999999999999999999")
continue
metric_labels = [METRIC_REGISTRY[m] for m in selected_metrics]
# Extract metrics per ROI for this condition
for (chromo, roi_name), roi_fir_curve in roi_curves.items():
metrics = compute_waveform_metrics(roi_fir_curve, fir_delays=fir_delays, selected_metrics=selected_metrics)
collapsed_features.extend(metrics)
# Prefix feature names with condition
feature_names.extend([f"{cond}_{roi_name}_{m}" for m in metric_labels])
feature_channels.extend([roi_name] * len(metric_labels))
# Enqueue plots for this specific condition
plot_and_enqueue_waveform_metrics(
roi_curves=roi_curves,
fir_delays=fir_delays,
selected_metrics=selected_metrics,
target_condition=cond,
png_queue=png_queue
)
except Exception as e:
logger.warning(f"Failed extracting FIR metrics for condition '{cond}': {e}")
return {
'features': np.array(collapsed_features),
'feature_names': feature_names,
'feature_channels': feature_channels
}
def compute_waveform_metrics(fir_curve, fir_delays, selected_metrics=('Peak_Amp',)):
peak_idx = np.argmax(fir_curve)
peak_amp = fir_curve[peak_idx]
ttp = fir_delays[peak_idx]
calculated_metrics = {}
if 'Peak_Amp' in selected_metrics:
calculated_metrics['Peak_Amp'] = peak_amp
if 'TTP' in selected_metrics:
calculated_metrics['TTP'] = ttp
if 'AUC' in selected_metrics:
trapz_fn = getattr(np, 'trapezoid', getattr(np, 'trapz', None))
calculated_metrics['AUC'] = trapz_fn(fir_curve, fir_delays)
if 'Rising_Slope' in selected_metrics:
calculated_metrics['Rising_Slope'] = (peak_amp - fir_curve[0]) / (ttp - fir_delays[0]) if ttp > fir_delays[0] else 0.0
if 'Recovery_Slope' in selected_metrics:
calculated_metrics['Recovery_Slope'] = (fir_curve[-1] - peak_amp) / (fir_delays[-1] - ttp) if fir_delays[-1] > ttp else 0.0
if 'FWHM' in selected_metrics:
half_max = peak_amp / 2.0
above_half = np.where(fir_curve >= half_max)[0]
calculated_metrics['FWHM'] = fir_delays[above_half[-1]] - fir_delays[above_half[0]] if len(above_half) > 1 else 0.0
if 'Onset_Latency' in selected_metrics:
onset_thresh = 0.2 * peak_amp
above_onset = np.where(fir_curve >= onset_thresh)[0]
calculated_metrics['Onset_Latency'] = fir_delays[above_onset[0]] if len(above_onset) > 0 else 0.0
if 'P2P_Amp' in selected_metrics:
calculated_metrics['P2P_Amp'] = peak_amp - np.min(fir_curve)
if 'Signal_Std' in selected_metrics:
calculated_metrics['Signal_Std'] = np.std(fir_curve)
if 'RMS' in selected_metrics:
calculated_metrics['RMS'] = np.sqrt(np.mean(fir_curve**2))
return [calculated_metrics[m] for m in selected_metrics]
if __name__ == "__main__":
print("This file has no functionality when not used in tandem with the FLARES application.")