Spaces:
Sleeping
Sleeping
File size: 6,148 Bytes
546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d dd5efa8 546b54d ea26380 546b54d ea26380 546b54d dd5efa8 546b54d dd5efa8 ea26380 546b54d dd5efa8 546b54d ea26380 546b54d ea26380 546b54d ea26380 dd5efa8 ea26380 546b54d ea26380 546b54d dd5efa8 ea26380 dd5efa8 546b54d ea26380 546b54d dd5efa8 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d ea26380 546b54d | 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 | import numpy as np
import torch
import tensorflow as tf
import pandas as pd
import joblib
import mne
from sklearn.metrics import confusion_matrix, f1_score, accuracy_score
from preprocessing import preprocess_eeg_file
from preprocessing_2dcnn import convert_epoch_to_spectrogram
from preprocessing_epilepsynet import *
from EpilepsyNet_model import TimeSeriesAttentionClassifier
from eegnet_model import EEGNet
def aggregate_predictions(spectrogram_list, model, threshold=0.5):
X = np.array([np.transpose(s, (1, 2, 0)) for s in spectrogram_list])
print(f'---Aggregating predictions from {len(spectrogram_list)} segments---')
preds = model.predict(X)
mean_prob = np.mean(preds[:, 1])
segment_probs = preds[:, 1].tolist()
final_label = 1 if mean_prob >= threshold else 0
return final_label, mean_prob, segment_probs
def predict_eeg_recording(edf_path, model_name='2DCNN', threshold=0.5):
if model_name == '2DCNN':
model = tf.keras.models.load_model('model1_2dcnn.h5')
channels = ["EEG FP1-REF", "EEG FP2-REF", "EEG F3-REF", "EEG F4-REF", "EEG C3-REF"]
preprocessed_df = preprocess_eeg_file(edf_path, fmin=1.0, fmax=45.0, segment_lenght=5, overlap=2, desired=channels)
if preprocessed_df is None or preprocessed_df.empty:
raise ValueError("EEG file could not be preprocessed or no valid segments found.")
spectrogram_list = preprocessed_df.apply(
lambda row: convert_epoch_to_spectrogram(row, channels, fs=250, nperseg=128, noverlap=64), axis=1
).tolist()
return aggregate_predictions(spectrogram_list, model, threshold)
elif model_name == 'EEGNet':
loaded = joblib.load("eegnet_model.joblib")
state_dict = loaded["model_state_dict"]
model = EEGNet(n_channels=21, n_samples=1250, num_classes=2)
model.load_state_dict(state_dict)
channels = [
'EEG FP1-REF', 'EEG FP2-REF', 'EEG F3-REF', 'EEG F4-REF', 'EEG C3-REF', 'EEG C4-REF',
'EEG P3-REF', 'EEG P4-REF', 'EEG O1-REF', 'EEG O2-REF', 'EEG F7-REF', 'EEG F8-REF',
'EEG T3-REF', 'EEG T4-REF', 'EEG T5-REF', 'EEG T6-REF', 'EEG FZ-REF', 'EEG CZ-REF',
'EEG PZ-REF', 'EEG ROC-REF', 'EEG LOC-REF'
]
preprocessed_df = preprocess_eeg_file(
edf_path, fmin=1.0, fmax=45.0, segment_lenght=5, overlap=0, desired=channels
)
if preprocessed_df is None or preprocessed_df.empty:
raise ValueError("EEG file could not be preprocessed or no valid segments found.")
timeseries_list = preprocessed_df.apply(
lambda row: convert_epoch_to_timeseries(row, channels), axis=1
).tolist()
return aggregate_predictions_EEGNET(timeseries_list, model, threshold)
elif model_name == 'EpilepsyNet':
raw = mne.io.read_raw_edf(edf_path, preload=True, verbose='ERROR')
eeg_cols = ['EEG FP1', 'EEG FP2', 'EEG F3', 'EEG F4', 'EEG C3', 'EEG C4', 'EEG P3', 'EEG P4',
'EEG O1', 'EEG O2', 'EEG F7', 'EEG F8', 'EEG T3', 'EEG T4', 'EEG T5', 'EEG T6',
'EEG T1', 'EEG T2', 'EEG FZ', 'EEG CZ', 'EEG PZ']
parameters = {
'eeg_cols': eeg_cols,
'segment_duration': 60.0,
'n_segments_per_file': 12,
'samples_per_segment': 1250,
'random_state': 42
}
X = process_raw_files(
raw_file=raw,
eeg_cols=eeg_cols,
segment_duration=parameters['segment_duration'],
n_segments_per_file=parameters['n_segments_per_file'],
random_state=parameters['random_state']
)
X_std = standardize_data(X)
corr_matrix = compute_correlation_matrix(X_std)
upper_triangle_matrix = extract_upper_triangle(corr_matrix)
X_tensor = torch.tensor(upper_triangle_matrix, dtype=torch.float32)
X_tensor = X_tensor.unsqueeze(0)
input_dim = 210
embed_dim = 256
num_heads = 16
model = TimeSeriesAttentionClassifier(input_dim, embed_dim, num_heads)
model.load_state_dict(torch.load('EpilepsyNet.pth'))
model.eval()
outputs, _ = model(X_tensor)
predicted = (outputs >= 0.5).float()
prob = outputs.float().squeeze().item()
return int(predicted), prob, [prob]
def convert_epoch_to_timeseries(epoch_row, channels):
ts_list = []
for ch in channels:
if ch in epoch_row:
ts = epoch_row[ch]
ts_list.append(ts)
return np.stack(ts_list, axis=0)
def aggregate_predictions_EEGNET(segment_list, model, threshold):
model.eval()
preds = []
with torch.no_grad():
for seg in segment_list:
seg_tensor = torch.tensor(seg, dtype=torch.float32)
seg_tensor = seg_tensor.unsqueeze(0).unsqueeze(0)
output = model(seg_tensor)
prob = torch.softmax(output, dim=1)[0].cpu().numpy()
preds.append(prob)
avg_pred = np.mean(preds, axis=0)
segment_probs = [float(p[1]) for p in preds]
final_class = int(avg_pred[1] >= threshold)
return final_class, float(avg_pred[1]), segment_probs
def predict_ensemble_eeg_recording(edf_path, ensemble_method, threshold=0.5):
pred_prob_list = []
votes = []
for model_name in ["2DCNN", "EEGNet", "EpilepsyNet"]:
pred_label, prob, _ = predict_eeg_recording(edf_path, model_name=model_name, threshold=threshold)
pred_prob_list.append(prob)
votes.append(int(prob >= threshold))
print(f"Prediction from {model_name}: label={pred_label}, probability={prob}")
if ensemble_method.lower() == "average":
avg_prob = np.mean(pred_prob_list)
final_class = int(avg_prob >= threshold)
print("Averaged probability:", avg_prob)
return final_class, avg_prob, []
elif ensemble_method.lower() == "voting":
final_class = int(round(np.mean(votes)))
print("Votes from each model:", votes)
return final_class, votes, []
else:
raise ValueError("Ensemble method must be either 'average' or 'voting'.")
|