{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":163569236,"sourceType":"kernelVersion"}],"dockerImageVersionId":30646,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os,sys\nimport shutil\nimport torch\n\nCODE_DIR = '/kaggle/working/code'\nos.makedirs(CODE_DIR, exist_ok=True)\n\nCODE_DIR_SRC = '/kaggle/input/datapreprocessing/code'\n# List all files in the source directory\nfiles = os.listdir(CODE_DIR_SRC)\n# Iterate over each file and copy it to the destination directory\nfor file in files:\n    # Construct full file paths\n    source_file = os.path.join(CODE_DIR_SRC, file)\n    if source_file.endswith('__pycache__'):\n        continue\n    destination_file = os.path.join(CODE_DIR, file)  \n    # Copy the file\n    shutil.copy(source_file, destination_file)\n\nsys.path.append(CODE_DIR)\n\nMODEL_DIR = '/kaggle/working/models'\nos.makedirs(MODEL_DIR, exist_ok=True)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\nTRAIN_FILE = '/kaggle/input/datapreprocessing/preprocessed/train.csv'\nSPC_DIR = '/kaggle/input/datapreprocessing/preprocessed/spectrogram'\n\nTEST_SPC_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms'\nTEST_PROCESSED_SPC_DIR = \"/kaggle/working/preprocessed/spectrogram\"\nos.makedirs(TEST_PROCESSED_SPC_DIR, exist_ok=True)\n\nTRAIN_EEG_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs'\nTRAIN_SPC_DIR = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms'\n\nTEST_FILE = '/kaggle/input/hms-harmful-brain-activity-classification/test.csv'\nSUBMISSION_FILE = '/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv'\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-22T11:11:52.442267Z","iopub.execute_input":"2024-02-22T11:11:52.442541Z","iopub.status.idle":"2024-02-22T11:11:54.597204Z","shell.execute_reply.started":"2024-02-22T11:11:52.442518Z","shell.execute_reply":"2024-02-22T11:11:54.596340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Hyperparameters\nBATCH_SIZE = 50\nHIDDEN_SIZE = 100\nLEARNING_RATE = 0.001\nNUM_EPOCHS = 100\nNUM_SPLITS = 10","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:11:59.081572Z","iopub.execute_input":"2024-02-22T11:11:59.082458Z","iopub.status.idle":"2024-02-22T11:11:59.086807Z","shell.execute_reply.started":"2024-02-22T11:11:59.082425Z","shell.execute_reply":"2024-02-22T11:11:59.085781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = pd.read_csv(TRAIN_FILE)\ndel data['Unnamed: 0']\ndata.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:02.257523Z","iopub.execute_input":"2024-02-22T11:12:02.258281Z","iopub.status.idle":"2024-02-22T11:12:02.637758Z","shell.execute_reply.started":"2024-02-22T11:12:02.258249Z","shell.execute_reply":"2024-02-22T11:12:02.636796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile {CODE_DIR}/spec_dataset.py\n\nimport os\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\n# Define custom dataset\nclass SpectrogramDataset(Dataset):\n    def __init__(self, train_df, SPC_DIR):\n        self.train = train_df  \n        self.SPC_DIR = SPC_DIR\n\n    def __len__(self):\n        return len(self.train)\n\n    def __getitem__(self, idx):\n        spc_file = f\"{self.train.spectrogram_id[idx]}.pt\"\n        spc_path = os.path.join(self.SPC_DIR, spc_file)\n        spc = torch.load(spc_path)\n        \n        spc_start = self.train.spec_start_idx[idx]\n        spc_end = self.train.spec_end_idx[idx]\n        spc = spc[spc_start:spc_end]\n        \n        labels = ['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote',\t'other_vote']\n        y = self.train.loc[idx,labels].astype('float').values\n        y = y/np.sum(y)\n        return spc, y","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:06.245171Z","iopub.execute_input":"2024-02-22T11:12:06.245573Z","iopub.status.idle":"2024-02-22T11:12:06.253163Z","shell.execute_reply.started":"2024-02-22T11:12:06.245530Z","shell.execute_reply":"2024-02-22T11:12:06.252272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile {CODE_DIR}/spec_test_dataset.py\n\nimport os\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\n\n# Define custom dataset\nclass SpectrogramTestDataset(Dataset):\n    def __init__(self, test_df, SPC_DIR):\n        self.test = test_df  \n        self.SPC_DIR = SPC_DIR\n\n    def __len__(self):\n        return len(self.test)\n\n    def __getitem__(self, idx):\n        spc_file = f\"{self.test.spectrogram_id[idx]}.pt\"\n        spc_path = os.path.join(self.SPC_DIR, spc_file)\n        spc = torch.load(spc_path)\n        return spc","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:14.460640Z","iopub.execute_input":"2024-02-22T11:12:14.461377Z","iopub.status.idle":"2024-02-22T11:12:14.467225Z","shell.execute_reply.started":"2024-02-22T11:12:14.461347Z","shell.execute_reply":"2024-02-22T11:12:14.466211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile {CODE_DIR}/spec_dataset_builder.py\n\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.preprocessing import OrdinalEncoder\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import DataLoader\n\nfrom spec_dataset import SpectrogramDataset\nfrom spec_test_dataset import SpectrogramTestDataset\n\nSPC_DIR = '/kaggle/input/datapreprocessing/preprocessed/spectrogram'\n\nclass DataSetBuilder:\n    def __init__(self, \n                 data, \n                 target_column= 'expert_consensus', \n                 n_splits=10, \n                 random_state=42, \n                 batch_size=32, \n                 SPC_DIR=SPC_DIR,\n                ):\n        self.data = data\n        self.target_column = target_column\n        self.n_splits = n_splits\n        self.random_state = random_state\n        self.batch_size = batch_size\n        self.SPC_DIR = SPC_DIR\n        self.encoder = OrdinalEncoder()\n\n    def stratified_split(self, data):\n        y = self.encoder.fit_transform(data[[self.target_column]])\n        X = data.drop(columns=[self.target_column])\n        skf = StratifiedKFold(n_splits=self.n_splits, shuffle=True, random_state=self.random_state)\n        stratified_splits = []\n        for train_index, _ in skf.split(X, y):\n            train_data = self.data.iloc[train_index].copy().reset_index(drop=True)\n            stratified_splits.append(train_data)\n        return stratified_splits\n    \n    def make_datasets(self, data):\n        data_splits = self.stratified_split(data)\n        dataset_list = [SpectrogramDataset(df, self.SPC_DIR) for df in data_splits]\n        return dataset_list\n    \n    def make_test_dataloader(test_data, TEST_SPC_DIR):\n        dataset = SpectrogramTestDataset(test_data, TEST_SPC_DIR)\n        dataloader = DataLoader(dataset, batch_size=1, pin_memory=True) \n        return dataloader\n        \n    def make_dataloaders(self):\n        k_datasets = self.make_datasets(self.data)\n        dataloaders = [DataLoader(dataset, batch_size=self.batch_size, shuffle=True, pin_memory=True) for dataset in k_datasets]\n        return dataloaders","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:18.534642Z","iopub.execute_input":"2024-02-22T11:12:18.535014Z","iopub.status.idle":"2024-02-22T11:12:18.542337Z","shell.execute_reply.started":"2024-02-22T11:12:18.534985Z","shell.execute_reply":"2024-02-22T11:12:18.541433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from spec_dataset_builder import DataSetBuilder\n\nspec_datamaker = DataSetBuilder(\n    data = data,\n    target_column = 'expert_consensus',\n    n_splits = NUM_SPLITS,\n    batch_size=BATCH_SIZE,\n    SPC_DIR = SPC_DIR\n)\n\ndata_loaders = spec_datamaker.make_dataloaders()","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:24.598301Z","iopub.execute_input":"2024-02-22T11:12:24.598962Z","iopub.status.idle":"2024-02-22T11:12:25.466693Z","shell.execute_reply.started":"2024-02-22T11:12:24.598932Z","shell.execute_reply":"2024-02-22T11:12:25.465674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from preprocess_specs import SpectrogramDatasetBuilder\n\ntest_data = pd.read_csv(TEST_FILE)\nspectrogram_builder = SpectrogramDatasetBuilder()\nspectrogram_builder.buildAll(TEST_SPC_DIR)\n\n#subm_spec_file = f\"{TEST_PROCESSED_SPC_DIR}/{test_df['spectrogram_id'][0]}.pt\"\n\nsub_dataloader = DataSetBuilder.make_test_dataloader(test_data, \n                                                     TEST_SPC_DIR = TEST_PROCESSED_SPC_DIR)","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:28.665695Z","iopub.execute_input":"2024-02-22T11:12:28.666234Z","iopub.status.idle":"2024-02-22T11:12:29.979103Z","shell.execute_reply.started":"2024-02-22T11:12:28.666179Z","shell.execute_reply":"2024-02-22T11:12:29.978164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile {CODE_DIR}/model.py\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\n# MODEL PARAMETERS\nINPUT_SIZE = 100\nOUTPUT_SIZE = 6\n\nclass GRUClassifier(nn.Module):\n    def __init__(self, \n                 input_size=INPUT_SIZE, \n                 hidden_size=INPUT_SIZE, \n                 output_size=OUTPUT_SIZE):\n        \n        super(GRUClassifier, self).__init__()\n        self.hidden_size = hidden_size\n        \n        self.gru_list = nn.ModuleList([nn.GRU(input_size, hidden_size, batch_first=True) for _ in range(4)])\n        self.ln_list = nn.ModuleList([nn.LayerNorm(hidden_size) for _ in range(4)])\n                    \n        self.fc = nn.Linear(4*hidden_size, output_size)\n        self.softmax = nn.Softmax(dim=1)\n\n    def forward(self, X):\n        X = X.permute(3,0,1,2)\n        # Initialize a list to store hidden states\n        hidden_states = []\n\n        # Iterate over each input\n        for i, x in enumerate(X):\n            out, h = self.gru_list[i](x)\n            # Apply layer normalization\n            out = self.ln_list[i](out)\n            # Append the last hidden state to the list\n            hidden_states.append(h[-1])\n\n        # Concatenate the hidden states along the last axis\n        hidden = torch.cat(hidden_states, dim=1)\n        \n        output = self.fc(hidden)\n        output = self.softmax(output)\n        return output","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:33.067388Z","iopub.execute_input":"2024-02-22T11:12:33.068664Z","iopub.status.idle":"2024-02-22T11:12:33.075743Z","shell.execute_reply.started":"2024-02-22T11:12:33.068629Z","shell.execute_reply":"2024-02-22T11:12:33.074672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from model import GRUClassifier\n\nmodel = GRUClassifier().to(device)\nfor X,y in data_loaders[0]:\n    X = X.to(device)\n    y = y.to(device)\n    labels = model(X)\n    print(labels.shape)\n    break\n    \nmodel.parameters","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:38.094861Z","iopub.execute_input":"2024-02-22T11:12:38.095252Z","iopub.status.idle":"2024-02-22T11:12:39.117410Z","shell.execute_reply.started":"2024-02-22T11:12:38.095215Z","shell.execute_reply":"2024-02-22T11:12:39.116473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(model, pred_loader):\n    preds = []\n    \n    for X in pred_loader:\n        X = X.to(device)\n        print(X.shape)\n        pred = model(X)[0]\n        preds.append(pred)\n        \n    return preds\n\n#test\npreds = predict(model, sub_dataloader)\nprint(preds)","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:46.390819Z","iopub.execute_input":"2024-02-22T11:12:46.391217Z","iopub.status.idle":"2024-02-22T11:12:46.561504Z","shell.execute_reply.started":"2024-02-22T11:12:46.391169Z","shell.execute_reply":"2024-02-22T11:12:46.560503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader\nfrom tqdm.auto import tqdm\n\ndef get_accuracy(model, test_loader):\n    device = next(model.parameters()).device \n    \n    score = 0\n    model.eval()\n    for testX, testY in tqdm(test_loader):\n        testX = testX.to(device)\n        testY = testY.to(device)\n        \n        preds = model(testX)\n        _, pred_labels = torch.max(preds, dim=1)\n        _, target_labels = torch.max(testY, dim=1)\n        score += torch.sum(pred_labels == target_labels)\n\n    return score/len(test_loader.dataset)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = data_loaders[9]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom tqdm.auto import tqdm\n\nmodel = GRUClassifier()\nmodel = model.to(device)\n\n# Define optimizer\noptimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)\n\n# Define loss function\ncriterion = nn.KLDivLoss(reduction='batchmean')\n\n# Training loop\ndef train_model(model, train_loader, optimizer, criterion):\n    model.train()\n    running_loss = 0.0\n    for inputs, targets in tqdm(train_loader):\n        optimizer.zero_grad()\n        inputs = inputs.to(device)\n        targets = targets.to(device)\n        outputs = model(inputs)\n        loss = criterion(outputs, targets)\n        loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), max_norm=2.0)\n        optimizer.step()\n        running_loss += loss.item() * inputs.size(0)\n        \n    return running_loss / len(train_loader.dataset)\n\n# Validation loop\ndef validate_model(model, valid_loader, criterion):\n    model.eval()\n    running_loss = 0.0\n    with torch.no_grad():\n        for inputs, targets in tqdm(valid_loader):\n            inputs = inputs.to(device)\n            outputs = model(inputs)\n            loss = criterion(outputs, targets)\n            running_loss += loss.item() * inputs.size(0)\n    return running_loss / len(valid_loader.dataset)\n\naccuracy = get_accuracy(model, test_loader)\nprint(f\"base accuracy: {accuracy}\")\n\n# Assuming you have train_loader and valid_loader DataLoader objects for training and validation data\nfor epoch in tqdm(range(NUM_EPOCHS)):\n    \n    loader_idx = epoch%10\n    train_loader = data_loaders[loader_idx]\n    \n    train_loss = train_model(model, train_loader, optimizer, criterion)\n    print(f\"Epoch {epoch+1}/{NUM_EPOCHS}, Train Loss: {train_loss}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"accuracy = get_accuracy(model, test_loader)\nprint(f\"Accuracy after {NUM_EPOCHS} epochs: {accuracy}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vote_cols = ['seizure_vote','lpd_vote','gpd_vote','lrda_vote','grda_vote','other_vote']\n\npreds = predict(model, sub_dataloader)[0].cpu().detach().numpy()\nsubmission = pd.read_csv(SUBMISSION_FILE)","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:12:57.412082Z","iopub.execute_input":"2024-02-22T11:12:57.412471Z","iopub.status.idle":"2024-02-22T11:12:57.428058Z","shell.execute_reply.started":"2024-02-22T11:12:57.412443Z","shell.execute_reply":"2024-02-22T11:12:57.427088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:13:25.218814Z","iopub.execute_input":"2024-02-22T11:13:25.219710Z","iopub.status.idle":"2024-02-22T11:13:25.226105Z","shell.execute_reply.started":"2024-02-22T11:13:25.219675Z","shell.execute_reply":"2024-02-22T11:13:25.224977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vote_cols = submission.columns[1:]\n\nfor i,col in enumerate(vote_cols):\n    submission.loc[0,col] = preds[i]","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:13:34.309254Z","iopub.execute_input":"2024-02-22T11:13:34.309642Z","iopub.status.idle":"2024-02-22T11:13:34.316041Z","shell.execute_reply.started":"2024-02-22T11:13:34.309610Z","shell.execute_reply":"2024-02-22T11:13:34.315105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !rm -rf /kaggle/working/*","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:13:37.809854Z","iopub.execute_input":"2024-02-22T11:13:37.810230Z","iopub.status.idle":"2024-02-22T11:13:38.774247Z","shell.execute_reply.started":"2024-02-22T11:13:37.810198Z","shell.execute_reply":"2024-02-22T11:13:38.773132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:13:41.462979Z","iopub.execute_input":"2024-02-22T11:13:41.463557Z","iopub.status.idle":"2024-02-22T11:13:41.472048Z","shell.execute_reply.started":"2024-02-22T11:13:41.463514Z","shell.execute_reply":"2024-02-22T11:13:41.471113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'gru_v2.pth')","metadata":{"execution":{"iopub.status.busy":"2024-02-22T11:13:47.471841Z","iopub.execute_input":"2024-02-22T11:13:47.472496Z","iopub.status.idle":"2024-02-22T11:13:47.480968Z","shell.execute_reply.started":"2024-02-22T11:13:47.472465Z","shell.execute_reply":"2024-02-22T11:13:47.480037Z"},"trusted":true},"execution_count":null,"outputs":[]}]}