{"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"}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# GNNs meets HMS","metadata":{}},{"cell_type":"code","source":"import os\n\nimport numpy as np\n\n#!pip install pyarrow\n#!pip install fastparquet\nimport pandas as pd\n\nimport torch\nfrom torch.utils.data import Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n!pip install torch_geometric\nimport torch.nn.functional as F\nfrom torch_geometric.loader import DataLoader\nfrom torch_geometric.nn import SGConv, ChebConv, GCNConv, GATConv\nfrom torch_geometric.nn.pool import global_mean_pool\n\n!pip install torcheeg\nfrom torcheeg.transforms import Compose, BandDifferentialEntropy, MeanStdNormalize\nfrom torcheeg.transforms.pyg import ToDynamicG\n\nfrom sklearn.model_selection import KFold\nfrom sklearn.preprocessing import StandardScaler\n\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:03.938085Z","iopub.execute_input":"2024-03-07T21:49:03.938731Z","iopub.status.idle":"2024-03-07T21:49:49.982934Z","shell.execute_reply.started":"2024-03-07T21:49:03.938698Z","shell.execute_reply":"2024-03-07T21:49:49.981975Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define config","metadata":{}},{"cell_type":"code","source":"torch.set_default_tensor_type('torch.FloatTensor')\nconfig = {\n    'batch_size': 2048,\n    'device': \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    'n_fold': 3,\n    'seed': 2**20-1,\n    'epochs': 10,\n    'in_channels': 10_000, #4\n    'hidden_channels': 32,\n    'num_conv_layers': 4,\n    'num_classes': 6,\n    'lr': 1e-2,\n    'weight_decay': 5e-4\n}","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:57:25.519642Z","iopub.execute_input":"2024-03-07T21:57:25.520001Z","iopub.status.idle":"2024-03-07T21:57:25.528117Z","shell.execute_reply.started":"2024-03-07T21:57:25.519974Z","shell.execute_reply":"2024-03-07T21:57:25.527124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get data","metadata":{}},{"cell_type":"code","source":"EEG_PATH_TRAIN = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/'\n\ndf = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\ndf","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.023132Z","iopub.execute_input":"2024-03-07T21:49:50.023827Z","iopub.status.idle":"2024-03-07T21:49:50.305242Z","shell.execute_reply.started":"2024-03-07T21:49:50.023799Z","shell.execute_reply":"2024-03-07T21:49:50.304324Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EEG_IDS = df.eeg_id.unique()\n\nTARGETS = df.columns[-6:]\nTARS = {'Seizure':0, 'LPD':1, 'GPD':2, 'LRDA':3, 'GRDA':4, 'Other':5}\nTARS_INV = {x:y for y,x in TARS.items()}\n\ntrain = df.groupby('eeg_id')[['patient_id']].agg('first')\n\ntmp = df.groupby('eeg_id')[TARGETS].agg('sum')\nfor t in TARGETS:\n    train[t] = tmp[t].values\n    \ny_data = train[TARGETS].values\ny_data = y_data / y_data.sum(axis=1,keepdims=True)\ntrain[TARGETS] = y_data\n\ntmp = df.groupby('eeg_id')[['expert_consensus']].agg('first')\ntrain['target'] = tmp\n\ntrain = train.reset_index()\ntrain = train.loc[train.eeg_id.isin(EEG_IDS)]\nprint('Train Data with unique eeg_id shape:', train.shape )","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.307441Z","iopub.execute_input":"2024-03-07T21:49:50.307750Z","iopub.status.idle":"2024-03-07T21:49:50.371258Z","shell.execute_reply.started":"2024-03-07T21:49:50.307726Z","shell.execute_reply":"2024-03-07T21:49:50.370284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eeg_from_parquet(parquet_path, display=False):\n    \n    # EXTRACT MIDDLE 50 SECONDS\n    eeg = pd.read_parquet(parquet_path)\n    rows = len(eeg)\n    offset = (rows-10_000)//2\n    eeg = eeg.iloc[offset:offset+10_000]\n    \n    if display: \n        plt.figure(figsize=(10,5))\n        offset = 0\n    \n    # CONVERT TO NUMPY\n    data = np.zeros((10_000,len(eeg.columns)))\n    for j,col in enumerate(eeg.columns):\n\n        x = eeg[col].values.astype('float32')\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        data[:,j] = x\n        \n        if display: \n            if j!=0: offset += x.max()\n            plt.plot(range(10_000),x-offset,label=col)\n            offset -= x.min()\n            \n    if display:\n        plt.legend()\n        name = parquet_path.split('/')[-1]\n        name = name.split('.')[0]\n        plt.title(f'EEG {name}',size=16)\n        plt.show()\n        \n    return data","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.372511Z","iopub.execute_input":"2024-03-07T21:49:50.372841Z","iopub.status.idle":"2024-03-07T21:49:50.382478Z","shell.execute_reply.started":"2024-03-07T21:49:50.372814Z","shell.execute_reply":"2024-03-07T21:49:50.381583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Example of one-row-analysis","metadata":{}},{"cell_type":"code","source":"def get_one_row_data(row_id, mode='train'):\n    row = train.iloc[row_id]\n    \n    path = EEG_PATH_TRAIN if mode == 'train' else EEG_PATH_TEST\n    \n    \n    return eeg_from_parquet(f'{path}{row.eeg_id}.parquet')\n    eeg = pd.read_parquet(f'{path}{row.eeg_id}.parquet')\n    eeg_offset = int(row.eeg_label_offset_seconds)\n    eeg = eeg.iloc[eeg_offset*200:(eeg_offset+50)*200]\n    \n    ekg = torch.Tensor(eeg['EKG'].to_numpy())\n    eeg = torch.Tensor(eeg.drop(columns='EKG').to_numpy())\n    \n    return (eeg, ekg)","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.383837Z","iopub.execute_input":"2024-03-07T21:49:50.384163Z","iopub.status.idle":"2024-03-07T21:49:50.395808Z","shell.execute_reply.started":"2024-03-07T21:49:50.384133Z","shell.execute_reply":"2024-03-07T21:49:50.394970Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"row_id = 0#5859\nnp.isnan(get_one_row_data(row_id)).any()","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.397143Z","iopub.execute_input":"2024-03-07T21:49:50.398061Z","iopub.status.idle":"2024-03-07T21:49:50.598135Z","shell.execute_reply.started":"2024-03-07T21:49:50.398027Z","shell.execute_reply":"2024-03-07T21:49:50.597199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Let's make a graph from each row data","metadata":{}},{"cell_type":"code","source":"class EEGDataset(Dataset):\n    def __init__(self, df, scaler=StandardScaler(), transform=ToDynamicG(\n        edge_func='absolute_pearson_correlation_coefficient', threshold=0.5, binary=True)):\n        \n        self.df = df\n        self.path = EEG_PATH_TRAIN\n        self.scaler = scaler\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n\n        eeg = eeg_from_parquet(f'{self.path}{row.eeg_id}.parquet')\n        #eeg = self.scaler.fit_transform(eeg)\n        eeg = torch.Tensor(eeg)\n        graph = self.transform(eeg=eeg.T)['eeg']\n        \n        y = np.array(row[-7:-1].values, 'float32').reshape(1,-1)\n        y = y / y.sum(axis=1, keepdims=True)\n        y = torch.Tensor(y)\n        graph.update({'y': y})\n\n        return graph","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.599370Z","iopub.execute_input":"2024-03-07T21:49:50.599690Z","iopub.status.idle":"2024-03-07T21:49:50.608210Z","shell.execute_reply.started":"2024-03-07T21:49:50.599666Z","shell.execute_reply":"2024-03-07T21:49:50.607269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms = Compose([\n    #BandDifferentialEntropy(),\n    MeanStdNormalize(),\n    ToDynamicG(edge_func='absolute_pearson_correlation_coefficient', threshold=0.5, binary=True)\n])","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.609390Z","iopub.execute_input":"2024-03-07T21:49:50.609715Z","iopub.status.idle":"2024-03-07T21:49:50.622502Z","shell.execute_reply.started":"2024-03-07T21:49:50.609691Z","shell.execute_reply":"2024-03-07T21:49:50.621729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5859","metadata":{}},{"cell_type":"code","source":"train.iloc[5859]","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.625111Z","iopub.execute_input":"2024-03-07T21:49:50.625379Z","iopub.status.idle":"2024-03-07T21:49:50.636869Z","shell.execute_reply.started":"2024-03-07T21:49:50.625357Z","shell.execute_reply":"2024-03-07T21:49:50.635983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = get_one_row_data(5859)\nde = BandDifferentialEntropy()\nde_x = de(eeg=x.T)['eeg']\nde_x_torch = torch.Tensor(de_x)\nt = ToDynamicG(edge_func='absolute_pearson_correlation_coefficient', threshold=0.5, binary=True)\nt(eeg=de_x_torch)","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.637936Z","iopub.execute_input":"2024-03-07T21:49:50.638554Z","iopub.status.idle":"2024-03-07T21:49:50.820356Z","shell.execute_reply.started":"2024-03-07T21:49:50.638521Z","shell.execute_reply":"2024-03-07T21:49:50.819430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = EEGDataset(train, transform=transforms)\ndataset.__getitem__(5859).x","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:49:50.821551Z","iopub.execute_input":"2024-03-07T21:49:50.822085Z","iopub.status.idle":"2024-03-07T21:49:50.953286Z","shell.execute_reply.started":"2024-03-07T21:49:50.822052Z","shell.execute_reply":"2024-03-07T21:49:50.952311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create GNN","metadata":{}},{"cell_type":"code","source":"class GNN(torch.nn.Module):\n    def __init__(self, in_channels=4, num_conv_layers=config['num_conv_layers'],\n                 hid_channels=config['hidden_channels'],\n                 num_classes=config['num_classes']):\n        super().__init__()\n        self.embedding = torch.nn.Sequential(\n            torch.nn.Linear(10_000, 4096),\n            torch.nn.ReLU(),\n            torch.nn.Linear(4096, 1024),\n            torch.nn.ReLU(),\n            torch.nn.Linear(1024, 32),\n        )\n        #self.conv1 = GATConv(in_channels, hid_channels)\n        self.convs = torch.nn.ModuleList()\n        for _ in range(num_conv_layers):\n            self.convs.append(ChebConv(hid_channels, hid_channels, K=2))\n        self.lin1 = torch.nn.Linear(hid_channels, hid_channels)\n        self.lin2 = torch.nn.Linear(hid_channels, num_classes)\n\n    def reset_parameters(self):\n        self.embedding.apply(lambda x: x.reset_parameters() if isinstance(x, torch.nn.Linear) else x)\n        #self.conv1.reset_parameters()\n        for conv in self.convs:\n            conv.reset_parameters()\n        self.lin1.reset_parameters()\n        self.lin2.reset_parameters()\n\n    def forward(self, data):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        \n        #x = F.relu(self.conv1(x, edge_index))\n        x = F.relu(self.embedding(x))\n        for conv in self.convs:\n            x = F.relu(conv(x, edge_index))\n        x = global_mean_pool(x, batch)\n        x = F.relu(self.lin1(x))\n        x = F.dropout(x, p=0.5, training=self.training)\n        x = self.lin2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:54:19.587548Z","iopub.execute_input":"2024-03-07T21:54:19.587917Z","iopub.status.idle":"2024-03-07T21:54:19.599473Z","shell.execute_reply.started":"2024-03-07T21:54:19.587890Z","shell.execute_reply":"2024-03-07T21:54:19.598399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"def criterion(logit, target):\n    log_prob = F.log_softmax(logit, dim=1)\n    return F.kl_div(log_prob, target, reduction=\"batchmean\")\n\ndef KL_loss(p,q):\n    epsilon=10**(-15)\n    p=torch.clip(p,epsilon,1-epsilon)\n    q = nn.functional.log_softmax(q,dim=1)\n    return torch.mean(torch.sum(p*(torch.log(p)-q),dim=1))\n\ndef compute_loss(model, data_loader):\n    model.eval()\n    l_loss = []\n    with torch.no_grad():\n        for data in data_loader:\n            data.to(config['device'])\n            y_pred = model(data)\n            loss = criterion(y_pred, data.y)\n            l_loss.append(loss.item())\n    return np.mean(l_loss) ","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:54:19.975745Z","iopub.execute_input":"2024-03-07T21:54:19.976102Z","iopub.status.idle":"2024-03-07T21:54:19.984246Z","shell.execute_reply.started":"2024-03-07T21:54:19.976074Z","shell.execute_reply":"2024-03-07T21:54:19.983013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nkf = KFold(n_splits=config['n_fold'], shuffle=True, random_state=config['seed'])\n\nl_best_loss = []\n\nfor fold, (iloc_train, iloc_valid) in enumerate(kf.split(train)):\n    print(f\"Fold {fold}:\")\n    \n    train_ds = EEGDataset(df=train.iloc[iloc_train], transform=transforms)\n    valid_ds = EEGDataset(df=train.iloc[iloc_valid], transform=transforms)\n    train_loader = DataLoader(dataset=train_ds, shuffle=True, batch_size=config['batch_size'], \n                              num_workers=os.cpu_count(), drop_last=True)\n    valid_loader = DataLoader(dataset=valid_ds, batch_size=config['batch_size'], \n                              num_workers=os.cpu_count())\n    \n    model = GNN(in_channels=10_000).to(config['device'])#EEG_Classifier(config['hidden_channels']).to(config['device'])\n    optimizer = torch.optim.Adam(model.parameters(), lr=config['lr'], \n                                 weight_decay=config['weight_decay'])\n    scheduler = CosineAnnealingLR(optimizer=optimizer, T_max=config['epochs'])\n    \n    best_loss = float(\"inf\")\n    history = []\n    \n    for epoch in range(config['epochs']):\n        model.train()\n        l_loss = []\n        for data in train_loader:\n            data.to(config['device'])\n            optimizer.zero_grad()\n            out = model(data)\n            loss = criterion(out, data.y)\n            l_loss.append(loss.item())\n\n            print(f\"epoch={epoch}\\t loss={loss}\")\n            loss.backward()\n            #optimizer.step()\n        train_loss = np.mean(l_loss)\n        valid_loss = compute_loss(model, valid_loader)\n        history.append((epoch, train_loss, valid_loss))\n        print(f\"Epoch {epoch}\")\n        print(f\"Train Loss: {train_loss:>10.6f}, Valid Loss: {valid_loss:>10.6}\")\n\n        if valid_loss < best_loss:\n                print(f\"Loss improves from {best_loss:>10.6f} to {valid_loss:>10.6}\")\n                torch.save(model.state_dict(), f\"{'basic_GNN'}__{fold}.pt\")\n                best_loss = valid_loss\n    print(f\"\\nBest loss Model training with {best_loss}\\n\")\n    l_best_loss.append(best_loss)\n    \n    history = pd.DataFrame(history, columns=[\"epoch\", \"loss\", \"val_loss\"]).set_index(\"epoch\")\n    history.plot(subplots=True, layout=(1, 2), sharey=\"row\", figsize=(14, 6))\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-03-07T21:57:28.553805Z","iopub.execute_input":"2024-03-07T21:57:28.554462Z"},"trusted":true},"execution_count":null,"outputs":[]}]}