{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":4043,"databundleVersionId":44567,"sourceType":"competition"},{"sourceId":10291929,"sourceType":"datasetVersion","datasetId":6369491}],"dockerImageVersionId":30822,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# unzip train and test\n\nimport zipfile\nimport os\n\nzip_file_path = '/kaggle/input/inria-bci-challenge/train.zip'\nextract_to = '/kaggle/working/train'\n\nos.makedirs(extract_to, exist_ok=True)\n\nwith zipfile.ZipFile(zip_file_path, 'r') as zip_ref:\n    zip_ref.extractall(extract_to)\n\nprint(f\"File extracted to {extract_to}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:22:59.382996Z","iopub.execute_input":"2024-12-26T08:22:59.383427Z","iopub.status.idle":"2024-12-26T08:24:55.793821Z","shell.execute_reply.started":"2024-12-26T08:22:59.383397Z","shell.execute_reply":"2024-12-26T08:24:55.792179Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import zipfile\nimport os\n\nzip_file_path = '/kaggle/input/inria-bci-challenge/test.zip'\nextract_to = '/kaggle/working/test'\n\nos.makedirs(extract_to, exist_ok=True)\n\nwith zipfile.ZipFile(zip_file_path, 'r') as zip_ref:\n    zip_ref.extractall(extract_to)\n\nprint(f\"File extracted to {extract_to}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:24:55.795811Z","iopub.execute_input":"2024-12-26T08:24:55.796162Z","iopub.status.idle":"2024-12-26T08:25:58.222390Z","shell.execute_reply.started":"2024-12-26T08:24:55.796137Z","shell.execute_reply":"2024-12-26T08:25:58.221169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport pandas as pd\nimport os\nimport glob\nimport mne\nimport matplotlib.pyplot as plt\nimport numpy as np\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:25:58.224821Z","iopub.execute_input":"2024-12-26T08:25:58.225234Z","iopub.status.idle":"2024-12-26T08:26:00.235998Z","shell.execute_reply.started":"2024-12-26T08:25:58.225203Z","shell.execute_reply":"2024-12-26T08:26:00.234474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(mne.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:26:00.237914Z","iopub.execute_input":"2024-12-26T08:26:00.238618Z","iopub.status.idle":"2024-12-26T08:26:00.357189Z","shell.execute_reply.started":"2024-12-26T08:26:00.238567Z","shell.execute_reply":"2024-12-26T08:26:00.355727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_data(data_path):\n    file_list = glob.glob(os.path.join(data_path, \"*.csv\"))\n    data = {}\n    for file in tqdm(file_list, desc=\"Loading files\"):\n        file_name = os.path.basename(file)\n        data[file_name] = pd.read_csv(file)\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:26:00.358479Z","iopub.execute_input":"2024-12-26T08:26:00.358954Z","iopub.status.idle":"2024-12-26T08:26:00.379637Z","shell.execute_reply.started":"2024-12-26T08:26:00.358908Z","shell.execute_reply":"2024-12-26T08:26:00.378136Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Raw_MNE object","metadata":{}},{"cell_type":"code","source":"def create_raw_objects(data, channel_location, sfreq, train_labels):\n    eeg_channel_names = list(channel_locations['Labels'])\n    \n    channel_names = eeg_channel_names + ['EOG']\n    \n    channel_types = ['eeg'] * 56 + ['eog']\n    \n    # Create montage\n    montage = mne.channels.make_dig_montage(\n        ch_pos=dict(zip(channel_locations['Labels'], \n                       zip(channel_locations['Radius'] * np.cos(channel_locations['Phi']),\n                           channel_locations['Radius'] * np.sin(channel_locations['Phi']),\n                           np.zeros(len(channel_locations['Labels']))))),\n        coord_frame='head'\n    )\n    \n    raw_objects = []\n    train_labels_split = [x[:-6] for x in train_labels[\"IdFeedBack\"]]\n    \n    for file_name, df in tqdm(data.items(), desc=\"Creating MNE objects\"):\n\n        eeg_data = df[channel_names[:-1]].to_numpy().T\n        eog_data = df['EOG'].to_numpy()[np.newaxis, :]\n\n        all_data = np.vstack([eeg_data, eog_data])\n        \n        # Create info object\n        info = mne.create_info(ch_names=channel_names, sfreq=sfreq, ch_types=channel_types)\n        \n        # Create raw object\n        raw = mne.io.RawArray(all_data, info)\n        raw.set_montage(montage)\n        \n        # Add annotations\n        feedback_indices = df[df['FeedBackEvent'] == 1].index\n        feedback_times = df.loc[feedback_indices, 'Time'].to_numpy()\n        \n        file_split = file_name[5:-4]\n        indx_train_label = [index for index, x in enumerate(train_labels_split) if x == file_split]\n        descriptions = [train_labels.loc[indx, \"Prediction\"] for indx in indx_train_label]\n        \n        annotations = mne.Annotations(onset=feedback_times, \n                                    duration=[0] * len(feedback_times), \n                                    description=descriptions)\n        raw.set_annotations(annotations)\n        \n        raw_objects.append(raw)\n    \n    return raw_objects","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:26:00.380986Z","iopub.execute_input":"2024-12-26T08:26:00.381362Z","iopub.status.idle":"2024-12-26T08:26:00.394376Z","shell.execute_reply.started":"2024-12-26T08:26:00.381333Z","shell.execute_reply":"2024-12-26T08:26:00.393225Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Preprocess Data","metadata":{}},{"cell_type":"code","source":"def preprocess_raw(raw, visualize=False):\n    raw_processed = raw.copy()\n    \n    raw_processed.filter(\n        l_freq=1,\n        h_freq=40,\n        picks=['eeg', 'eog']\n    )\n    \n    if visualize:\n        raw_processed.plot_psd(fmax=60)\n\n    raw_processed.notch_filter(\n        freqs=50,\n        picks=['eeg']\n    )\n    if visualize:\n        raw_processed.plot_psd(fmax=60)\n\n    ica = mne.preprocessing.ICA(\n        n_components=20,\n        random_state=42,\n        max_iter=1000\n    )\n\n    ica.fit(\n        raw_processed,\n        picks='eeg'\n    )\n    if visualize:\n        ica.plot_components()\n        ica.plot_sources(raw_processed, show_scrollbars=False)\n\n    eog_indices, eog_scores = ica.find_bads_eog(\n        raw_processed,\n        ch_name='EOG'\n    )\n    if visualize:\n        \n        ica.plot_scores(eog_scores)\n        ica.plot_properties(raw_processed, picks=eog_indices)\n    \n    ica.exclude = eog_indices\n    \n    ica.apply(raw_processed)\n    \n    return raw_processed","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:26:00.395538Z","iopub.execute_input":"2024-12-26T08:26:00.395893Z","iopub.status.idle":"2024-12-26T08:26:00.425879Z","shell.execute_reply.started":"2024-12-26T08:26:00.395863Z","shell.execute_reply":"2024-12-26T08:26:00.424512Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Get epochs of each event from each file","metadata":{}},{"cell_type":"code","source":"def extract_epochs(raw, tmin=-0.2, tmax=1.0):\n\n    events = mne.events_from_annotations(raw)\n    epochs = mne.Epochs(raw, events[0], event_id=events[1], tmin=tmin, tmax=tmax,\n                       baseline=(None, 0), preload=True)\n    return epochs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:26:00.429967Z","iopub.execute_input":"2024-12-26T08:26:00.430406Z","iopub.status.idle":"2024-12-26T08:26:00.458653Z","shell.execute_reply.started":"2024-12-26T08:26:00.430365Z","shell.execute_reply":"2024-12-26T08:26:00.457181Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data_path = r\"/kaggle/working/train\"\ndata = load_data(data_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:26:00.461122Z","iopub.execute_input":"2024-12-26T08:26:00.461487Z","iopub.status.idle":"2024-12-26T08:29:45.306183Z","shell.execute_reply.started":"2024-12-26T08:26:00.461452Z","shell.execute_reply":"2024-12-26T08:29:45.305050Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create MNE train object","metadata":{}},{"cell_type":"code","source":"channel_locations = pd.read_csv('/kaggle/input/inria-bci-challenge/ChannelsLocation.csv')\ntrain_labels = pd.read_csv('/kaggle/input/inria-bci-challenge/TrainLabels.csv')\n\nraw_objects = create_raw_objects(data, channel_locations, sfreq=200, train_labels=train_labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:29:45.307485Z","iopub.execute_input":"2024-12-26T08:29:45.307913Z","iopub.status.idle":"2024-12-26T08:30:00.661545Z","shell.execute_reply.started":"2024-12-26T08:29:45.307882Z","shell.execute_reply":"2024-12-26T08:30:00.660390Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Do visualization for the first object on train dataset","metadata":{}},{"cell_type":"code","source":"raw_objects[0].info","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:30:00.662955Z","iopub.execute_input":"2024-12-26T08:30:00.663729Z","iopub.status.idle":"2024-12-26T08:30:00.925741Z","shell.execute_reply.started":"2024-12-26T08:30:00.663643Z","shell.execute_reply":"2024-12-26T08:30:00.924114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"raw_objects[0].annotations\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:30:00.926972Z","iopub.execute_input":"2024-12-26T08:30:00.927489Z","iopub.status.idle":"2024-12-26T08:30:00.934746Z","shell.execute_reply.started":"2024-12-26T08:30:00.927444Z","shell.execute_reply":"2024-12-26T08:30:00.933657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"raw_objects[0].plot_psd(fmax=60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:30:00.935961Z","iopub.execute_input":"2024-12-26T08:30:00.936401Z","iopub.status.idle":"2024-12-26T08:30:02.762770Z","shell.execute_reply.started":"2024-12-26T08:30:00.936360Z","shell.execute_reply":"2024-12-26T08:30:02.761236Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"preprocess_raw(raw_objects[0], visualize=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:30:02.764116Z","iopub.execute_input":"2024-12-26T08:30:02.764542Z","iopub.status.idle":"2024-12-26T08:30:31.104480Z","shell.execute_reply.started":"2024-12-26T08:30:02.764497Z","shell.execute_reply":"2024-12-26T08:30:31.103031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Do preprocessing for train datataset","metadata":{}},{"cell_type":"code","source":"\nX = [] \ny = []  \n\nfor raw in tqdm(raw_objects, desc=\"Processing data\"):\n    raw_processed = preprocess_raw(raw)\n    \n    epochs = extract_epochs(raw_processed)\n    \n    epoch_data = epochs.get_data()\n    \n    X.append(epoch_data)\n    \n    labels = epochs.events[:, 2]\n    \n    y.append(labels)\n\nX = np.concatenate(X, axis=0)\ny = np.concatenate(y, axis=0) \n\ny = y - 1\nprint(f\"Features shape: {X.shape}\")\nprint(f\"Labels shape: {y.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:30:31.105674Z","iopub.execute_input":"2024-12-26T08:30:31.106208Z","iopub.status.idle":"2024-12-26T08:53:54.940358Z","shell.execute_reply.started":"2024-12-26T08:30:31.106160Z","shell.execute_reply":"2024-12-26T08:53:54.938871Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"Features shape: {X.shape}\")\nprint(f\"Labels shape: {y.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:53:54.942744Z","iopub.execute_input":"2024-12-26T08:53:54.943238Z","iopub.status.idle":"2024-12-26T08:53:54.949709Z","shell.execute_reply.started":"2024-12-26T08:53:54.943203Z","shell.execute_reply":"2024-12-26T08:53:54.948423Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Split X and y for training","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nX = X.transpose(0, 2, 1)\nX_train, X_val, y_train, y_val = train_test_split(X, y, test_size=0.2, random_state=42)\n\nprint(f\"Training set shape: X_train: {X_train.shape}, y_train: {y_train.shape}\")\nprint(f\"Validation set shape: X_val: {X_val.shape}, y_val: {y_val.shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:53:54.951033Z","iopub.execute_input":"2024-12-26T08:53:54.951451Z","iopub.status.idle":"2024-12-26T08:53:55.131675Z","shell.execute_reply.started":"2024-12-26T08:53:54.951411Z","shell.execute_reply":"2024-12-26T08:53:55.130534Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load test dataset","metadata":{}},{"cell_type":"code","source":"test_data_path = r\"/kaggle/working/test\"\ntest_data = load_data(test_data_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:53:55.132778Z","iopub.execute_input":"2024-12-26T08:53:55.133127Z","iopub.status.idle":"2024-12-26T08:56:05.187855Z","shell.execute_reply.started":"2024-12-26T08:53:55.133058Z","shell.execute_reply":"2024-12-26T08:56:05.186821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df2 = pd.read_csv('/kaggle/input/true-labels/true_labels.csv', header=None)\ndf2.columns = ['Prediction']\ndf2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:56:05.189000Z","iopub.execute_input":"2024-12-26T08:56:05.189376Z","iopub.status.idle":"2024-12-26T08:56:05.222647Z","shell.execute_reply.started":"2024-12-26T08:56:05.189348Z","shell.execute_reply":"2024-12-26T08:56:05.221675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"desired_order = [\n    'Data_S01_Sess01.csv', 'Data_S01_Sess02.csv', 'Data_S01_Sess03.csv', 'Data_S01_Sess04.csv', 'Data_S01_Sess05.csv',\n    'Data_S03_Sess01.csv', 'Data_S03_Sess02.csv', 'Data_S03_Sess03.csv', 'Data_S03_Sess04.csv', 'Data_S03_Sess05.csv',\n    'Data_S04_Sess01.csv', 'Data_S04_Sess02.csv', 'Data_S04_Sess03.csv', 'Data_S04_Sess04.csv', 'Data_S04_Sess05.csv',\n    'Data_S05_Sess01.csv', 'Data_S05_Sess02.csv', 'Data_S05_Sess03.csv', 'Data_S05_Sess04.csv', 'Data_S05_Sess05.csv',\n    'Data_S08_Sess01.csv', 'Data_S08_Sess02.csv', 'Data_S08_Sess03.csv', 'Data_S08_Sess04.csv', 'Data_S08_Sess05.csv',\n    'Data_S09_Sess01.csv', 'Data_S09_Sess02.csv', 'Data_S09_Sess03.csv', 'Data_S09_Sess04.csv', 'Data_S09_Sess05.csv',\n    'Data_S10_Sess01.csv', 'Data_S10_Sess02.csv', 'Data_S10_Sess03.csv', 'Data_S10_Sess04.csv', 'Data_S10_Sess05.csv',\n    'Data_S15_Sess01.csv', 'Data_S15_Sess02.csv', 'Data_S15_Sess03.csv', 'Data_S15_Sess04.csv', 'Data_S15_Sess05.csv',\n    'Data_S19_Sess01.csv', 'Data_S19_Sess02.csv', 'Data_S19_Sess03.csv', 'Data_S19_Sess04.csv', 'Data_S19_Sess05.csv',\n    'Data_S25_Sess01.csv', 'Data_S25_Sess02.csv', 'Data_S25_Sess03.csv', 'Data_S25_Sess04.csv', 'Data_S25_Sess05.csv'\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:56:05.223935Z","iopub.execute_input":"2024-12-26T08:56:05.224322Z","iopub.status.idle":"2024-12-26T08:56:05.230233Z","shell.execute_reply.started":"2024-12-26T08:56:05.224278Z","shell.execute_reply":"2024-12-26T08:56:05.229168Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = []\ncounter = 0\nfor file_name in desired_order:\n    subject = file_name[5:8]\n    session = file_name[9:15]\n    n_feedbacks = 100 if session == 'Sess05' else 60\n    for fb_num in range(0, n_feedbacks):\n        fb_id = f\"{subject}_{session}_FB{fb_num + 1:03d}\"\n        pred = df2.iloc[counter].values[0]\n        predictions.append({\n                'IdFeedBack': fb_id,\n                'Prediction': pred\n            })\n        counter += 1\npredictions_df = pd.DataFrame(predictions)\npredictions_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:56:05.231634Z","iopub.execute_input":"2024-12-26T08:56:05.232041Z","iopub.status.idle":"2024-12-26T08:56:05.370788Z","shell.execute_reply.started":"2024-12-26T08:56:05.232000Z","shell.execute_reply":"2024-12-26T08:56:05.369614Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create MNE test object","metadata":{}},{"cell_type":"code","source":"test_raw_objects = create_raw_objects(test_data, channel_locations, sfreq=200, train_labels=predictions_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:56:05.372178Z","iopub.execute_input":"2024-12-26T08:56:05.372580Z","iopub.status.idle":"2024-12-26T08:56:15.807962Z","shell.execute_reply.started":"2024-12-26T08:56:05.372541Z","shell.execute_reply":"2024-12-26T08:56:15.806719Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"## Preprocessing test dataset","metadata":{}},{"cell_type":"code","source":"X_test = [] \ny_test = []  \n\nfor raw in tqdm(test_raw_objects, desc=\"Processing data\"):\n    raw_processed = preprocess_raw(raw)\n    \n    epochs = extract_epochs(raw_processed)\n    \n    epoch_data = epochs.get_data()\n    \n    X_test.append(epoch_data)\n    \n    labels = epochs.events[:, 2]\n    \n    y_test.append(labels)\n\nX_test = np.concatenate(X_test, axis=0)\ny_test= np.concatenate(y_test, axis=0) \n\ny_test = y_test - 1\nprint(f\"Features shape: {X_test.shape}\")\nprint(f\"Labels shape: {y_test.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T08:56:15.813234Z","iopub.execute_input":"2024-12-26T08:56:15.813565Z","iopub.status.idle":"2024-12-26T09:07:46.730171Z","shell.execute_reply.started":"2024-12-26T08:56:15.813536Z","shell.execute_reply":"2024-12-26T09:07:46.728762Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model 1 using cnn and lstm","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.nn import BatchNorm1d","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:11:42.510366Z","iopub.execute_input":"2024-12-26T09:11:42.510734Z","iopub.status.idle":"2024-12-26T09:11:42.516050Z","shell.execute_reply.started":"2024-12-26T09:11:42.510704Z","shell.execute_reply":"2024-12-26T09:11:42.514827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:38:25.929525Z","iopub.execute_input":"2024-12-26T09:38:25.929920Z","iopub.status.idle":"2024-12-26T09:38:25.935884Z","shell.execute_reply.started":"2024-12-26T09:38:25.929885Z","shell.execute_reply":"2024-12-26T09:38:25.934799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EEGClassifier(nn.Module):\n    def __init__(self, num_classes=2, input_channels=56):\n        super(EEGClassifier, self).__init__()\n        \n        self.conv1d = nn.Sequential(\n            nn.Conv1d(input_channels, 32, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm1d(32),\n            nn.LeakyReLU(0.2),\n            nn.Dropout(0.2),\n            nn.MaxPool1d(kernel_size=2),\n            \n            nn.Conv1d(32, 64, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm1d(64),\n            nn.LeakyReLU(0.2),\n            nn.Dropout(0.3),\n            nn.MaxPool1d(kernel_size=2),\n            \n            nn.Conv1d(64, 128, kernel_size=5, stride=1, padding=2),\n            nn.BatchNorm1d(128),\n            nn.LeakyReLU(0.2),\n            nn.Dropout(0.4),\n            nn.MaxPool1d(kernel_size=2),\n        )\n\n        self.lstm = nn.LSTM(\n            input_size=128,\n            hidden_size=64,\n            num_layers=1,\n            batch_first=True,\n            bidirectional=True\n        )\n\n        self.classifier = nn.Sequential(\n            nn.Linear(128, 64),  # 128 because of bidirectional (64*2)\n            nn.LeakyReLU(0.2),\n            nn.Dropout(0.5),\n            nn.Linear(64, num_classes)\n        )\n\n    def forward(self, x):\n        batch_size, num_channels, num_samples = x.size()\n        features = self.conv1d(x)\n        features = features.permute(0, 2, 1)\n        lstm_out, _ = self.lstm(features)\n        last_time_step = lstm_out[:, -1, :]\n        output = self.classifier(last_time_step)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:21:50.738927Z","iopub.execute_input":"2024-12-26T09:21:50.739407Z","iopub.status.idle":"2024-12-26T09:21:50.749329Z","shell.execute_reply.started":"2024-12-26T09:21:50.739376Z","shell.execute_reply":"2024-12-26T09:21:50.748052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(self, features, labels, transform=None):\n        self.features = torch.tensor(self.normalize_data(features), dtype=torch.float32)\n        self.labels = torch.tensor(labels, dtype=torch.long)\n        self.transform = transform\n\n    def normalize_data(self, data):\n        mean = np.mean(data, axis=2, keepdims=True)\n        std = np.std(data, axis=2, keepdims=True)\n        return (data - mean) / (std + 1e-8)\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        features = self.features[idx]\n        if self.transform:\n            features = self.transform(features)\n        return features, self.labels[idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:35:10.486016Z","iopub.execute_input":"2024-12-26T09:35:10.486497Z","iopub.status.idle":"2024-12-26T09:35:10.494155Z","shell.execute_reply.started":"2024-12-26T09:35:10.486462Z","shell.execute_reply":"2024-12-26T09:35:10.493029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, train_loader, criterion, optimizer, scheduler, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    for inputs, labels in train_loader:\n        inputs, labels = inputs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        \n\n        l2_lambda = 0.001\n        l2_norm = sum(p.pow(2.0).sum() for p in model.parameters())\n        loss = loss + l2_lambda * l2_norm\n        \n        loss.backward()\n        \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        \n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n    \n    scheduler.step()\n    \n    return running_loss / len(train_loader), 100. * correct / total\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:42:45.043799Z","iopub.execute_input":"2024-12-26T09:42:45.044209Z","iopub.status.idle":"2024-12-26T09:42:45.051985Z","shell.execute_reply.started":"2024-12-26T09:42:45.044176Z","shell.execute_reply":"2024-12-26T09:42:45.050698Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate_one_epoch(model, val_loader, criterion, device):\n\n    model.eval()  \n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    with torch.no_grad(): \n        for inputs, labels in val_loader:\n\n            inputs, labels = inputs.to(device), labels.to(device)\n\n\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n\n    val_loss = running_loss / len(val_loader)\n    val_accuracy = 100. * correct / total\n\n    return val_loss, val_accuracy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:42:47.377898Z","iopub.execute_input":"2024-12-26T09:42:47.378320Z","iopub.status.idle":"2024-12-26T09:42:47.384721Z","shell.execute_reply.started":"2024-12-26T09:42:47.378287Z","shell.execute_reply":"2024-12-26T09:42:47.383555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_and_validate(model, train_loader, val_loader, criterion, optimizer, num_epochs, device):\n    \n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=0.001,\n        epochs=num_epochs,\n        steps_per_epoch=len(train_loader),\n        pct_start=0.3\n    )    \n    \n    history = {\"train_loss\": [], \"train_accuracy\": [], \"val_loss\": [], \"val_accuracy\": []}\n    \n    best_val_acc = 0\n    patience = 10\n    patience_counter = 0\n    \n    for epoch in range(num_epochs):\n        train_loss, train_accuracy = train_one_epoch(\n            model, train_loader, criterion, optimizer,scheduler, device\n        )\n        \n        val_loss, val_accuracy = validate_one_epoch(model, val_loader, criterion, device)\n        \n        if val_accuracy > best_val_acc:\n            best_val_acc = val_accuracy\n            patience_counter = 0\n            # Save best model\n            torch.save(model.state_dict(), 'best_model.pth')\n        else:\n            patience_counter += 1\n            \n        if patience_counter >= patience:\n            print(f\"Early stopping triggered at epoch {epoch+1}\")\n            break\n            \n        print(f\"Epoch [{epoch+1}/{num_epochs}]\")\n        print(f\"    Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.2f}%\")\n        print(f\"    Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.2f}%\")\n        print(f\"    Learning Rate: {scheduler.get_last_lr()[0]:.6f}\")\n        \n    return history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:42:49.694219Z","iopub.execute_input":"2024-12-26T09:42:49.694612Z","iopub.status.idle":"2024-12-26T09:42:49.702726Z","shell.execute_reply.started":"2024-12-26T09:42:49.694577Z","shell.execute_reply":"2024-12-26T09:42:49.701341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch_size = 64\nepochs = 50\n\ntrain_dataset = EEGDataset(X_train, y_train)\nval_dataset = EEGDataset(X_val, y_val)\n\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n\n\n\nmodel = EEGClassifier(num_classes=2, input_channels=X_train.shape[1])  \nmodel.to(device)\n\ncriterion = torch.nn.CrossEntropyLoss()\n    \noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\nhistory = train_and_validate(\n    model=model,\n    train_loader=train_loader,\n    val_loader=val_loader,\n    criterion = criterion,\n    optimizer = optimizer,\n    num_epochs=epochs,\n    device=device\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:42:52.320197Z","iopub.execute_input":"2024-12-26T09:42:52.320715Z","iopub.status.idle":"2024-12-26T09:44:18.562775Z","shell.execute_reply.started":"2024-12-26T09:42:52.320668Z","shell.execute_reply":"2024-12-26T09:44:18.561487Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test(model, test_loader, criterion, device):\n    model.eval()\n\n    total_loss = 0\n    total_correct = 0\n    total_samples = 0\n\n    with torch.no_grad():\n        for inputs, labels in test_loader:\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            outputs = model(inputs)\n\n            loss = criterion(outputs, labels)\n            total_loss += loss.item()\n\n            _, predicted = torch.max(outputs, 1)\n            total_correct += (predicted == labels).sum().item()\n\n            total_samples += labels.size(0)\n\n    avg_loss = total_loss / len(test_loader)\n    accuracy = total_correct / total_samples\n\n    return avg_loss, accuracy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:44:24.968908Z","iopub.execute_input":"2024-12-26T09:44:24.969345Z","iopub.status.idle":"2024-12-26T09:44:24.976744Z","shell.execute_reply.started":"2024-12-26T09:44:24.969310Z","shell.execute_reply":"2024-12-26T09:44:24.975421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dataset = EEGDataset(X_test, y_test)\ntest_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n\ntest_loss, test_accuracy = test(\n    model=model,\n    test_loader=test_loader,\n    criterion=criterion,\n    device=device\n)\n\nprint(f\"Test Loss: {test_loss:.4f}, Test Accuracy: {test_accuracy*100:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:44:27.339733Z","iopub.execute_input":"2024-12-26T09:44:27.340102Z","iopub.status.idle":"2024-12-26T09:44:28.839744Z","shell.execute_reply.started":"2024-12-26T09:44:27.340054Z","shell.execute_reply":"2024-12-26T09:44:28.838611Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model 2","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.models import Sequential\nfrom keras.layers import Conv1D, MaxPooling1D, Dropout, Flatten, Dense, BatchNormalization\nfrom tensorflow.keras.optimizers import SGD\nimport datetime","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:07:22.775132Z","iopub.execute_input":"2024-12-26T10:07:22.775551Z","iopub.status.idle":"2024-12-26T10:07:22.781498Z","shell.execute_reply.started":"2024-12-26T10:07:22.775520Z","shell.execute_reply":"2024-12-26T10:07:22.780134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X_train2 = X_train.transpose(0, 2, 1) \nX_val2 = X_val.transpose(0, 2, 1)\nX_test2 = X_test.transpose(0, 2, 1) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:05:20.251485Z","iopub.execute_input":"2024-12-26T10:05:20.251892Z","iopub.status.idle":"2024-12-26T10:05:20.257985Z","shell.execute_reply.started":"2024-12-26T10:05:20.251861Z","shell.execute_reply":"2024-12-26T10:05:20.256620Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_data(data):\n    mean = np.mean(data, axis=(0, 1), keepdims=True)\n    std = np.std(data, axis=(0, 1), keepdims=True)\n    return (data - mean) / (std + 1e-8)\n\nX_train2 = normalize_data(X_train2)\nX_val2 = normalize_data(X_val2)\nX_test2 = normalize_data(X_test2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:05:22.339117Z","iopub.execute_input":"2024-12-26T10:05:22.339485Z","iopub.status.idle":"2024-12-26T10:05:24.775149Z","shell.execute_reply.started":"2024-12-26T10:05:22.339456Z","shell.execute_reply":"2024-12-26T10:05:24.773736Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_conv_model(lr, input_shape, num_classes=2):\n    model = Sequential([\n        # First Conv Block\n        Conv1D(filters=32, kernel_size=5, activation='relu', input_shape=input_shape,\n               kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Conv1D(filters=32, kernel_size=5, activation='relu',\n               kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Dropout(0.3),\n        MaxPooling1D(pool_size=2),\n        \n        # Second Conv Block\n        Conv1D(filters=64, kernel_size=3, activation='relu',\n               kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Conv1D(filters=64, kernel_size=3, activation='relu',\n               kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Dropout(0.4),\n        MaxPooling1D(pool_size=2),\n        \n        # Third Conv Block\n        Conv1D(filters=128, kernel_size=3, activation='relu',\n               kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Conv1D(filters=128, kernel_size=3, activation='relu',\n               kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Dropout(0.5),\n        MaxPooling1D(pool_size=2),\n        \n        # Dense Layers\n        Flatten(),\n        Dense(128, activation='relu', kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Dropout(0.5),\n        Dense(64, activation='relu', kernel_regularizer=tf.keras.regularizers.l2(0.01)),\n        BatchNormalization(),\n        Dropout(0.5),\n        Dense(num_classes, activation='softmax')\n    ])\n    \n    # Use Adam optimizer with learning rate scheduling\n    optimizer = tf.keras.optimizers.Adam(\n        learning_rate=lr,\n        beta_1=0.9,\n        beta_2=0.999,\n        epsilon=1e-07\n    )\n    \n    model.compile(\n        loss='sparse_categorical_crossentropy',\n        optimizer=optimizer,\n        metrics=['accuracy']\n    )\n    \n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:05:47.622509Z","iopub.execute_input":"2024-12-26T10:05:47.622909Z","iopub.status.idle":"2024-12-26T10:05:47.636434Z","shell.execute_reply.started":"2024-12-26T10:05:47.622875Z","shell.execute_reply":"2024-12-26T10:05:47.635383Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"input_shape = (X_train2.shape[1], X_train2.shape[2])\nmodel = create_conv_model(lr=0.001, input_shape=input_shape)\n\n# Model summary\nmodel.summary()\n\n# Callbacks\nreduce_lr = tf.keras.callbacks.ReduceLROnPlateau(\n    monitor='val_loss',\n    factor=0.5,\n    patience=3,\n    min_lr=1e-6,\n    verbose=1\n)\n\n# Create unique log directory for TensorBoard\nlog_dir = \"logs/fit/\" + datetime.datetime.now().strftime(\"%Y%m%d-%H%M%S\")\ntensorboard_callback = tf.keras.callbacks.TensorBoard(\n    log_dir=log_dir,\n    histogram_freq=1,\n    update_freq='epoch'\n)\n\n\n# Early stopping\nearly_stopping = tf.keras.callbacks.EarlyStopping(\n    monitor='val_loss',\n    patience=10,\n    restore_best_weights=True,\n    verbose=1\n)\n\n# Train the model\nhistory = model.fit(\n    x=X_train2,\n    y=y_train,\n    epochs=100,\n    batch_size=32,\n    validation_data=(X_val2, y_val),\n    class_weight={0: 1, 1: 1},\n    callbacks=[\n        tensorboard_callback,\n        early_stopping,\n        reduce_lr\n    ],\n    verbose=1\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:07:29.006046Z","iopub.execute_input":"2024-12-26T10:07:29.006471Z","iopub.status.idle":"2024-12-26T10:10:36.439295Z","shell.execute_reply.started":"2024-12-26T10:07:29.006438Z","shell.execute_reply":"2024-12-26T10:10:36.438020Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:23:35.986719Z","iopub.execute_input":"2024-12-26T10:23:35.987191Z","iopub.status.idle":"2024-12-26T10:23:35.992437Z","shell.execute_reply.started":"2024-12-26T10:23:35.987157Z","shell.execute_reply":"2024-12-26T10:23:35.990976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = model.predict(X_test2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:22:21.485303Z","iopub.execute_input":"2024-12-26T10:22:21.485766Z","iopub.status.idle":"2024-12-26T10:22:23.440878Z","shell.execute_reply.started":"2024-12-26T10:22:21.485735Z","shell.execute_reply":"2024-12-26T10:22:23.439588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predicted_classes = np.argmax(predictions, axis=1)\naccuracy = accuracy_score(y_test, predicted_classes)\nprint(f\"Test Accuracy: {accuracy*100:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T10:23:46.102245Z","iopub.execute_input":"2024-12-26T10:23:46.102591Z","iopub.status.idle":"2024-12-26T10:23:46.111412Z","shell.execute_reply.started":"2024-12-26T10:23:46.102566Z","shell.execute_reply":"2024-12-26T10:23:46.110178Z"}},"outputs":[],"execution_count":null}]}