{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"none","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"A custom modification of GATConv source code available at\n\nhttps://pytorch-geometric.readthedocs.io/en/latest/_modules/torch_geometric/nn/conv/gat_conv.html#GATConv)\n\nto apply a double convolution encoding to the eegs nodes instead of the original plane linear transformation to a general set node features.\n\nI think this kind of processing makes more sense to the eeg specific case. We'll see.","metadata":{}},{"cell_type":"code","source":"import typing\nfrom typing import Optional, Tuple, Union\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import Tensor\nfrom torch.nn import Parameter\n\ntry:\n    from torch_geometric.nn.conv import MessagePassing\n    from torch_geometric.nn.dense.linear import Linear\n    from torch_geometric.nn.inits import glorot, zeros\n    from torch_geometric.typing import (\n        Adj,\n        NoneType,\n        OptPairTensor,\n        OptTensor,\n        Size,\n        SparseTensor,\n        torch_sparse,\n    )\n    from torch_geometric.utils import (\n        add_self_loops,\n        is_torch_sparse_tensor,\n        remove_self_loops,\n        softmax,\n    )\n    from torch_geometric.utils.sparse import set_sparse_value\n    \nexcept:\n    !pip install torch_geometric\n    from torch_geometric.nn.conv import MessagePassing\n    from torch_geometric.nn.dense.linear import Linear\n    from torch_geometric.nn.inits import glorot, zeros\n    from torch_geometric.typing import (\n        Adj,\n        NoneType,\n        OptPairTensor,\n        OptTensor,\n        Size,\n        SparseTensor,\n        torch_sparse,\n    )\n    from torch_geometric.utils import (\n        add_self_loops,\n        is_torch_sparse_tensor,\n        remove_self_loops,\n        softmax,\n    )\n    from torch_geometric.utils.sparse import set_sparse_value\n\nif typing.TYPE_CHECKING:\n    from typing import overload\nelse:\n    from torch.jit import _overload_method as overload","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://pytorch-geometric.readthedocs.io/en/latest/_modules/torch_geometric/nn/conv/gat_conv.html#GATConv\n# Modified GATconv to apply double conv encoding to nodes instead of linear\nclass myGATConv(MessagePassing):\n    def __init__(\n        self,\n        width: int,\n        in_channels: int,\n        mid_channels: int,\n        out_channels: int,\n        kernel_size: int,\n        heads: int = 1,\n        concat: bool = True,\n        negative_slope: float = 0.2,\n        dropout: float = 0.0,\n        add_self_loops: bool = True,\n        edge_dim: Optional[int] = None,\n        fill_value: Union[float, Tensor, str] = 'mean',\n        bias: bool = True,\n        **kwargs,\n    ):\n        kwargs.setdefault('aggr', 'add')\n        super().__init__(node_dim=0, **kwargs)\n\n        self.W = width\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.heads = heads\n        self.concat = concat\n        self.negative_slope = negative_slope\n        self.dropout = dropout\n        self.add_self_loops = add_self_loops\n        self.edge_dim = edge_dim\n        self.fill_value = fill_value\n        # Feel free to change the encoding\n        # I've chosen the classical double conv + pooling\n        self.ENCODER = nn.Sequential(\n                nn.Conv1d(in_channels, mid_channels, kernel_size, padding='same'),\n                nn.BatchNorm1d(mid_channels),\n                nn.ReLU(inplace=True),\n                nn.Conv1d(mid_channels, heads * out_channels, kernel_size, padding='same'),\n                nn.BatchNorm1d(heads * out_channels),\n                nn.ReLU(inplace=True),\n                nn.MaxPool1d(2)\n        )\n        # The learnable parameters to compute attention coefficients:\n        self.att_src = Parameter(torch.empty(1, heads, out_channels * width//2))\n        self.att_dst = Parameter(torch.empty(1, heads, out_channels * width//2))\n\n        self.lin_edge = None\n        self.register_parameter('att_edge', None)\n\n        if bias and concat:\n            self.bias = Parameter(torch.empty(heads * out_channels * width//2))\n        elif bias and not concat:\n            self.bias = Parameter(torch.empty(out_channels * width//2))\n        else:\n            self.register_parameter('bias', None)\n\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        super().reset_parameters()\n        if self.ENCODER is not None:\n            for L in self.ENCODER:\n                try:\n                    L.reset_parameters()\n                except:\n                    None\n        if self.lin_edge is not None:\n            self.lin_edge.reset_parameters()\n        glorot(self.att_src)\n        glorot(self.att_dst)\n        glorot(self.att_edge)\n        zeros(self.bias)\n\n\n    @overload\n    def forward(\n        self,\n        x: Union[Tensor, OptPairTensor],\n        edge_index: Adj,\n        edge_attr: OptTensor = None,\n        size: Size = None,\n        return_attention_weights: NoneType = None,\n    ) -> Tensor:\n        pass\n\n    @overload\n    def forward(  # noqa: F811\n        self,\n        x: Union[Tensor, OptPairTensor],\n        edge_index: Tensor,\n        edge_attr: OptTensor = None,\n        size: Size = None,\n        return_attention_weights: bool = None,\n    ) -> Tuple[Tensor, Tuple[Tensor, Tensor]]:\n        pass\n\n    @overload\n    def forward(  # noqa: F811\n        self,\n        x: Union[Tensor, OptPairTensor],\n        edge_index: SparseTensor,\n        edge_attr: OptTensor = None,\n        size: Size = None,\n        return_attention_weights: bool = None,\n    ) -> Tuple[Tensor, SparseTensor]:\n        pass\n\n    def forward(  # noqa: F811\n        self,\n        x: Union[Tensor, OptPairTensor],\n        edge_index: Adj,\n        edge_attr: OptTensor = None,\n        size: Size = None,\n        return_attention_weights: Optional[bool] = None,\n    ) -> Union[\n            Tensor,\n            Tuple[Tensor, Tuple[Tensor, Tensor]],\n            Tuple[Tensor, SparseTensor],\n    ]:\n        H, F = self.heads, self.out_channels * self.W//2\n        # We first transform the input node features.\n        # BATCHxNODESxCHANNELSxWIDTH\n        assert x.dim() == 4, \"BxNxCxW\"\n        B,N,C,W = x.shape\n        x_src = x_dst = self.ENCODER(x.view(B*N,C,W)).view(-1, H, F)\n        x = (x_src, x_dst)\n\n        # Next, we compute node-level attention coefficients, both for source\n        # and target nodes (if present):\n        alpha_src = (x_src * self.att_src).sum(-1)\n        alpha_dst = (x_dst * self.att_dst).sum(-1)\n        alpha = (alpha_src, alpha_dst)\n\n        if self.add_self_loops:\n            if isinstance(edge_index, Tensor):\n                # We only want to add self-loops for nodes that appear both as\n                # source and target nodes:\n                edge_index, edge_attr = remove_self_loops(\n                    edge_index, edge_attr)\n                edge_index, edge_attr = add_self_loops(\n                    edge_index, edge_attr, fill_value=self.fill_value,\n                    num_nodes=B*N)\n            elif isinstance(edge_index, SparseTensor):\n                if self.edge_dim is None:\n                    edge_index = torch_sparse.set_diag(edge_index)\n                else:\n                    raise NotImplementedError(\n                        \"The usage of 'edge_attr' and 'add_self_loops' \"\n                        \"simultaneously is currently not yet supported for \"\n                        \"'edge_index' in a 'SparseTensor' form\")\n\n        # edge_updater_type: (alpha: OptPairTensor, edge_attr: OptTensor)\n        alpha = self.edge_updater(edge_index, alpha=alpha, edge_attr=edge_attr,\n                                  size=size)\n\n        # propagate_type: (x: OptPairTensor, alpha: Tensor)\n        out = self.propagate(edge_index, x=x, alpha=alpha, size=size)\n\n        if self.concat:\n            out = out.view(-1, self.heads * F)\n        else:\n            out = out.mean(dim=1)\n\n        if self.bias is not None:\n            out = out + self.bias\n            \n        out = out.view(B,N,self.out_channels,self.W//2)\n\n        if isinstance(return_attention_weights, bool):\n            if isinstance(edge_index, Tensor):\n                if is_torch_sparse_tensor(edge_index):\n                    # TODO TorchScript requires to return a tuple\n                    adj = set_sparse_value(edge_index, alpha)\n                    return out, (adj, alpha)\n                else:\n                    return out, (edge_index, alpha)\n            elif isinstance(edge_index, SparseTensor):\n                return out, edge_index.set_value(alpha, layout='coo')\n        else:\n            return out\n\n\n    def edge_update(self, alpha_j: Tensor, alpha_i: OptTensor,\n                    edge_attr: OptTensor, index: Tensor, ptr: OptTensor,\n                    dim_size: Optional[int]) -> Tensor:\n        # Given edge-level attention coefficients for source and target nodes,\n        # we simply need to sum them up to \"emulate\" concatenation:\n        alpha = alpha_j if alpha_i is None else alpha_j + alpha_i\n        if index.numel() == 0:\n            return alpha\n        if edge_attr is not None and self.lin_edge is not None:\n            if edge_attr.dim() == 1:\n                edge_attr = edge_attr.view(-1, 1)\n            edge_attr = self.lin_edge(edge_attr)\n            edge_attr = edge_attr.view(-1, self.heads, self.out_channels)\n            alpha_edge = (edge_attr * self.att_edge).sum(dim=-1)\n            alpha = alpha + alpha_edge\n\n        alpha = F.leaky_relu(alpha, self.negative_slope)\n        alpha = softmax(alpha, index, ptr, dim_size)\n        alpha = F.dropout(alpha, p=self.dropout, training=self.training)\n        return alpha\n\n    def message(self, x_j: Tensor, alpha: Tensor) -> Tensor:\n        return alpha.unsqueeze(-1) * x_j\n\n    def __repr__(self) -> str:\n        return (f'{self.__class__.__name__}({self.in_channels}, '\n                f'{self.out_channels}, heads={self.heads})')","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:50.665916Z","iopub.execute_input":"2024-02-22T10:59:50.666892Z","iopub.status.idle":"2024-02-22T10:59:50.702127Z","shell.execute_reply.started":"2024-02-22T10:59:50.666842Z","shell.execute_reply":"2024-02-22T10:59:50.701115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndf = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/train.csv')\ndf.head(5)\nvotes = [c for c in df.columns if '_vote' in c]","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:50.706708Z","iopub.execute_input":"2024-02-22T10:59:50.707341Z","iopub.status.idle":"2024-02-22T10:59:51.268726Z","shell.execute_reply.started":"2024-02-22T10:59:50.707294Z","shell.execute_reply":"2024-02-22T10:59:51.267551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"electrodes = [\n    'Fp1',\n    'F3',\n    'C3',\n    'P3',\n    'F7',\n    'T3',\n    'T5',\n    'O1',\n    'Fz',\n    'Cz',\n    'Pz',\n    'Fp2',\n    'F4',\n    'C4',\n    'P4',\n    'F8',\n    'T4',\n    'T6',\n    'O2',\n    'EKG'\n]","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:51.270247Z","iopub.execute_input":"2024-02-22T10:59:51.270604Z","iopub.status.idle":"2024-02-22T10:59:51.279566Z","shell.execute_reply.started":"2024-02-22T10:59:51.270576Z","shell.execute_reply":"2024-02-22T10:59:51.275957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:51.281121Z","iopub.execute_input":"2024-02-22T10:59:51.281790Z","iopub.status.idle":"2024-02-22T10:59:51.292725Z","shell.execute_reply.started":"2024-02-22T10:59:51.281751Z","shell.execute_reply":"2024-02-22T10:59:51.291401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nclass eeg_Dataset(torch.utils.data.Dataset):\n    '''\n    '''  \n    def __init__(self, df, W=2048, VALID=False):\n        self.START = (10000 - W)//2\n        self.END = self.START + W\n        self.VALID = VALID\n        if VALID:\n            self.data = df\n        else:\n            self.data = list(df.groupby(['patient_id','eeg_id','expert_consensus']))\n   \n    def __len__(self):\n        return len(self.data)\n    \n    def __getdata__(self):\n        return self.data\n        \n    def __getitem__(self, idx):\n        PATH = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs/'\n\n        if self.VALID:\n            row = self.data.iloc[idx]\n        else:\n            df = self.data[idx][1]\n            row = df.iloc[np.random.randint(len(df))]\n\n        eeg = pd.read_parquet(f'{PATH}{int(row.eeg_id)}.parquet')\n        eeg_offset = int( row.eeg_label_offset_seconds )\n        eeg = eeg.iloc[eeg_offset*200+self.START:eeg_offset*200+self.END]\n        eeg = eeg[electrodes].values.T\n        eeg[np.isnan(eeg)] = np.nanmean(eeg)\n        eeg = torch.from_numpy(eeg).to(device)\n\n        labels = row[votes].values\n        labels /= labels.sum()\n        labels = torch.from_numpy(labels.astype(np.float32)).to(device)\n\n        return [eeg.unsqueeze(-2),labels]","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:51.294252Z","iopub.execute_input":"2024-02-22T10:59:51.294779Z","iopub.status.idle":"2024-02-22T10:59:51.308080Z","shell.execute_reply.started":"2024-02-22T10:59:51.294742Z","shell.execute_reply":"2024-02-22T10:59:51.306695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = eeg_Dataset(df)","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:51.309846Z","iopub.execute_input":"2024-02-22T10:59:51.310495Z","iopub.status.idle":"2024-02-22T10:59:53.216882Z","shell.execute_reply.started":"2024-02-22T10:59:51.310456Z","shell.execute_reply":"2024-02-22T10:59:53.215727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:53.218152Z","iopub.execute_input":"2024-02-22T10:59:53.218541Z","iopub.status.idle":"2024-02-22T10:59:53.224773Z","shell.execute_reply.started":"2024-02-22T10:59:53.218510Z","shell.execute_reply":"2024-02-22T10:59:53.223163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BS = 256","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:53.226207Z","iopub.execute_input":"2024-02-22T10:59:53.226554Z","iopub.status.idle":"2024-02-22T10:59:53.233996Z","shell.execute_reply.started":"2024-02-22T10:59:53.226524Z","shell.execute_reply":"2024-02-22T10:59:53.233099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"I = []\nJ = []\nN = len(electrodes)\nfor k in range(BS):\n    K = N*list(range(N*k,N*k+N))\n    I += K\n    J += sorted(K)\n    \nedges = torch.tensor([I,J])\nedges.shape","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:53.235247Z","iopub.execute_input":"2024-02-22T10:59:53.236119Z","iopub.status.idle":"2024-02-22T10:59:53.295874Z","shell.execute_reply.started":"2024-02-22T10:59:53.236084Z","shell.execute_reply":"2024-02-22T10:59:53.294723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dl = DataLoader(ds,\n                batch_size=BS,\n                shuffle=True,\n                drop_last=True)","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:53.297362Z","iopub.execute_input":"2024-02-22T10:59:53.298055Z","iopub.status.idle":"2024-02-22T10:59:53.307039Z","shell.execute_reply.started":"2024-02-22T10:59:53.298018Z","shell.execute_reply":"2024-02-22T10:59:53.305497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for b in dl:\n    print('input shape:  ',b[0].shape)\n    x = myGATConv(\n        2048,# WIDTH\n        1,   # IN_CHANNELS\n        16,  # MID_CHANNELS\n        16,  # OUT_CHANNELS\n        3    # KERNEL_SIZE\n    ).to(device)(b[0],edges)\n    print('output shape: ',x.shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2024-02-22T10:59:53.310261Z","iopub.execute_input":"2024-02-22T10:59:53.310685Z","iopub.status.idle":"2024-02-22T11:00:10.539590Z","shell.execute_reply.started":"2024-02-22T10:59:53.310654Z","shell.execute_reply":"2024-02-22T11:00:10.538139Z"},"trusted":true},"execution_count":null,"outputs":[]}]}