Wojszm/eeg
0
1import os2import torch3import torch.nn as nn4import numpy as np5import scipy.io as sio6import gradio as gr7 8class Discriminator1D(nn.Module):9 def __init__(self, channels_eeg, features_d, signal_length):10 super(Discriminator1D, self).__init__()11 self.disc = nn.Sequential(12 nn.Conv1d(channels_eeg, features_d, kernel_size=4, stride=2, padding=1),13 nn.LeakyReLU(0.2, inplace=True),14 self._block(features_d, features_d * 2, kernel_size=4, stride=2, padding=1),15 self._block(features_d * 2, features_d * 4, kernel_size=4, stride=2, padding=1),16 self._block(features_d * 4, features_d * 8, kernel_size=4, stride=2, padding=1),17 self._block(features_d * 8, features_d * 16, kernel_size=4, stride=2, padding=1),18 nn.Conv1d(features_d * 16, 1, kernel_size=8, stride=1, padding=0),19 )20 21 def _block(self, in_channels, out_channels, kernel_size, stride, padding):22 return nn.Sequential(23 nn.Conv1d(in_channels, out_channels, kernel_size, stride, padding, bias=False),24 nn.InstanceNorm1d(out_channels, affine=True),25 nn.LeakyReLU(0.2, inplace=True),26 )27 28 def forward(self, x):29 return self.disc(x)30 31 32def classify_signal_multi(segment, disc_dict, device):33 34 if isinstance(segment, np.ndarray):35 segment = torch.tensor(segment, dtype=torch.float32)36 37 if segment.dim() == 1:38 segment = segment.unsqueeze(0) # (1, 256)39 segment = segment.unsqueeze(0).to(device) # (1, 1, 256)40 41 scores = {}42 with torch.no_grad():43 for freq_label, disc_model in disc_dict.items():44 score = disc_model(segment).view(-1).item()45 scores[freq_label] = score46 47 best_freq = max(scores, key=scores.get)48 return best_freq, scores49 50 51 52device = torch.device("cuda" if torch.cuda.is_available() else "cpu")53channels_eeg = 154features_d = 6455signal_length = 25656 57 58model_paths = {59 '5Hz': "discriminator_5Hz.pth",60 '6Hz': "discriminator_6Hz.pth",61 '7Hz': "discriminator_7Hz.pth",62 '8Hz': "discriminator_8Hz.pth"63}64 65 66def load_discriminator(model_path):67 model = Discriminator1D(channels_eeg, features_d, signal_length).to(device)68 model.load_state_dict(torch.load(model_path, map_location=device))69 model.eval()70 return model71 72 73disc_dict = {freq: load_discriminator(path) for freq, path in model_paths.items()}74 75 76def process_mat_file(mat_file):77 78 output_str = ""79 80 81 try:82 mat_contents = sio.loadmat(mat_file.name)83 except Exception as e:84 return f"Nie udało się wczytać pliku: {e}"85 86 87 if 'X' not in mat_contents:88 return "Plik nie zawiera klucza 'X'."89 90 eeg_data = mat_contents['X']91 92 93 eeg_data = eeg_data.T94 95 96 if eeg_data.shape[1] <= 14:97 return "Plik ma zbyt mało kanałów. Nie można wybrać kanału 14."98 99 100 first_channel_data = eeg_data[14, :]101 output_str += f"<b>Długość sygnału (kanał 15):</b> {len(first_channel_data)}\n<br>"102 103 104 sampling_rate = 256105 segment_length = sampling_rate 106 hop_size = 10 107 108 109 segments_list = []110 for start_idx in range(0, len(first_channel_data) - segment_length + 1, hop_size):111 end_idx = start_idx + segment_length112 segment = first_channel_data[start_idx:end_idx]113 segments_list.append(segment)114 115 if not segments_list:116 return "Brak wyciętych segmentów."117 118 119 all_data = np.array(segments_list)120 121 122 min_val = np.min(all_data, axis=1, keepdims=True)123 max_val = np.max(all_data, axis=1, keepdims=True)124 all_data = 2.0 * (all_data - min_val) / (max_val - min_val + 1e-8) - 1.0125 126 127 segment_classifications = []128 for seg in all_data:129 best_freq, scores_dict = classify_signal_multi(seg, disc_dict, device=device)130 segment_classifications.append(best_freq)131 132 133 freq_counts = {134 '5Hz': segment_classifications.count('5Hz'),135 '6Hz': segment_classifications.count('6Hz'),136 '7Hz': segment_classifications.count('7Hz'),137 '8Hz': segment_classifications.count('8Hz'),138 }139 total_segments = len(all_data)140 dominant_freq = max(freq_counts, key=freq_counts.get)141 142 143 output_str += f"<b>Liczba segmentów:</b> {total_segments}<br>"144 output_str += f"<b>Wynik klasyfikacji:</b> {freq_counts}<br>"145 output_str += f"<b style='color: green;'>Dominująca częstotliwość:</b> <b style='color: green;'>{dominant_freq}</b><br>"146 147 return output_str148 149 150 151iface = gr.Interface(152 fn=process_mat_file,153 inputs=gr.File(label="Wgraj plik .mat", file_types=[".mat"]),154 outputs=gr.HTML(elem_id="output_html"),155 title="Klasyfikator sygnału EEG 5-8Hz - praca magisterska",156 description="""157 \n Aplikacja dokonuje segmentacji sygnału z kanału nr 15 (elektroda OZ) [index 14], z przesunięciem okna co 10 próbek, oraz klasyfikuje poszczególne segmenty przy użyciu czterech wytrenowanych dyskryminatorów (5 Hz, 6 Hz, 7 Hz, 8 Hz).158 \n Do trenowania modelu wykorzystano architekturę WGAN-GP.159 \n Model został wytrenowany na zbiorze danych pochodzącym z następującego artykułu: https://www.researchgate.net/publication/290096626_Dataset_BCI_EEG_SSVEP_for_four_classes_of_stimuli160 \n Struktura wgrywanej macierzy powinna być następująca: próbki × kanały, gdzie wiersze odpowiadają próbkom, a kolumny reprezentują kanały EEG.161 \n Przypisz wgrywanej macierzy klucz "X".162 \n Wgraj plik .mat zawierający dane EEG. 163 """,164 css="""165 /* Zwiększenie wysokości okienka wyjściowego */166 #output_html {167 height: 500px;168 overflow: auto;169 }170 """171 172 173)174 175if __name__ == "__main__":176 iface.launch(share=True)