Spaces:
Sleeping
Sleeping
| 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'.") | |