{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":11689210,"sourceType":"datasetVersion","datasetId":7336709}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Gerekli İmportlar\nimport os\nimport pandas as pd\nimport numpy as np\nfrom tensorflow.keras import layers, models\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.applications import DenseNet121\nfrom tensorflow.keras.applications import ResNet50V2\nfrom tensorflow.keras.applications import NASNetMobile\nfrom tensorflow.keras.applications import InceptionResNetV2\nfrom tensorflow.keras.applications import DenseNet121\nfrom tensorflow.keras.layers import Input, Conv2D, MaxPooling2D, Flatten, Dense, Dropout, Concatenate, GlobalAveragePooling2D\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout, BatchNormalization\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Flatten, Dense, Dropout, BatchNormalization\nimport tensorflow as tf\nfrom sklearn.utils import class_weight\nfrom tensorflow.keras.callbacks import ModelCheckpoint\nfrom tensorflow.keras.models import load_model\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\nfrom tensorflow.keras.layers import Input, Concatenate, GlobalAveragePooling2D, Dense\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.applications import ResNet50V2\nimport tensorflow as tf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:11:39.094993Z","iopub.execute_input":"2025-05-11T11:11:39.095160Z","iopub.status.idle":"2025-05-11T11:11:53.628301Z","shell.execute_reply.started":"2025-05-11T11:11:39.095144Z","shell.execute_reply":"2025-05-11T11:11:53.627779Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Eldeki tüm veriyi ayrıca import edelim\nBASE_PATH = \"/kaggle/input/hms-harmful-brain-activity-classification\"\neeg_path = BASE_PATH+\"/\"+\"train_eegs\"\npd.read_parquet(eeg_path+\"/\"+os.listdir(eeg_path)[0])\n\ncsv = pd.read_csv(BASE_PATH+\"/train.csv\")\nunique_eeg_ids_df = csv.drop_duplicates(subset='eeg_id')\nunique_eeg_ids_df = unique_eeg_ids_df[['eeg_id', 'expert_consensus']]\nunique_values = unique_eeg_ids_df['expert_consensus'].unique()\nprint(unique_values)\nlabels = {0:\"Seizure\",1:\"GPD\",2:\"LRDA\",3:\"LPD\",4:\"GRDA\",5:\"Other\"}\n# Invert the labels dictionary to map string labels to their numeric values\nlabel_map = {v: k for k, v in labels.items()}\n\n# Replace the string values in the expert_consensus column with their numeric values\nunique_eeg_ids_df['expert_consensus'] = unique_eeg_ids_df['expert_consensus'].map(label_map)\nunique_eeg_ids_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:12:41.955136Z","iopub.execute_input":"2025-05-11T11:12:41.955759Z","iopub.status.idle":"2025-05-11T11:12:42.107149Z","shell.execute_reply.started":"2025-05-11T11:12:41.955734Z","shell.execute_reply":"2025-05-11T11:12:42.106351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df = pd.read_csv(\"/kaggle/input/hms-dataset-split/test_df.csv\")\nuse_df = pd.read_csv(\"/kaggle/input/hms-dataset-split/use_df.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:13:19.432886Z","iopub.execute_input":"2025-05-11T11:13:19.433158Z","iopub.status.idle":"2025-05-11T11:13:19.463180Z","shell.execute_reply.started":"2025-05-11T11:13:19.433138Z","shell.execute_reply":"2025-05-11T11:13:19.462686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pywt\nprint(\"The wavelet functions we can use:\")\nprint(pywt.wavelist())\n\nUSE_WAVELET = None #or \"db8\" or anything below\n\n# DENOISE FUNCTION\ndef maddest(d, axis=None):\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\ndef denoise(x, wavelet='haar', level=1):    \n    coeff = pywt.wavedec(x, wavelet, mode=\"per\")\n    sigma = (1/0.6745) * maddest(coeff[-level])\n\n    uthresh = sigma * np.sqrt(2*np.log(len(x)))\n    coeff[1:] = (pywt.threshold(i, value=uthresh, mode='hard') for i in coeff[1:])\n\n    ret=pywt.waverec(coeff, wavelet, mode='per')\n    \n    return ret\n\nimport librosa\n\ndef spectrogram_from_eeg(parquet_path, display=False):\n    parquet_path = BASE_PATH+\"/train_eegs/\"+str(parquet_path)+\".parquet\"\n    # LOAD MIDDLE 50 SECONDS OF EEG SERIES\n    eeg = pd.read_parquet(parquet_path)\n    middle = (len(eeg)-10_000)//2\n    eeg = eeg.iloc[middle:middle+10_000]\n    \n    # VARIABLE TO HOLD SPECTROGRAM\n    img = np.zeros((128,256,4),dtype='float32')\n    \n    if display: plt.figure(figsize=(10,7))\n    signals = []\n    for k in range(4):\n        COLS = FEATS[k]\n        \n        for kk in range(4):\n        \n            # COMPUTE PAIR DIFFERENCES\n            x = eeg[COLS[kk]].values - eeg[COLS[kk+1]].values\n\n            # FILL NANS\n            m = np.nanmean(x)\n            if np.isnan(x).mean()<1: x = np.nan_to_num(x,nan=m)\n            else: x[:] = 0\n\n            # DENOISE\n            if USE_WAVELET:\n                x = denoise(x, wavelet=USE_WAVELET)\n            signals.append(x)\n\n            # RAW SPECTROGRAM\n            mel_spec = librosa.feature.melspectrogram(y=x, sr=200, hop_length=len(x)//256, \n                  n_fft=1024, n_mels=128, fmin=0, fmax=20, win_length=128)\n\n            # LOG TRANSFORM\n            width = (mel_spec.shape[1]//32)*32\n            mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[:,:width]\n\n            # STANDARDIZE TO -1 TO 1\n            mel_spec_db = (mel_spec_db+40)/40 \n            img[:,:,k] += mel_spec_db\n                \n        # AVERAGE THE 4 MONTAGE DIFFERENCES\n        img[:,:,k] /= 4.0\n        \n        if display:\n            plt.subplot(2,2,k+1)\n            plt.imshow(img[:,:,k],aspect='auto',origin='lower')\n            plt.title(f'EEG {eeg_id} - Spectrogram {NAMES[k]}')\n            \n    if display: \n        plt.show()\n        plt.figure(figsize=(10,5))\n        offset = 0\n        for k in range(4):\n            if k>0: offset -= signals[3-k].min()\n            plt.plot(range(10_000),signals[k]+offset,label=NAMES[3-k])\n            offset += signals[3-k].max()\n        plt.legend()\n        plt.title(f'EEG {eeg_id} Signals')\n        plt.show()\n        print(); print('#'*25); print()\n        \n    return img\n\ndef spectrogram_from_eeg_2d(eeg_id, display=False):\n    data = spectrogram_from_eeg(eeg_id)\n    concatenated_image = np.vstack((np.hstack((data[:, :, 0], data[:, :, 1])), \n                                    np.hstack((data[:, :, 2], data[:, :, 3]))))\n    return concatenated_image\n\nNAMES = ['LL','LP','RP','RR']\n\nFEATS = [['Fp1','F7','T3','T5','O1'],\n         ['Fp1','F3','C3','P3','O1'],\n         ['Fp2','F8','T4','T6','O2'],\n         ['Fp2','F4','C4','P4','O2']]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:14:19.327117Z","iopub.execute_input":"2025-05-11T11:14:19.327657Z","iopub.status.idle":"2025-05-11T11:14:19.437410Z","shell.execute_reply.started":"2025-05-11T11:14:19.327607Z","shell.execute_reply":"2025-05-11T11:14:19.436693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# veriyi hazırlayalım\ntry:\n    os.mkdir(\"np_files\")\nexcept:\n    pass\n\n\ncount = 0\nfor index, row in unique_eeg_ids_df.iterrows():\n    count += 1\n    data = spectrogram_from_eeg_2d(row[\"eeg_id\"])\n    if count % 100 == 0:\n        print(count, \",\", end=\"\")\n    np.save(f\"np_files/{row['eeg_id']}.npy\", data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:14:51.081856Z","iopub.execute_input":"2025-05-11T11:14:51.082829Z","iopub.status.idle":"2025-05-11T11:49:46.259923Z","shell.execute_reply.started":"2025-05-11T11:14:51.082800Z","shell.execute_reply":"2025-05-11T11:49:46.259078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Modeli oluştur\nmodel = Sequential()\n\n# İlk Convolutional katmanı ve MaxPooling katmanı\nmodel.add(Conv2D(32, (3, 3), activation='relu', input_shape=(512, 256, 1)))\nmodel.add(MaxPooling2D((2, 2)))\n\n# İkinci Convolutional katmanı ve MaxPooling katmanı\nmodel.add(Conv2D(64, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D((2, 2)))\n\n# Üçüncü Convolutional katmanı ve MaxPooling katmanı\nmodel.add(Conv2D(128, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D((2, 2)))\n\n# Dördüncü Convolutional katmanı ve MaxPooling katmanı\nmodel.add(Conv2D(256, (3, 3), activation='relu'))\nmodel.add(MaxPooling2D((2, 2)))\n\n# Flatten katmanı\nmodel.add(Flatten())\n\n# Tam bağlantılı katmanlar\nmodel.add(Dense(512, activation='relu'))\nmodel.add(Dropout(0.5))\nmodel.add(Dense(256, activation='relu'))\nmodel.add(Dropout(0.5))\n\n# Çıkış katmanı\nmodel.add(Dense(2, activation='softmax'))\n\n\n# Derleme\nmodel.compile(\n    optimizer='adam',\n    loss='categorical_crossentropy',\n    metrics=['accuracy']\n)\n\n# Özet\n# model.summary()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import keras.utils\n\nclass EEGDataGenerator(keras.utils.Sequence):\n    \"\"\"\n    Data generator for EEG spectrograms for Keras.\n    Converts EEG IDs to spectrograms using the provided function and returns batches.\n    \"\"\"\n    \n    def __init__(self, dataframe, spectrogram_function, batch_size=32, \n                 shuffle=True, seed=None, is_test=False):\n        \"\"\"\n        Initialize the data generator.\n        \n        Args:\n            dataframe (pd.DataFrame): DataFrame containing 'eeg_id' and 'expert_consensus' columns\n            spectrogram_function (callable): Function that converts eeg_id to spectrogram array\n            batch_size (int): Size of batches to generate\n            shuffle (bool): Whether to shuffle the data after each epoch\n            seed (int): Random seed for reproducibility\n            is_test (bool): If True, don't return labels (for prediction)\n        \"\"\"\n        self.df = dataframe.copy()\n        self.batch_size = batch_size\n        self.spectrogram_function = spectrogram_function\n        self.shuffle = shuffle\n        self.seed = seed\n        self.is_test = is_test\n        \n        # Generate indices\n        self.indices = np.arange(len(self.df))\n        \n        # Class mapping if needed\n        self.classes = sorted(self.df['expert_consensus'].unique())\n        self.class_indices = {cls: i for i, cls in enumerate(self.classes)}\n        \n        # Initial shuffle\n        if self.shuffle:\n            np.random.seed(self.seed)\n            np.random.shuffle(self.indices)\n    \n    def __len__(self):\n        \"\"\"Denotes the number of batches per epoch\"\"\"\n        return int(np.ceil(len(self.df) / self.batch_size))\n    \n    def __getitem__(self, index):\n        \"\"\"Generate one batch of data\"\"\"\n        # Generate indices of the batch\n        batch_indices = self.indices[index * self.batch_size:(index + 1) * self.batch_size]\n        \n        # Get batch data\n        batch_df = self.df.iloc[batch_indices]\n        \n        # Generate spectrograms\n        batch_x = np.array([\n            self.spectrogram_function(eeg_id) \n            for eeg_id in batch_df['eeg_id']\n        ])\n        \n        if self.is_test:\n            return batch_x\n        \n        # Generate labels (one-hot encoded)\n        batch_y = np.array([\n            self.class_indices[label] \n            for label in batch_df['expert_consensus']\n        ])\n        \n        return batch_x, tf.keras.utils.to_categorical(batch_y, num_classes=len(self.classes))\n    \n    def on_epoch_end(self):\n        \"\"\"Updates indices after each epoch\"\"\"\n        if self.shuffle:\n            np.random.seed(self.seed)\n            np.random.shuffle(self.indices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:58:04.581012Z","iopub.execute_input":"2025-05-11T11:58:04.581731Z","iopub.status.idle":"2025-05-11T11:58:04.590987Z","shell.execute_reply.started":"2025-05-11T11:58:04.581697Z","shell.execute_reply":"2025-05-11T11:58:04.590332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_img(eeg_id):\n    data = np.load(f\"np_files/{eeg_id}.npy\")\n    return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:57:57.500868Z","iopub.execute_input":"2025-05-11T11:57:57.501509Z","iopub.status.idle":"2025-05-11T11:57:57.505070Z","shell.execute_reply.started":"2025-05-11T11:57:57.501486Z","shell.execute_reply":"2025-05-11T11:57:57.504330Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_eeg_generators(train_df, val_df, test_df, spectrogram_from_eeg, \n                          batch_size=32, seed=42):\n\n    # Create generators\n    train_generator = EEGDataGenerator(\n        dataframe=train_df,\n        spectrogram_function=get_img,\n        batch_size=batch_size,\n        shuffle=True,\n        seed=seed,\n        is_test=False\n    )\n    \n    val_generator = EEGDataGenerator(\n        dataframe=val_df,\n        spectrogram_function=get_img,\n        batch_size=batch_size,\n        shuffle=False,\n        seed=seed,\n        is_test=False\n    )\n    \n    test_generator = EEGDataGenerator(\n        dataframe=test_df,\n        spectrogram_function=get_img,\n        batch_size=batch_size,\n        shuffle=False,\n        seed=seed,\n        is_test=True\n    )\n    \n    return train_generator, val_generator, test_generator\n\n\ntrain_df, val_df = train_test_split(use_df, test_size=0.2, random_state=42) # %30 test + validasyon\ntrain_generator, val_generator, test_generator = create_eeg_generators(train_df, val_df, test_df, get_img, batch_size=32, seed=42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T11:58:05.846244Z","iopub.execute_input":"2025-05-11T11:58:05.846735Z","iopub.status.idle":"2025-05-11T11:58:05.856837Z","shell.execute_reply.started":"2025-05-11T11:58:05.846712Z","shell.execute_reply":"2025-05-11T11:58:05.856053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"GPU Available:\", tf.config.list_physical_devices('GPU'))\n\n# Checkpoint callback: her epoch sonunda modeli kaydeder\ncheckpoint_cb = ModelCheckpoint(\n    filepath='model_epoch_{epoch:02d}.keras',  # örnek: model_epoch_01.keras\n    save_freq='epoch',\n    save_weights_only=False,  # modeli tam olarak kaydetsin (ağırlık + yapı)\n    verbose=1\n)\n\n# Modeli eğit\nhistory = model.fit(\n    train_generator,\n    validation_data=val_generator,\n    epochs=10,\n    callbacks=[checkpoint_cb]\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport torch\nclass NumpyImageDataset(Dataset):\n    def __init__(self, images, labels):\n        self.images = images\n        self.labels = labels\n        self.transform = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((224, 224)),\n            transforms.ToTensor(),  # 0-255 → 0-1, shape: (C, H, W)\n        ])\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        img = self.images[idx]\n\n        # Eğer tek kanal ise, 3 kanala genişlet\n        if img.ndim == 2:\n            img = np.stack([img] * 3, axis=-1)  # (H, W) → (H, W, 3)\n        elif img.shape[2] == 1:\n            img = np.repeat(img, 3, axis=2)  # (H, W, 1) → (H, W, 3)\n\n        img = self.transform(img)\n        label = self.labels[idx]\n        return img, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:01:22.616253Z","iopub.execute_input":"2025-05-11T12:01:22.616544Z","iopub.status.idle":"2025-05-11T12:01:22.622613Z","shell.execute_reply.started":"2025-05-11T12:01:22.616524Z","shell.execute_reply":"2025-05-11T12:01:22.621886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import ViTForImageClassification, ViTFeatureExtractor\n\nfeature_extractor = ViTFeatureExtractor.from_pretrained(\"google/vit-base-patch16-224-in21k\")\n\nmodel = ViTForImageClassification.from_pretrained(\n    \"google/vit-base-patch16-224-in21k\",\n    num_labels=6  # 6 sınıf\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:01:32.800082Z","iopub.execute_input":"2025-05-11T12:01:32.800355Z","iopub.status.idle":"2025-05-11T12:01:33.062097Z","shell.execute_reply.started":"2025-05-11T12:01:32.800337Z","shell.execute_reply":"2025-05-11T12:01:33.061523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Val verilerini hazırla\nval_images = [get_img(eeg_id) for eeg_id in val_df['eeg_id'].values]\nval_labels = val_df['expert_consensus'].values\n\n# Dataset ve DataLoader oluştur\nval_dataset = NumpyImageDataset(val_images, val_labels)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:01:35.076430Z","iopub.execute_input":"2025-05-11T12:01:35.076721Z","iopub.status.idle":"2025-05-11T12:01:36.690942Z","shell.execute_reply.started":"2025-05-11T12:01:35.076699Z","shell.execute_reply":"2025-05-11T12:01:36.690113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Val verilerini hazırla\ntrain_images = [get_img(eeg_id) for eeg_id in train_df['eeg_id'].values]\ntrain_labels = train_df['expert_consensus'].values\n\n# Dataset ve DataLoader oluştur\ntrain_dataset = NumpyImageDataset(train_images, train_labels)\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:01:36.692164Z","iopub.execute_input":"2025-05-11T12:01:36.692463Z","iopub.status.idle":"2025-05-11T12:01:53.002406Z","shell.execute_reply.started":"2025-05-11T12:01:36.692446Z","shell.execute_reply":"2025-05-11T12:01:53.001844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(model, val_loader, device='cpu'):\n    model.eval()\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images).logits\n            _, predicted = torch.max(outputs, 1)\n            correct += (predicted == labels).sum().item()\n            total += labels.size(0)\n    model.train()\n    return correct / total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:01:53.003482Z","iopub.execute_input":"2025-05-11T12:01:53.003740Z","iopub.status.idle":"2025-05-11T12:01:53.008246Z","shell.execute_reply.started":"2025-05-11T12:01:53.003717Z","shell.execute_reply":"2025-05-11T12:01:53.007661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\n\nloss_fn = nn.CrossEntropyLoss()\n\nfrom torch.optim import Adam\n\noptimizer = Adam(model.parameters(), lr=1e-4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:02:36.028311Z","iopub.execute_input":"2025-05-11T12:02:36.028584Z","iopub.status.idle":"2025-05-11T12:02:36.033276Z","shell.execute_reply.started":"2025-05-11T12:02:36.028563Z","shell.execute_reply":"2025-05-11T12:02:36.032556Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\nfor epoch in range(10):\n    print(f\"\\nEpoch {epoch+1}/10\")\n    model.train()\n\n    total_loss = 0\n    for images, labels in tqdm(train_loader, desc=f\"Training Epoch {epoch+1}\"):\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = model(images).logits\n        loss = loss_fn(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n\n        total_loss += loss.item()\n\n    avg_loss = total_loss / len(train_loader)\n    val_acc = evaluate(model, val_loader, device)\n    print(f\"Epoch {epoch+1} finished. Avg Loss: {avg_loss:.4f} - Val Accuracy: {val_acc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T12:06:17.386330Z","iopub.execute_input":"2025-05-11T12:06:17.386986Z","iopub.status.idle":"2025-05-11T13:31:50.729668Z","shell.execute_reply.started":"2025-05-11T12:06:17.386955Z","shell.execute_reply":"2025-05-11T13:31:50.728952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_images = np.stack([get_img(eeg_id) for eeg_id in test_df[\"eeg_id\"]])\ntest_labels = test_df[\"expert_consensus\"].values\n\ntest_dataset = NumpyImageDataset(test_images, test_labels)\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)\n\ndef evaluate(model, dataloader, device):\n    model.eval()\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for images, labels in dataloader:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images).logits\n            _, preds = torch.max(outputs, 1)\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n    return correct / total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T13:33:19.975083Z","iopub.execute_input":"2025-05-11T13:33:19.975813Z","iopub.status.idle":"2025-05-11T13:33:26.080240Z","shell.execute_reply.started":"2025-05-11T13:33:19.975788Z","shell.execute_reply":"2025-05-11T13:33:26.079686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_acc = evaluate(model, test_loader, device)\nprint(f\"Test Accuracy: {test_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T13:33:27.996244Z","iopub.execute_input":"2025-05-11T13:33:27.996517Z","iopub.status.idle":"2025-05-11T13:33:51.290150Z","shell.execute_reply.started":"2025-05-11T13:33:27.996495Z","shell.execute_reply":"2025-05-11T13:33:51.289536Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay\nimport matplotlib.pyplot as plt\n\nall_preds = []\nall_labels = []\n\nmodel.eval()\nwith torch.no_grad():\n    for images, labels in test_loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = model(images).logits\n        preds = torch.argmax(outputs, dim=1)\n\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\ncm = confusion_matrix(all_labels, all_preds)\ndisp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=[\n    \"Seizure\", \"GPD\", \"LRDA\", \"LPD\", \"GRDA\", \"Other\"\n])\n\nfig, ax = plt.subplots(figsize=(8, 6))\ndisp.plot(ax=ax, cmap=\"Blues\", values_format='d')\nplt.title(\"Confusion Matrix on Test Set\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T13:46:16.025048Z","iopub.execute_input":"2025-05-11T13:46:16.025365Z","iopub.status.idle":"2025-05-11T13:46:41.302267Z","shell.execute_reply.started":"2025-05-11T13:46:16.025344Z","shell.execute_reply":"2025-05-11T13:46:41.301422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"torch.save(model.state_dict(), \"vit_model_weights.pth\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T13:44:28.799346Z","iopub.execute_input":"2025-05-11T13:44:28.800064Z","iopub.status.idle":"2025-05-11T13:44:29.249896Z","shell.execute_reply.started":"2025-05-11T13:44:28.800039Z","shell.execute_reply":"2025-05-11T13:44:29.249206Z"}},"outputs":[],"execution_count":null}]}