CoolFace
Apppublic

Wojszm/eeg

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
app.py176 linesDownload Raw Back to root
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)