{"metadata":{"accelerator":"GPU","colab":{"gpuType":"T4","machine_shape":"hm","provenance":[{"file_id":"1hEYvoxNLgFcjmkYBO0iQJup7FJjKyhGN","timestamp":1720312097607},{"file_id":"1fMLs4yH1L6lGxQvPsoZt5fOkCOfk1gmb","timestamp":1720304723557}]},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":8753774,"sourceType":"datasetVersion","datasetId":5258539},{"sourceId":8772913,"sourceType":"datasetVersion","datasetId":5272293}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.12"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport pandas as pd\nimport sklearn\nimport matplotlib.pyplot as plt\nimport gc\nimport tqdm\n\nimport kora.install.rdkit\nfrom rdkit import Chem\nfrom rdkit.Chem import AllChem\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nfrom torch.utils.data import Dataset, DataLoader\nfrom fastai.vision.all import *\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"id":"LrJvGyoEqeAq"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"executionInfo":{"elapsed":6,"status":"ok","timestamp":1720312179424,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"2rit3sK3rNHU","outputId":"ed803bfa-ed58-4fa6-e8de-f8fcf231e004"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BS = 512\nLR = 2.5e-5\nEPOCHS = 1\nHEADS = 2\nMASK_NODES = .2\nMASK_GC = .2\nMASK_EDGES = .2\nDROPOUT = 0\nLmax = 75\nNBH = 50\nFOLDS = [0,1,2,3,4]\nCV = 5\nff = 1\nT = 'binds_HSA'","metadata":{"id":"R3ikzFC-qeAt"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Shuffled from https://www.kaggle.com/datasets/shlomoron/belka-shrunken-train-set\ndtypes = {\n    'buildingblock1_smiles': np.int16,\n    'buildingblock2_smiles': np.int16,\n    'buildingblock3_smiles': np.int16,\n    'binds_BRD4':np.byte,\n    'binds_HSA':np.byte,\n    'binds_sEH':np.byte\n}\ntrain = pd.read_csv(\"C:/Users/Angel/kaggle/train.csv\", dtype=dtypes)\ntrain.head()","metadata":{"executionInfo":{"elapsed":156650,"status":"ok","timestamp":1720312381178,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"XPFNOqWXqeAt","outputId":"7b38848b-d419-4115-eb8c-e178bf66fbbf"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"elements = {\n    15: 'P', # 66:'Dy\n     6: 'C',\n     8: 'O',\n     7: 'N',\n    17:'Cl',\n    16: 'S',\n     9: 'F',\n    35:'Br',\n    53: 'I',\n     5: 'B',\n    14:'Si'\n}","metadata":{"id":"aoxXc-Q0qeAu"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"i = 0\nfor Z in elements:\n    elements[Z] = i\n    i += 1\nelements","metadata":{"executionInfo":{"elapsed":8,"status":"ok","timestamp":1720312381179,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"KESyGuBMqeAu","outputId":"f1612e67-1eed-489e-d5eb-a6972fe13678"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hyb = {}\nfor h in ['SP','SP2','SP3']:\n    hyb[h] = i\n    i += 1\nhyb","metadata":{"executionInfo":{"elapsed":8,"status":"ok","timestamp":1720312381179,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"TAy50s_7qeAv","outputId":"635f9cf9-6485-4f51-96b6-54a35b1f2657"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bonds = {}\nfor b in [1,2,3,4]:\n    bonds[b] = i\n    i += 1\nbonds","metadata":{"executionInfo":{"elapsed":8,"status":"ok","timestamp":1720312381180,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"F-FrEd1zqeAv","outputId":"c476d438-1717-456d-abac-8d2e02362560"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ring = {}\nfor r in [0,1]:\n    ring[r] = i\n    i += 1\nring","metadata":{"executionInfo":{"elapsed":8,"status":"ok","timestamp":1720312381180,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"iYCEv2aAqeAw","outputId":"cd605ba1-194c-4d38-cffc-7cd10ca48278"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"arom = {}\nfor a in [0,1]:\n    arom[a] = i\n    i += 1\narom","metadata":{"executionInfo":{"elapsed":7,"status":"ok","timestamp":1720312381180,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"-XpisuTiqeAw","outputId":"55cbb113-e943-4477-ee7c-92a7ab040c04"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"hyd = {}\nfor H in [0,1,2,3]:\n    hyd[H] = i\n    i += 1\nhyd","metadata":{"executionInfo":{"elapsed":6,"status":"ok","timestamp":1720312381180,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"XI8eq9O5qeAw","outputId":"a03884e9-e7f0-473c-9897-e24a89af9afd"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bond_type = {\n    'SINGLE':0,\n    'DOUBLE':1,\n    'TRIPLE':2,\n    'AROMATIC':3\n}","metadata":{"id":"kL8u2V6IqeAx"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"id":"PGfr5j7rqeAx"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BELKA_DS(torch.utils.data.Dataset):\n    '''\n    '''\n    def __init__(self, df, f, ff, VALID=False):\n\n        positives = df[df[T]==1]\n        negatives = df[df[T]==0]\n\n        N = 271//CV\n        START = f*N\n        if VALID:\n            self.MASK_NODES = 0\n            self.MASK_EDGES = 0\n            self.MASK_GC = 0\n\n            positives = positives[positives['buildingblock1_smiles'].isin(np.arange(START,START+N).tolist())]\n            negatives = negatives[negatives['buildingblock1_smiles'].isin(np.arange(START,START+N).tolist())]\n\n            data = pd.concat([positives[:BS-BS//2],negatives[:BS//2]])\n            self.data = data.sample(len(data)).reset_index(drop=True)\n\n        else:\n            self.MASK_NODES = MASK_NODES\n            self.MASK_EDGES = MASK_EDGES\n            self.MASK_GC = MASK_GC\n\n            positives = positives[~positives['buildingblock1_smiles'].isin(np.arange(START,START+N).tolist())]\n            negatives = negatives[~negatives['buildingblock1_smiles'].isin(np.arange(START,START+N).tolist())]\n\n            num_pos = len(positives)\n            num_neg = ((2*num_pos)//BS + 1)*BS - num_pos\n\n            START = ff*num_neg\n            data = pd.concat([positives,negatives[START:START+num_neg]])\n            self.data = data.sample(len(data)).reset_index(drop=True)\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n\n        df = self.data\n\n        m = Chem.MolFromSmiles(df['molecule_smiles'][idx].replace('Dy','P'))\n        label = torch.tensor([df[T][idx]],device=device).unsqueeze(0).float()\n\n        GC = np.random.rand() > self.MASK_GC\n        if GC: AllChem.ComputeGasteigerCharges(m)\n\n        N = m.GetNumAtoms()\n        empty = torch.ones((N,N),dtype=bool,device=device)\n        AM0 = torch.zeros((N,N),device=device)\n        AM = torch.zeros((N,N),device=device)\n        idx = torch.arange(N)\n        empty[idx,idx] = False\n\n        nodes = torch.zeros((Lmax,27))#.to(device)\n        edges = []\n        edge_attr = []\n        for i in range(N):\n            atom = m.GetAtomWithIdx(i)\n            Z = atom.GetAtomicNum()\n            if GC:\n                Q_i = float(atom.GetProp('_GasteigerCharge'))\n                nodes[i,26] = Q_i\n\n            if Z == 15:\n                nodes[i,elements[Z]] = 1\n            else:\n                if np.random.rand() > self.MASK_NODES:\n                    nodes[i,elements[Z]] = 1\n                else:\n                    nodes[i,1:11]  = .1\n\n                if np.random.rand() > self.MASK_NODES:\n                    nodes[i,hyb[str(atom.GetHybridization())]] = 1\n                else:\n                    nodes[i,11:14] = 1./3\n\n                if np.random.rand() > self.MASK_NODES:\n                    nodes[i,bonds[atom.GetDegree()]] = 1\n                else:\n                    nodes[i,14:18] = .25\n\n                if np.random.rand() > self.MASK_NODES:\n                    nodes[i,ring[atom.IsInRing()]] = 1\n                else:\n                    nodes[i,18:20] = .5\n\n                if np.random.rand() > self.MASK_NODES:\n                    nodes[i,arom[atom.GetIsAromatic()]] = 1\n                else:\n                    nodes[i,20:22] = .5\n\n                if np.random.rand() > self.MASK_NODES:\n                    nodes[i,hyd[atom.GetTotalNumHs()]] = 1\n                else:\n                    nodes[i,22:26] = .25\n\n            for j in range(i):\n                bond = m.GetBondBetweenAtoms(i, j)\n                if bond is not None:\n                    edge = torch.zeros(2,6)\n                    atom = m.GetAtomWithIdx(j)\n                    if GC:\n                        Q_j = float(atom.GetProp('_GasteigerCharge'))\n                        Q = Q_j > Q_i\n                        edge[0,-1 - Q] = 1\n                        edge[1,-2 + Q] = 1\n                    else:\n                        edge[:,-2:] = .5\n\n                    if np.random.rand() > self.MASK_EDGES:\n                        edge[:,bond_type[str(bond.GetBondType())]] = 1\n                    else:\n                        edge[:,:4] = .25\n\n                    edges.append(torch.tensor([i,j]))\n                    edges.append(torch.tensor([j,i]))\n                    edge_attr.append(edge)\n\n                    AM0[i,j] = 1\n                    AM0[j,i] = 1\n\n        i = 0\n        new = (torch.matmul((~empty).float(),AM0) > 0)*(empty)\n        while torch.sum(new)>0:\n            i += 1\n            AM[new] = i\n            empty[new] = False\n            new = (torch.matmul((~empty).float(),AM0) > 0)*(empty)\n\n        edges = torch.stack(edges).to(device)\n        edge_attr = torch.cat(edge_attr).to(device)\n        mask = torch.zeros(Lmax,dtype=bool).to(device)\n        mask[N:] =  True\n        AMN = -torch.ones(Lmax,Lmax).long().to(device)\n        AMN[:N,:N] = AM\n\n        return nodes.to(device),AMN.unsqueeze(0),mask,edges,edge_attr,torch.tensor(len(edges)).unsqueeze(0).to(device),label","metadata":{"id":"CZmVYHVtqeAx"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def my_collate_fn(data):\n    collation = [torch.cat(s) for s in zip(*data)]\n    acce = 0\n    for i in range(1,len(collation[0].view(-1,Lmax,27))):\n            acce += collation[5][i-1]\n            collation[3][acce:] += Lmax\n\n    return collation[:5],collation[-1]","metadata":{"id":"w5bGvzJEqeAx"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Original from https://pytorch-geometric.readthedocs.io/en/latest/_modules/torch_geometric/nn/conv/gat_conv.html#GATConv\n# Edge features EDIT\n# In a Chemistry structure edge features should take part into aggregation\n# Here we virtually concatenate edge features on src nodes\n\nimport typing\nfrom typing import Optional, Tuple, Union\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import Tensor\nfrom torch.nn import Parameter\n\nfrom torch_geometric.nn.conv import MessagePassing\nfrom torch_geometric.nn.dense.linear import Linear as gLinear\nfrom torch_geometric.nn.inits import glorot, zeros\nfrom torch_geometric.typing import (\n    Adj,\n    NoneType,\n    OptPairTensor,\n    OptTensor,\n    Size,\n    SparseTensor,\n    torch_sparse,\n)\nfrom torch_geometric.utils import (\n    add_self_loops,\n    is_torch_sparse_tensor,\n    remove_self_loops,\n    softmax,\n)\nfrom 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\n\n\nclass myGATConv(MessagePassing):\n    def __init__(\n        self,\n        in_channels: Union[int, Tuple[int, int]],\n        out_channels: int,\n        heads: int = HEADS,\n        reduce: str = 'max',\n        negative_slope: float = 0.2,\n        dropout: float = 0.0,\n        add_self_loops: bool = True,\n        edge_dim: Optional[int] = 6,\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.in_channels = in_channels\n        self.out_channels = out_channels\n        self.heads = heads\n        self.reduce = reduce\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\n        self.V = gLinear(in_channels, out_channels, bias=False, weight_initializer='glorot')\n        # Query and Key transformations\n        self.K = gLinear(in_channels, out_channels, bias=False, weight_initializer='glorot')\n        self.Q = gLinear(in_channels, out_channels, bias=False, weight_initializer='glorot')\n\n        # We will only use lin_edge to virtually concatenate edge_attr to src\n        self.V_edge = gLinear(edge_dim, out_channels, bias=False, weight_initializer='glorot')\n        self.K_edge = gLinear(edge_dim, out_channels, bias=False, weight_initializer='glorot')\n\n        if bias:\n            self.bias = Parameter(torch.empty(out_channels))\n        else:\n            self.register_parameter('bias', None)\n\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        super().reset_parameters()\n        self.V.reset_parameters()\n        self.V_edge.reset_parameters()\n        self.K_edge.reset_parameters()\n        self.K.reset_parameters()\n        self.Q.reset_parameters()\n        zeros(self.bias)\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        inputs: tuple,\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        # NOTE: attention weights will be returned whenever\n        # `return_attention_weights` is set to a value, regardless of its\n        # actual value (might be `True` or `False`). This is a current somewhat\n        # hacky workaround to allow for TorchScript support via the\n        # `torch.jit._overload` decorator, as we can only change the output\n        # arguments conditioned on type (`None` or `bool`), not based on its\n        # actual value.\n#==================================================================================\n        x,edge_index,edge_attr = inputs\n        H, C = self.heads, self.out_channels//self.heads\n        x_dst = self.V(x)\n        num_nodes = x.size(0)\n        num_edges = edge_index.size(0)\n        EDGE_INDEX = torch.cat([\n            edge_index,\n            torch.cat([\n                torch.arange(num_nodes).unsqueeze(-1),\n                torch.arange(num_nodes).unsqueeze(-1)\n            ],-1).to(device)\n        ])\n\n        Q = self.Q(x).view(-1,H,C)\n\n        x_src = x_dst[EDGE_INDEX[:,0]]\n        K = self.K(x[EDGE_INDEX[:,0]])\n#       we simply need to sum them up to \"emulate\" concatenation:\n        x_src[:num_edges] = x_src[:num_edges] + self.V_edge(edge_attr)\n        K[:num_edges] = K[:num_edges] + self.K_edge(edge_attr)\n        x_dst = x_dst.view(-1,H,C)\n        x_src = x_src.view(-1,H,C)\n        K = K.view(-1,H,C)\n\n        alpha_src = (K*Q[EDGE_INDEX[:,1]]).sum(-1)\n        alpha_dst =torch.zeros(num_nodes).to(device)\n\n        EDGE_INDEX = torch.cat([\n            torch.arange(len(EDGE_INDEX)).unsqueeze(0).to(device),\n            EDGE_INDEX[:,1].unsqueeze(0)\n        ])\n\n        out = torch.zeros((num_nodes,H,C)).to(device)\n        for i in range(H):\n            alpha = (alpha_src[:,i], alpha_dst)\n        #   edge_updater_type: (alpha: OptPairTensor, edge_attr: OptTensor)\n            alpha = self.edge_updater(EDGE_INDEX, alpha=alpha, edge_attr=edge_attr,size=size)\n        #   propagate_type: (x: OptPairTensor, alpha: Tensor)\n            x = (x_src[:,i], x_dst[:,i])\n            out[:,i] = self.propagate(EDGE_INDEX, x=x, alpha=alpha, size=size)\n\n        out = out.view(-1, self.out_channels)\n\n        if self.bias is not None:\n            out = out + self.bias\n\n        return out,edge_index,edge_attr\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        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":{"id":"oYCo_PD0qeAx"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://github.com/pytorch/pytorch/blob/main/torch/nn/functional.py\n\"\"\"Functional interface.\"\"\"\n\nimport importlib\nimport math\nimport warnings\nfrom typing import Callable, List, Optional, Tuple, TYPE_CHECKING, Union\n\nimport torch\n# from torch import _VF, sym_int as _sym_int, Tensor\nfrom torch._C import _add_docstr, _infer_size\nfrom torch._jit_internal import (\n    _overload,\n    boolean_dispatch,\n    BroadcastingList1,\n    BroadcastingList2,\n    BroadcastingList3,\n)\nfrom torch._torch_docs import reproducibility_notes, sparse_support_notes, tf32_notes\nfrom torch.nn import _reduction as _Reduction, grad  # noqa: F401\nfrom torch.nn.modules.utils import _list_with_default, _pair, _single, _triple\nfrom torch.overrides import (\n    handle_torch_function,\n    has_torch_function,\n    has_torch_function_unary,\n    has_torch_function_variadic,\n)\n\nlinear = torch._C._nn.linear\n\nif TYPE_CHECKING:\n    from torch.types import _dtype as DType\nelse:\n    # The JIT doesn't understand Union, nor torch.dtype here\n    DType = int\n\ndef pad(\n    input: Tensor,\n    pad: List[int],\n    mode: str = \"constant\",\n    value: Optional[float] = None,\n) -> Tensor:\n\n    if has_torch_function_unary(input):\n        return handle_torch_function(\n            torch.nn.functional.pad, (input,), input, pad, mode=mode, value=value\n        )\n    if not torch.jit.is_scripting():\n        if torch.are_deterministic_algorithms_enabled() and input.is_cuda:\n            if mode == \"replicate\":\n                # Use slow decomp whose backward will be in terms of index_put.\n                # importlib is required because the import cannot be top level\n                # (cycle) and cannot be nested (TS doesn't support)\n                return importlib.import_module(\n                    \"torch._decomp.decompositions\"\n                )._replication_pad(input, pad)\n    return torch._C._nn.pad(input, pad, mode, value)\n\ndef _none_or_dtype(input: Optional[Tensor]) -> Optional[DType]:\n    if input is None:\n        return None\n    elif isinstance(input, torch.Tensor):\n        return input.dtype\n    raise RuntimeError(\"input to _none_or_dtype() must be None or torch.Tensor\")\n\ndef _mha_shape_check(\n    query: Tensor,\n    key: Tensor,\n    value: Tensor,\n    key_padding_mask: Optional[Tensor],\n    attn_mask: Optional[Tensor],\n    num_heads: int,\n):\n    # Verifies the expected shape for `query, `key`, `value`, `key_padding_mask` and `attn_mask`\n    # and returns if the input is batched or not.\n    # Raises an error if `query` is not 2-D (unbatched) or 3-D (batched) tensor.\n\n    # Shape check.\n    if query.dim() == 3:\n        # Batched Inputs\n        is_batched = True\n        assert key.dim() == 3 and value.dim() == 3, (\n            \"For batched (3-D) `query`, expected `key` and `value` to be 3-D\"\n            f\" but found {key.dim()}-D and {value.dim()}-D tensors respectively\"\n        )\n        if key_padding_mask is not None:\n            assert key_padding_mask.dim() == 2, (\n                \"For batched (3-D) `query`, expected `key_padding_mask` to be `None` or 2-D\"\n                f\" but found {key_padding_mask.dim()}-D tensor instead\"\n            )\n        if attn_mask is not None:\n            assert attn_mask.dim() in (2, 3), (\n                \"For batched (3-D) `query`, expected `attn_mask` to be `None`, 2-D or 3-D\"\n                f\" but found {attn_mask.dim()}-D tensor instead\"\n            )\n    elif query.dim() == 2:\n        # Unbatched Inputs\n        is_batched = False\n        assert key.dim() == 2 and value.dim() == 2, (\n            \"For unbatched (2-D) `query`, expected `key` and `value` to be 2-D\"\n            f\" but found {key.dim()}-D and {value.dim()}-D tensors respectively\"\n        )\n\n        if key_padding_mask is not None:\n            assert key_padding_mask.dim() == 1, (\n                \"For unbatched (2-D) `query`, expected `key_padding_mask` to be `None` or 1-D\"\n                f\" but found {key_padding_mask.dim()}-D tensor instead\"\n            )\n\n        if attn_mask is not None:\n            assert attn_mask.dim() in (2, 3), (\n                \"For unbatched (2-D) `query`, expected `attn_mask` to be `None`, 2-D or 3-D\"\n                f\" but found {attn_mask.dim()}-D tensor instead\"\n            )\n            if attn_mask.dim() == 3:\n                expected_shape = (num_heads, query.shape[0], key.shape[0])\n                assert (\n                    attn_mask.shape == expected_shape\n                ), f\"Expected `attn_mask` shape to be {expected_shape} but got {attn_mask.shape}\"\n    else:\n        raise AssertionError(\n            f\"query should be unbatched 2D or batched 3D tensor but received {query.dim()}-D query tensor\"\n        )\n\n    return is_batched\n\ndef _canonical_mask(\n    mask: Optional[Tensor],\n    mask_name: str,\n    other_type: Optional[DType],\n    other_name: str,\n    target_type: DType,\n    check_other: bool = True,\n) -> Optional[Tensor]:\n    if mask is not None:\n        _mask_dtype = mask.dtype\n        _mask_is_float = torch.is_floating_point(mask)\n        if _mask_dtype != torch.bool and not _mask_is_float:\n            raise AssertionError(\n                f\"only bool and floating types of {mask_name} are supported\"\n            )\n        if check_other and other_type is not None:\n            if _mask_dtype != other_type:\n                warnings.warn(\n                    f\"Support for mismatched {mask_name} and {other_name} \"\n                    \"is deprecated. Use same type for both instead.\"\n                )\n        if not _mask_is_float:\n            mask = torch.zeros_like(mask, dtype=target_type).masked_fill_(\n                mask, float(\"-inf\")\n            )\n    return mask\n\ndef _in_projection_packed(\n    q: Tensor,\n    k: Tensor,\n    v: Tensor,\n    w: Tensor,\n    b: Optional[Tensor] = None,\n) -> List[Tensor]:\n\n    E = q.size(-1)\n    if k is v:\n        if q is k:\n            # self-attention\n            proj = linear(q, w, b)\n            # reshape to 3, E and not E, 3 is deliberate for better memory coalescing and keeping same order as chunk()\n            proj = (\n                proj.unflatten(-1, (3, E))\n                .unsqueeze(0)\n                .transpose(0, -2)\n                .squeeze(-2)\n                .contiguous()\n            )\n            return proj[0], proj[1], proj[2]\n        else:\n            # encoder-decoder attention\n            w_q, w_kv = w.split([E, E * 2])\n            if b is None:\n                b_q = b_kv = None\n            else:\n                b_q, b_kv = b.split([E, E * 2])\n            q_proj = linear(q, w_q, b_q)\n            kv_proj = linear(k, w_kv, b_kv)\n            # reshape to 2, E and not E, 2 is deliberate for better memory coalescing and keeping same order as chunk()\n            kv_proj = (\n                kv_proj.unflatten(-1, (2, E))\n                .unsqueeze(0)\n                .transpose(0, -2)\n                .squeeze(-2)\n                .contiguous()\n            )\n            return (q_proj, kv_proj[0], kv_proj[1])\n    else:\n        w_q, w_k, w_v = w.chunk(3)\n        if b is None:\n            b_q = b_k = b_v = None\n        else:\n            b_q, b_k, b_v = b.chunk(3)\n        return linear(q, w_q, b_q), linear(k, w_k, b_k), linear(v, w_v, b_v)\n\n\ndef _in_projection(\n    q: Tensor,\n    k: Tensor,\n    v: Tensor,\n    w_q: Tensor,\n    w_k: Tensor,\n    w_v: Tensor,\n    b_q: Optional[Tensor] = None,\n    b_k: Optional[Tensor] = None,\n    b_v: Optional[Tensor] = None,\n) -> Tuple[Tensor, Tensor, Tensor]:\n\n    Eq, Ek, Ev = q.size(-1), k.size(-1), v.size(-1)\n    assert w_q.shape == (\n        Eq,\n        Eq,\n    ), f\"expecting query weights shape of {(Eq, Eq)}, but got {w_q.shape}\"\n    assert w_k.shape == (\n        Eq,\n        Ek,\n    ), f\"expecting key weights shape of {(Eq, Ek)}, but got {w_k.shape}\"\n    assert w_v.shape == (\n        Eq,\n        Ev,\n    ), f\"expecting value weights shape of {(Eq, Ev)}, but got {w_v.shape}\"\n    assert b_q is None or b_q.shape == (\n        Eq,\n    ), f\"expecting query bias shape of {(Eq,)}, but got {b_q.shape}\"\n    assert b_k is None or b_k.shape == (\n        Eq,\n    ), f\"expecting key bias shape of {(Eq,)}, but got {b_k.shape}\"\n    assert b_v is None or b_v.shape == (\n        Eq,\n    ), f\"expecting value bias shape of {(Eq,)}, but got {b_v.shape}\"\n    return linear(q, w_q, b_q), linear(k, w_k, b_k), linear(v, w_v, b_v)","metadata":{"id":"Blpd9wP8qeAy"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://github.com/pytorch/pytorch/blob/dc4db95540da06623c747bf0f2bf9f4af3d2925a/torch/nn/functional.py\n\nfrom torch import _VF\n\ndef dropout(input, p=0.5, training=True, inplace=False):\n    # type: (Tensor, float, bool, bool) -> Tensor\n    r\"\"\"\n    During training, randomly zeroes some of the elements of the input\n    tensor with probability :attr:`p` using samples from a Bernoulli\n    distribution.\n\n    See :class:`~torch.nn.Dropout` for details.\n\n    Args:\n        p: probability of an element to be zeroed. Default: 0.5\n        training: apply dropout if is ``True``. Default: ``True``\n        inplace: If set to ``True``, will do this operation in-place. Default: ``False``\n    \"\"\"\n    if not torch.jit.is_scripting():\n        if type(input) is not Tensor and has_torch_function((input,)):\n            return handle_torch_function(\n                dropout, (input,), input, p=p, training=training, inplace=inplace)\n    if p < 0. or p > 1.:\n        raise ValueError(\"dropout probability has to be between 0 and 1, \"\n                         \"but got {}\".format(p))\n    return (_VF.dropout_(input, p, training)\n            if inplace\n            else _VF.dropout(input, p, training))\n\ndef multi_head_attention_forward(query: Tensor,\n                                 key: Tensor,\n                                 value: Tensor,\n                                 A: Tensor,\n                                 B: Tensor,\n                                 embed_dim_to_check: int,\n                                 num_heads: int,\n                                 in_proj_weight: Tensor,\n                                 in_proj_bias: Tensor,\n                                 bias_k: Optional[Tensor],\n                                 bias_v: Optional[Tensor],\n                                 add_zero_attn: bool,\n                                 dropout_p: float,\n                                 out_proj_weight: Tensor,\n                                 out_proj_bias: Tensor,\n                                 training: bool = True,\n                                 key_padding_mask: Optional[Tensor] = None,\n                                 need_weights: bool = True,\n                                 attn_mask: Optional[Tensor] = None,\n                                 use_separate_proj_weight: bool = False,\n                                 q_proj_weight: Optional[Tensor] = None,\n                                 k_proj_weight: Optional[Tensor] = None,\n                                 v_proj_weight: Optional[Tensor] = None,\n                                 static_k: Optional[Tensor] = None,\n                                 static_v: Optional[Tensor] = None\n                                 ) -> Tuple[Tensor, Optional[Tensor]]:\n    r\"\"\"\n    Args:\n        query, key, value: map a query and a set of key-value pairs to an output.\n            See \"Attention Is All You Need\" for more details.\n        embed_dim_to_check: total dimension of the model.\n        num_heads: parallel attention heads.\n        in_proj_weight, in_proj_bias: input projection weight and bias.\n        bias_k, bias_v: bias of the key and value sequences to be added at dim=0.\n        add_zero_attn: add a new batch of zeros to the key and\n                       value sequences at dim=1.\n        dropout_p: probability of an element to be zeroed.\n        out_proj_weight, out_proj_bias: the output projection weight and bias.\n        training: apply dropout if is ``True``.\n        key_padding_mask: if provided, specified padding elements in the key will\n            be ignored by the attention. This is an binary mask. When the value is True,\n            the corresponding value on the attention layer will be filled with -inf.\n        need_weights: output attn_output_weights.\n        attn_mask: 2D or 3D mask that prevents attention to certain positions. A 2D mask will be broadcasted for all\n            the batches while a 3D mask allows to specify a different mask for the entries of each batch.\n        use_separate_proj_weight: the function accept the proj. weights for query, key,\n            and value in different forms. If false, in_proj_weight will be used, which is\n            a combination of q_proj_weight, k_proj_weight, v_proj_weight.\n        q_proj_weight, k_proj_weight, v_proj_weight, in_proj_bias: input projection weight and bias.\n        static_k, static_v: static key and value used for attention operators.\n\n\n    Shape:\n        Inputs:\n        - query: :math:`(L, N, E)` where L is the target sequence length, N is the batch size, E is\n          the embedding dimension.\n        - key: :math:`(S, N, E)`, where S is the source sequence length, N is the batch size, E is\n          the embedding dimension.\n        - value: :math:`(S, N, E)` where S is the source sequence length, N is the batch size, E is\n          the embedding dimension.\n        - key_padding_mask: :math:`(N, S)` where N is the batch size, S is the source sequence length.\n          If a ByteTensor is provided, the non-zero positions will be ignored while the zero positions\n          will be unchanged. If a BoolTensor is provided, the positions with the\n          value of ``True`` will be ignored while the position with the value of ``False`` will be unchanged.\n        - attn_mask: 2D mask :math:`(L, S)` where L is the target sequence length, S is the source sequence length.\n          3D mask :math:`(N*num_heads, L, S)` where N is the batch size, L is the target sequence length,\n          S is the source sequence length. attn_mask ensures that position i is allowed to attend the unmasked\n          positions. If a ByteTensor is provided, the non-zero positions are not allowed to attend\n          while the zero positions will be unchanged. If a BoolTensor is provided, positions with ``True``\n          are not allowed to attend while ``False`` values will be unchanged. If a FloatTensor\n          is provided, it will be added to the attention weight.\n        - static_k: :math:`(N*num_heads, S, E/num_heads)`, where S is the source sequence length,\n          N is the batch size, E is the embedding dimension. E/num_heads is the head dimension.\n        - static_v: :math:`(N*num_heads, S, E/num_heads)`, where S is the source sequence length,\n          N is the batch size, E is the embedding dimension. E/num_heads is the head dimension.\n\n        Outputs:\n        - attn_output: :math:`(L, N, E)` where L is the target sequence length, N is the batch size,\n          E is the embedding dimension.\n        - attn_output_weights: :math:`(N, L, S)` where N is the batch size,\n          L is the target sequence length, S is the source sequence length.\n    \"\"\"\n    if not torch.jit.is_scripting():\n        tens_ops = (query, key, value, in_proj_weight, in_proj_bias, bias_k, bias_v,\n                    out_proj_weight, out_proj_bias)\n        if any([type(t) is not Tensor for t in tens_ops]) and has_torch_function(tens_ops):\n            return handle_torch_function(\n                multi_head_attention_forward, tens_ops, query, key, value,\n                embed_dim_to_check, num_heads, in_proj_weight, in_proj_bias,\n                bias_k, bias_v, add_zero_attn, dropout_p, out_proj_weight,\n                out_proj_bias, training=training, key_padding_mask=key_padding_mask,\n                need_weights=need_weights, attn_mask=attn_mask,\n                use_separate_proj_weight=use_separate_proj_weight,\n                q_proj_weight=q_proj_weight, k_proj_weight=k_proj_weight,\n                v_proj_weight=v_proj_weight, static_k=static_k, static_v=static_v)\n    tgt_len, bsz, embed_dim = query.size()\n    assert embed_dim == embed_dim_to_check\n    # allow MHA to have different sizes for the feature dimension\n    assert key.size(0) == value.size(0) and key.size(1) == value.size(1)\n\n    head_dim = embed_dim // num_heads\n    assert head_dim * num_heads == embed_dim, \"embed_dim must be divisible by num_heads\"\n    scaling = float(head_dim) ** -0.5\n\n    if not use_separate_proj_weight:\n        if (query is key or torch.equal(query, key)) and (key is value or torch.equal(key, value)):\n            # self-attention\n            q, k, v = linear(query, in_proj_weight, in_proj_bias).chunk(3, dim=-1)\n\n        elif (key is value or torch.equal(key, value)):\n            # encoder-decoder attention\n            # This is inline in_proj function with in_proj_weight and in_proj_bias\n            _b = in_proj_bias\n            _start = 0\n            _end = embed_dim\n            _w = in_proj_weight[_start:_end, :]\n            if _b is not None:\n                _b = _b[_start:_end]\n            q = linear(query, _w, _b)\n\n            if key is None:\n                assert value is None\n                k = None\n                v = None\n            else:\n\n                # This is inline in_proj function with in_proj_weight and in_proj_bias\n                _b = in_proj_bias\n                _start = embed_dim\n                _end = None\n                _w = in_proj_weight[_start:, :]\n                if _b is not None:\n                    _b = _b[_start:]\n                k, v = linear(key, _w, _b).chunk(2, dim=-1)\n\n        else:\n            # This is inline in_proj function with in_proj_weight and in_proj_bias\n            _b = in_proj_bias\n            _start = 0\n            _end = embed_dim\n            _w = in_proj_weight[_start:_end, :]\n            if _b is not None:\n                _b = _b[_start:_end]\n            q = linear(query, _w, _b)\n\n            # This is inline in_proj function with in_proj_weight and in_proj_bias\n            _b = in_proj_bias\n            _start = embed_dim\n            _end = embed_dim * 2\n            _w = in_proj_weight[_start:_end, :]\n            if _b is not None:\n                _b = _b[_start:_end]\n            k = linear(key, _w, _b)\n\n            # This is inline in_proj function with in_proj_weight and in_proj_bias\n            _b = in_proj_bias\n            _start = embed_dim * 2\n            _end = None\n            _w = in_proj_weight[_start:, :]\n            if _b is not None:\n                _b = _b[_start:]\n            v = linear(value, _w, _b)\n    else:\n        q_proj_weight_non_opt = torch.jit._unwrap_optional(q_proj_weight)\n        len1, len2 = q_proj_weight_non_opt.size()\n        assert len1 == embed_dim and len2 == query.size(-1)\n\n        k_proj_weight_non_opt = torch.jit._unwrap_optional(k_proj_weight)\n        len1, len2 = k_proj_weight_non_opt.size()\n        assert len1 == embed_dim and len2 == key.size(-1)\n\n        v_proj_weight_non_opt = torch.jit._unwrap_optional(v_proj_weight)\n        len1, len2 = v_proj_weight_non_opt.size()\n        assert len1 == embed_dim and len2 == value.size(-1)\n\n        if in_proj_bias is not None:\n            q = linear(query, q_proj_weight_non_opt, in_proj_bias[0:embed_dim])\n            k = linear(key, k_proj_weight_non_opt, in_proj_bias[embed_dim:(embed_dim * 2)])\n            v = linear(value, v_proj_weight_non_opt, in_proj_bias[(embed_dim * 2):])\n        else:\n            q = linear(query, q_proj_weight_non_opt, in_proj_bias)\n            k = linear(key, k_proj_weight_non_opt, in_proj_bias)\n            v = linear(value, v_proj_weight_non_opt, in_proj_bias)\n    q = q * scaling\n\n    if attn_mask is not None:\n        assert attn_mask.dtype == torch.float32 or attn_mask.dtype == torch.float64 or \\\n            attn_mask.dtype == torch.float16 or attn_mask.dtype == torch.uint8 or attn_mask.dtype == torch.bool, \\\n            'Only float, byte, and bool types are supported for attn_mask, not {}'.format(attn_mask.dtype)\n        if attn_mask.dtype == torch.uint8:\n            warnings.warn(\"Byte tensor for attn_mask in nn.MultiheadAttention is deprecated. Use bool tensor instead.\")\n            attn_mask = attn_mask.to(torch.bool)\n\n        if attn_mask.dim() == 2:\n            attn_mask = attn_mask.unsqueeze(0)\n            if list(attn_mask.size()) != [1, query.size(0), key.size(0)]:\n                raise RuntimeError('The size of the 2D attn_mask is not correct.')\n        elif attn_mask.dim() == 3:\n            if list(attn_mask.size()) != [bsz * num_heads, query.size(0), key.size(0)]:\n                raise RuntimeError('The size of the 3D attn_mask is not correct.')\n        else:\n            raise RuntimeError(\"attn_mask's dimension {} is not supported\".format(attn_mask.dim()))\n        # attn_mask's dim is 3 now.\n\n    # convert ByteTensor key_padding_mask to bool\n    if key_padding_mask is not None and key_padding_mask.dtype == torch.uint8:\n        warnings.warn(\"Byte tensor for key_padding_mask in nn.MultiheadAttention is deprecated. Use bool tensor instead.\")\n        key_padding_mask = key_padding_mask.to(torch.bool)\n\n    if bias_k is not None and bias_v is not None:\n        if static_k is None and static_v is None:\n            k = torch.cat([k, bias_k.repeat(1, bsz, 1)])\n            v = torch.cat([v, bias_v.repeat(1, bsz, 1)])\n            if attn_mask is not None:\n                attn_mask = pad(attn_mask, (0, 1))\n            if key_padding_mask is not None:\n                key_padding_mask = pad(key_padding_mask, (0, 1))\n        else:\n            assert static_k is None, \"bias cannot be added to static key.\"\n            assert static_v is None, \"bias cannot be added to static value.\"\n    else:\n        assert bias_k is None\n        assert bias_v is None\n\n    q = q.contiguous().view(tgt_len, bsz * num_heads, head_dim).transpose(0, 1)\n    if k is not None:\n        k = k.contiguous().view(-1, bsz * num_heads, head_dim).transpose(0, 1)\n    if v is not None:\n        v = v.contiguous().view(-1, bsz * num_heads, head_dim).transpose(0, 1)\n\n    if static_k is not None:\n        assert static_k.size(0) == bsz * num_heads\n        assert static_k.size(2) == head_dim\n        k = static_k\n\n    if static_v is not None:\n        assert static_v.size(0) == bsz * num_heads\n        assert static_v.size(2) == head_dim\n        v = static_v\n\n    src_len = k.size(1)\n\n    if key_padding_mask is not None:\n        assert key_padding_mask.size(0) == bsz\n        assert key_padding_mask.size(1) == src_len\n\n    if add_zero_attn:\n        src_len += 1\n        k = torch.cat([k, torch.zeros((k.size(0), 1) + k.size()[2:], dtype=k.dtype, device=k.device)], dim=1)\n        v = torch.cat([v, torch.zeros((v.size(0), 1) + v.size()[2:], dtype=v.dtype, device=v.device)], dim=1)\n        if attn_mask is not None:\n            attn_mask = pad(attn_mask, (0, 1))\n        if key_padding_mask is not None:\n            key_padding_mask = pad(key_padding_mask, (0, 1))\n\n    attn_output_weights = torch.bmm(q, k.transpose(1, 2))\n    assert list(attn_output_weights.size()) == [bsz * num_heads, tgt_len, src_len]\n\n    if attn_mask is not None:\n        if attn_mask.dtype == torch.bool:\n            attn_output_weights.masked_fill_(attn_mask, float('-inf'))\n        else:\n            attn_output_weights += attn_mask\n\n    attn_output_weights = attn_output_weights.view(bsz, num_heads, tgt_len, src_len)\n\n    key_padding_mask = ((key_padding_mask.unsqueeze(-1)@key_padding_mask.unsqueeze(1))==0).unsqueeze(1).tile((1,num_heads,1,1))\n\n    attn_output_weights = attn_output_weights.masked_fill(\n            ~key_padding_mask,\n            float('-inf'),\n    )\n\n    attn_output_weights[key_padding_mask] = A.transpose(0,1)[key_padding_mask]*attn_output_weights[key_padding_mask]\n\n    attn_output_weights[key_padding_mask] = attn_output_weights[key_padding_mask] + B.transpose(0,1)[key_padding_mask]\n\n    attn_output_weights = attn_output_weights.view(bsz * num_heads, tgt_len, src_len)\n\n    key_padding_mask = key_padding_mask.view(bsz * num_heads, tgt_len, src_len)[:,:,0]\n\n    attn_output_weights[key_padding_mask] = torch.softmax(\n        attn_output_weights[key_padding_mask], dim=-1)\n    attn_output_weights[~key_padding_mask] = 0\n\n    attn_output_weights = dropout(attn_output_weights, p=dropout_p, training=training)\n\n    attn_output = torch.bmm(attn_output_weights, v)\n    assert list(attn_output.size()) == [bsz * num_heads, tgt_len, head_dim]\n    attn_output = attn_output.transpose(0, 1).contiguous().view(tgt_len, bsz, embed_dim)\n    attn_output = linear(attn_output, out_proj_weight, out_proj_bias)\n\n    if need_weights:\n        # average attention weights over heads\n        attn_output_weights = attn_output_weights.view(bsz, num_heads, tgt_len, src_len)\n        return attn_output, attn_output_weights.sum(dim=1) / num_heads\n    else:\n        return attn_output, None","metadata":{"id":"eXcaYc7rqeAz"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://pytorch.org/docs/stable/_modules/torch/backends/mha.html#get_fastpath_enabled\n_is_fastpath_enabled: bool = True\n\ndef get_fastpath_enabled() -> bool:\n    if not torch.jit.is_scripting():\n        return _is_fastpath_enabled\n    return True","metadata":{"id":"x1Uc_JO1qeAz"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://github.com/pytorch/pytorch/blob/main/torch/nn/modules/linear.py\nfrom torch.nn import functional as F, init\n\nclass _Linear(Module):\n\n    __constants__ = [\"in_features\", \"out_features\"]\n    in_features: int\n    out_features: int\n    weight: Tensor\n\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        bias: bool = True,\n        device=None,\n        dtype=None,\n    ) -> None:\n        factory_kwargs = {\"device\": device, \"dtype\": dtype}\n        super().__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.weight = Parameter(\n            torch.empty((out_features, in_features), **factory_kwargs)\n        )\n        if bias:\n            self.bias = Parameter(torch.empty(out_features, **factory_kwargs))\n        else:\n            self.register_parameter(\"bias\", None)\n        self.reset_parameters()\n\n    def reset_parameters(self) -> None:\n        # Setting a=sqrt(5) in kaiming_uniform is the same as initializing with\n        # uniform(-1/sqrt(in_features), 1/sqrt(in_features)). For details, see\n        # https://github.com/pytorch/pytorch/issues/57109\n        init.kaiming_uniform_(self.weight, a=math.sqrt(5))\n        if self.bias is not None:\n            fan_in, _ = init._calculate_fan_in_and_fan_out(self.weight)\n            bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0\n            init.uniform_(self.bias, -bound, bound)\n\n    def forward(self, input: Tensor) -> Tensor:\n        return F.linear(input, self.weight, self.bias)\n\n    def extra_repr(self) -> str:\n        return f\"in_features={self.in_features}, out_features={self.out_features}, bias={self.bias is not None}\"\n\nclass NonDynamicallyQuantizableLinear(_Linear):\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        bias: bool = True,\n        device=None,\n        dtype=None,\n    ) -> None:\n        super().__init__(\n            in_features, out_features, bias=bias, device=device, dtype=dtype\n        )\n\n# https://pytorch.org/docs/stable/_modules/torch/nn/init.html#xavier_uniform_\nfrom typing import Optional as _Optional\n\ndef _no_grad_normal_(tensor, mean, std, generator=None):\n    with torch.no_grad():\n        return tensor.normal_(mean, std, generator=generator)\n\ndef xavier_normal_(\n    tensor: Tensor,\n    gain: float = 1.0,\n    generator: _Optional[torch.Generator] = None,\n) -> Tensor:\n\n    fan_in, fan_out = _calculate_fan_in_and_fan_out(tensor)\n    std = gain * math.sqrt(2.0 / float(fan_in + fan_out))\n\n    return _no_grad_normal_(tensor, 0., std, generator)\n\ndef _no_grad_fill_(tensor, val):\n    with torch.no_grad():\n        return tensor.fill_(val)\n\ndef constant_(tensor: Tensor, val: float) -> Tensor:\n\n    if torch.overrides.has_torch_function_variadic(tensor):\n        return torch.overrides.handle_torch_function(constant_, (tensor,), tensor=tensor, val=val)\n    return _no_grad_fill_(tensor, val)\n\ndef _no_grad_uniform_(tensor, a, b, generator=None):\n    with torch.no_grad():\n        return tensor.uniform_(a, b, generator=generator)\n\ndef _calculate_fan_in_and_fan_out(tensor):\n    dimensions = tensor.dim()\n    if dimensions < 2:\n        raise ValueError(\"Fan in and fan out can not be computed for tensor with fewer than 2 dimensions\")\n\n    num_input_fmaps = tensor.size(1)\n    num_output_fmaps = tensor.size(0)\n    receptive_field_size = 1\n    if tensor.dim() > 2:\n        # math.prod is not always available, accumulate the product manually\n        # we could use functools.reduce but that is not supported by TorchScript\n        for s in tensor.shape[2:]:\n            receptive_field_size *= s\n    fan_in = num_input_fmaps * receptive_field_size\n    fan_out = num_output_fmaps * receptive_field_size\n\n    return fan_in, fan_out\n\ndef xavier_uniform_(\n    tensor: Tensor, gain: float = 1.0, generator: _Optional[torch.Generator] = None\n) -> Tensor:\n\n    fan_in, fan_out = _calculate_fan_in_and_fan_out(tensor)\n    std = gain * math.sqrt(2.0 / float(fan_in + fan_out))\n    a = math.sqrt(3.0) * std  # Calculate uniform bounds from standard deviation\n\n    return _no_grad_uniform_(tensor, -a, a, generator)\n\n# https://pytorch.org/docs/stable/_modules/torch/nn/modules/activation.html#MultiheadAttention\ndef _is_make_fx_tracing():\n    if not torch.jit.is_scripting():\n        torch_dispatch_mode_stack = torch.utils._python_dispatch._get_current_dispatch_mode_stack()\n        return any(type(x) == torch.fx.experimental.proxy_tensor.ProxyTorchDispatchMode for x in torch_dispatch_mode_stack)\n    else:\n        return False\n\ndef _check_arg_device(x: Optional[torch.Tensor]) -> bool:\n    if x is not None:\n        return x.device.type in [\"cpu\", \"cuda\"]\n    return True\n\ndef _arg_requires_grad(x: Optional[torch.Tensor]) -> bool:\n    if x is not None:\n        return x.requires_grad\n    return False\n\nclass MultiheadAttention(Module):\n\n    __constants__ = ['batch_first']\n    bias_k: Optional[torch.Tensor]\n    bias_v: Optional[torch.Tensor]\n\n    def __init__(self, embed_dim, num_heads, dropout=0., bias=True, add_bias_kv=False, add_zero_attn=False,\n                 kdim=None, vdim=None, batch_first=False, device=None, dtype=None) -> None:\n        if embed_dim <= 0 or num_heads <= 0:\n            raise ValueError(\n                f\"embed_dim and num_heads must be greater than 0,\"\n                f\" got embed_dim={embed_dim} and num_heads={num_heads} instead\"\n            )\n        factory_kwargs = {'device': device, 'dtype': dtype}\n        super().__init__()\n        self.embed_dim = embed_dim\n        self.kdim = kdim if kdim is not None else embed_dim\n        self.vdim = vdim if vdim is not None else embed_dim\n        self._qkv_same_embed_dim = self.kdim == embed_dim and self.vdim == embed_dim\n\n        self.num_heads = num_heads\n        self.dropout = dropout\n        self.batch_first = batch_first\n        self.head_dim = embed_dim // num_heads\n        assert self.head_dim * num_heads == self.embed_dim, \"embed_dim must be divisible by num_heads\"\n\n        if not self._qkv_same_embed_dim:\n            self.q_proj_weight = Parameter(torch.empty((embed_dim, embed_dim), **factory_kwargs))\n            self.k_proj_weight = Parameter(torch.empty((embed_dim, self.kdim), **factory_kwargs))\n            self.v_proj_weight = Parameter(torch.empty((embed_dim, self.vdim), **factory_kwargs))\n            self.register_parameter('in_proj_weight', None)\n        else:\n#           y = Ax + B\n            self.nbh_A = Parameter(torch.ones((num_heads, NBH), **factory_kwargs))\n            self.nbh_B = Parameter(torch.zeros((num_heads, NBH), **factory_kwargs))\n            self.in_proj_weight = Parameter(torch.empty((3 * embed_dim, embed_dim), **factory_kwargs))\n            self.register_parameter('q_proj_weight', None)\n            self.register_parameter('k_proj_weight', None)\n            self.register_parameter('v_proj_weight', None)\n\n        if bias:\n            self.in_proj_bias = Parameter(torch.empty(3 * embed_dim, **factory_kwargs))\n        else:\n            self.register_parameter('in_proj_bias', None)\n        self.out_proj = NonDynamicallyQuantizableLinear(embed_dim, embed_dim, bias=bias, **factory_kwargs)\n\n        if add_bias_kv:\n            self.bias_k = Parameter(torch.empty((1, 1, embed_dim), **factory_kwargs))\n            self.bias_v = Parameter(torch.empty((1, 1, embed_dim), **factory_kwargs))\n        else:\n            self.bias_k = self.bias_v = None\n\n        self.add_zero_attn = add_zero_attn\n\n        self._reset_parameters()\n\n    def _reset_parameters(self):\n        if self._qkv_same_embed_dim:\n            xavier_uniform_(self.in_proj_weight)\n        else:\n            xavier_uniform_(self.q_proj_weight)\n            xavier_uniform_(self.k_proj_weight)\n            xavier_uniform_(self.v_proj_weight)\n\n        if self.in_proj_bias is not None:\n            constant_(self.in_proj_bias, 0.)\n            constant_(self.out_proj.bias, 0.)\n        if self.bias_k is not None:\n            xavier_normal_(self.bias_k)\n        if self.bias_v is not None:\n            xavier_normal_(self.bias_v)\n\n    def __setstate__(self, state):\n        # Support loading old MultiheadAttention checkpoints generated by v1.1.0\n        if '_qkv_same_embed_dim' not in state:\n            state['_qkv_same_embed_dim'] = True\n\n        super().__setstate__(state)\n\n    def forward(\n            self,\n            query: Tensor,\n            key: Tensor,\n            value: Tensor,\n            AM: Tensor,\n            key_padding_mask: Optional[Tensor] = None,\n            need_weights: bool = True,\n            attn_mask: Optional[Tensor] = None,\n            average_attn_weights: bool = True,\n            is_causal : bool = False) -> Tuple[Tensor, Optional[Tensor]]:\n\n        why_not_fast_path = ''\n        if ((attn_mask is not None and torch.is_floating_point(attn_mask))\n           or (key_padding_mask is not None) and torch.is_floating_point(key_padding_mask)):\n            why_not_fast_path = \"floating-point masks are not supported for fast path.\"\n\n        is_batched = query.dim() == 3\n\n        key_padding_mask = _canonical_mask(\n            mask=key_padding_mask,\n            mask_name=\"key_padding_mask\",\n            other_type=_none_or_dtype(attn_mask),\n            other_name=\"attn_mask\",\n            target_type=query.dtype\n        )\n\n        attn_mask = _canonical_mask(\n            mask=attn_mask,\n            mask_name=\"attn_mask\",\n            other_type=None,\n            other_name=\"\",\n            target_type=query.dtype,\n            check_other=False,\n        )\n\n        is_fastpath_enabled = get_fastpath_enabled()\n\n        if not is_fastpath_enabled:\n            why_not_fast_path = \"torch.backends.mha.get_fastpath_enabled() was not True\"\n        elif not is_batched:\n            why_not_fast_path = f\"input not batched; expected query.dim() of 3 but got {query.dim()}\"\n        elif query is not key or key is not value:\n            # When lifting this restriction, don't forget to either\n            # enforce that the dtypes all match or test cases where\n            # they don't!\n            why_not_fast_path = \"non-self attention was used (query, key, and value are not the same Tensor)\"\n        elif self.in_proj_bias is not None and query.dtype != self.in_proj_bias.dtype:\n            why_not_fast_path = f\"dtypes of query ({query.dtype}) and self.in_proj_bias ({self.in_proj_bias.dtype}) don't match\"\n        elif self.in_proj_weight is None:\n            why_not_fast_path = \"in_proj_weight was None\"\n        elif query.dtype != self.in_proj_weight.dtype:\n            # this case will fail anyway, but at least they'll get a useful error message.\n            why_not_fast_path = f\"dtypes of query ({query.dtype}) and self.in_proj_weight ({self.in_proj_weight.dtype}) don't match\"\n        elif self.training:\n            why_not_fast_path = \"training is enabled\"\n        elif (self.num_heads % 2) != 0:\n            why_not_fast_path = \"self.num_heads is not even\"\n        elif not self.batch_first:\n            why_not_fast_path = \"batch_first was not True\"\n        elif self.bias_k is not None:\n            why_not_fast_path = \"self.bias_k was not None\"\n        elif self.bias_v is not None:\n            why_not_fast_path = \"self.bias_v was not None\"\n        elif self.add_zero_attn:\n            why_not_fast_path = \"add_zero_attn was enabled\"\n        elif not self._qkv_same_embed_dim:\n            why_not_fast_path = \"_qkv_same_embed_dim was not True\"\n        elif query.is_nested and (key_padding_mask is not None or attn_mask is not None):\n            why_not_fast_path = \"supplying both src_key_padding_mask and src_mask at the same time \\\n                                 is not supported with NestedTensor input\"\n        elif torch.is_autocast_enabled():\n            why_not_fast_path = \"autocast is enabled\"\n\n        if not why_not_fast_path:\n            tensor_args = (\n                query,\n                key,\n                value,\n                self.in_proj_weight,\n                self.in_proj_bias,\n                self.out_proj.weight,\n                self.out_proj.bias,\n            )\n            # We have to use list comprehensions below because TorchScript does not support\n            # generator expressions.\n            if torch.overrides.has_torch_function(tensor_args):\n                why_not_fast_path = \"some Tensor argument has_torch_function\"\n            elif _is_make_fx_tracing():\n                why_not_fast_path = \"we are running make_fx tracing\"\n            elif not all(_check_arg_device(x) for x in tensor_args):\n                why_not_fast_path = (\"some Tensor argument's device is neither one of \"\n                                     f\"cpu or cuda\")\n            elif torch.is_grad_enabled() and any(_arg_requires_grad(x) for x in tensor_args):\n                why_not_fast_path = (\"grad is enabled and at least one of query or the \"\n                                     \"input/output projection weights or biases requires_grad\")\n            if not why_not_fast_path:\n                merged_mask, mask_type = self.merge_masks(attn_mask, key_padding_mask, query)\n\n                if self.in_proj_bias is not None and self.in_proj_weight is not None:\n                    return torch._native_multi_head_attention(\n                        query,\n                        key,\n                        value,\n                        self.embed_dim,\n                        self.num_heads,\n                        self.in_proj_weight,\n                        self.in_proj_bias,\n                        self.out_proj.weight,\n                        self.out_proj.bias,\n                        merged_mask,\n                        need_weights,\n                        average_attn_weights,\n                        mask_type)\n\n        any_nested = query.is_nested or key.is_nested or value.is_nested\n        assert not any_nested, (\"MultiheadAttention does not support NestedTensor outside of its fast path. \" +\n                                f\"The fast path was not hit because {why_not_fast_path}\")\n\n        if self.batch_first and is_batched:\n            # make sure that the transpose op does not affect the \"is\" property\n            if key is value:\n                if query is key:\n                    query = key = value = query.transpose(1, 0)\n                else:\n                    query, key = (x.transpose(1, 0) for x in (query, key))\n                    value = key\n            else:\n                query, key, value = (x.transpose(1, 0) for x in (query, key, value))\n\n        if not self._qkv_same_embed_dim:\n            attn_output, attn_output_weights = multi_head_attention_forward(#multi_head_attention_forward(#F.multi_head_attention_forward(\n                query, key, value, self.embed_dim, self.num_heads,\n                self.in_proj_weight, self.in_proj_bias,\n                self.bias_k, self.bias_v, self.add_zero_attn,\n                self.dropout, self.out_proj.weight, self.out_proj.bias,\n                training=self.training,\n                key_padding_mask=key_padding_mask, need_weights=need_weights,\n                attn_mask=attn_mask,\n                use_separate_proj_weight=True,\n                q_proj_weight=self.q_proj_weight, k_proj_weight=self.k_proj_weight,\n                v_proj_weight=self.v_proj_weight)#,\n                #average_attn_weights=average_attn_weights,\n                #is_causal=is_causal)\n        else:\n#           print(self.nbh_B)\n#           print('SANITY CHECK')\n            attn_output, attn_output_weights = multi_head_attention_forward(#multi_head_attention_forward(#F.multi_head_attention_forward(\n                query, key, value,\n                self.nbh_A[:,AM],self.nbh_B[:,AM],#self.nbh_A[:,AM],\n                self.embed_dim, self.num_heads,\n                self.in_proj_weight, self.in_proj_bias,\n                self.bias_k, self.bias_v, self.add_zero_attn,\n                self.dropout, self.out_proj.weight, self.out_proj.bias,\n                training=self.training,\n                key_padding_mask=key_padding_mask,\n                need_weights=need_weights,\n                attn_mask=attn_mask)#,\n                #average_attn_weights=average_attn_weights,\n                #is_causal=is_causal)\n        if self.batch_first and is_batched:\n            return attn_output.transpose(1, 0), attn_output_weights\n        else:\n            return attn_output, attn_output_weights\n\n    def merge_masks(self, attn_mask: Optional[Tensor], key_padding_mask: Optional[Tensor],\n                    query: Tensor) -> Tuple[Optional[Tensor], Optional[int]]:\n        r\"\"\"Determine mask type and combine masks if necessary.\n\n        If only one mask is provided, that mask\n        and the corresponding mask type will be returned. If both masks are provided, they will be both\n        expanded to shape ``(batch_size, num_heads, seq_len, seq_len)``, combined with logical ``or``\n        and mask type 2 will be returned\n        Args:\n            attn_mask: attention mask of shape ``(seq_len, seq_len)``, mask type 0\n            key_padding_mask: padding mask of shape ``(batch_size, seq_len)``, mask type 1\n            query: query embeddings of shape ``(batch_size, seq_len, embed_dim)``\n        Returns:\n            merged_mask: merged mask\n            mask_type: merged mask type (0, 1, or 2)\n        \"\"\"\n        mask_type: Optional[int] = None\n        merged_mask: Optional[Tensor] = None\n\n        if key_padding_mask is not None:\n            mask_type = 1\n            merged_mask = key_padding_mask\n\n        if attn_mask is not None:\n            # In this branch query can't be a nested tensor, so it has a shape\n            batch_size, seq_len, _ = query.shape\n            mask_type = 2\n\n            # Always expands attn_mask to 4D\n            if attn_mask.dim() == 3:\n                attn_mask_expanded = attn_mask.view(batch_size, -1, seq_len, seq_len)\n            else:  # attn_mask.dim() == 2:\n                attn_mask_expanded = attn_mask.view(1, 1, seq_len, seq_len).expand(batch_size, self.num_heads, -1, -1)\n            merged_mask = attn_mask_expanded\n\n            if key_padding_mask is not None:\n                key_padding_mask_expanded = key_padding_mask.view(batch_size, 1, 1, seq_len).expand(-1, self.num_heads, -1, -1)\n                merged_mask = attn_mask_expanded + key_padding_mask_expanded\n\n        # no attn_mask and no key_padding_mask, returns None, None\n        return merged_mask, mask_type","metadata":{"id":"kpez3YG8qeAz"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://pytorch.org/docs/stable/_modules/torch/nn/modules/transformer.html#TransformerEncoderLayer\nLinear = torch.nn.Linear\nDropout = torch.nn.Dropout\n\n_shape_t = Union[int, List[int], Size]\n\nclass LayerNorm(Module):\n\n    __constants__ = ['normalized_shape', 'eps', 'elementwise_affine']\n    normalized_shape: Tuple[int, ...]\n    eps: float\n    elementwise_affine: bool\n\n    def __init__(self, normalized_shape: _shape_t, eps: float = 1e-5, elementwise_affine: bool = True,\n                 bias: bool = True, device=None, dtype=None) -> None:\n        factory_kwargs = {'device': device, 'dtype': dtype}\n        super().__init__()\n        if isinstance(normalized_shape, numbers.Integral):\n            # mypy error: incompatible types in assignment\n            normalized_shape = (normalized_shape,)  # type: ignore[assignment]\n        self.normalized_shape = tuple(normalized_shape)  # type: ignore[arg-type]\n        self.eps = eps\n        self.elementwise_affine = elementwise_affine\n        if self.elementwise_affine:\n            self.weight = Parameter(torch.empty(self.normalized_shape, **factory_kwargs))\n            if bias:\n                self.bias = Parameter(torch.empty(self.normalized_shape, **factory_kwargs))\n            else:\n                self.register_parameter('bias', None)\n        else:\n            self.register_parameter('weight', None)\n            self.register_parameter('bias', None)\n\n        self.reset_parameters()\n\n    def reset_parameters(self) -> None:\n        if self.elementwise_affine:\n            init.ones_(self.weight)\n            if self.bias is not None:\n                init.zeros_(self.bias)\n\n    def forward(self, input: Tensor) -> Tensor:\n        return F.layer_norm(\n            input, self.normalized_shape, self.weight, self.bias, self.eps)\n\n    def extra_repr(self) -> str:\n        return '{normalized_shape}, eps={eps}, ' \\\n            'elementwise_affine={elementwise_affine}'.format(**self.__dict__)\n\ndef _get_activation_fn(activation: str) -> Callable[[Tensor], Tensor]:\n    if activation == \"relu\":\n        return F.relu\n    elif activation == \"gelu\":\n        return F.gelu\n\n    raise RuntimeError(f\"activation should be relu/gelu, not {activation}\")\n\nclass TransformerEncoderLayer(Module):\n\n    __constants__ = ['norm_first']\n\n    def __init__(self, d_model: int, nhead: int, dim_feedforward: int = 2048, dropout: float = 0.1,\n                 activation: Union[str, Callable[[Tensor], Tensor]] = F.relu,\n                 layer_norm_eps: float = 1e-5, batch_first: bool = False, norm_first: bool = False,\n                 bias: bool = True, device=None, dtype=None) -> None:\n        factory_kwargs = {'device': device, 'dtype': dtype}\n        super().__init__()\n        self.self_attn = MultiheadAttention(d_model, nhead, dropout=dropout,\n                                            bias=bias, batch_first=batch_first,\n                                            **factory_kwargs)\n        # Implementation of Feedforward model\n        self.linear1 = Linear(d_model, dim_feedforward, bias=bias, **factory_kwargs)\n        self.dropout = Dropout(dropout)\n        self.linear2 = Linear(dim_feedforward, d_model, bias=bias, **factory_kwargs)\n\n        self.norm_first = norm_first\n        self.norm1 = LayerNorm(d_model, eps=layer_norm_eps, bias=bias, **factory_kwargs)\n        self.norm2 = LayerNorm(d_model, eps=layer_norm_eps, bias=bias, **factory_kwargs)\n        self.dropout1 = Dropout(dropout)\n        self.dropout2 = Dropout(dropout)\n\n        # Legacy string support for activation function.\n        if isinstance(activation, str):\n            activation = _get_activation_fn(activation)\n\n        # We can't test self.activation in forward() in TorchScript,\n        # so stash some information about it instead.\n        if activation is F.relu or isinstance(activation, torch.nn.ReLU):\n            self.activation_relu_or_gelu = 1\n        elif activation is F.gelu or isinstance(activation, torch.nn.GELU):\n            self.activation_relu_or_gelu = 2\n        else:\n            self.activation_relu_or_gelu = 0\n        self.activation = activation\n\n    def __setstate__(self, state):\n        super().__setstate__(state)\n        if not hasattr(self, 'activation'):\n            self.activation = F.relu\n\n\n    def forward(\n            self,\n            src: Tensor,\n            AM: Tensor,\n            src_mask: Optional[Tensor] = None,\n            src_key_padding_mask: Optional[Tensor] = None,\n            is_causal: bool = False) -> Tensor:\n\n        src_key_padding_mask = _canonical_mask(\n            mask=src_key_padding_mask,\n            mask_name=\"src_key_padding_mask\",\n            other_type=_none_or_dtype(src_mask),\n            other_name=\"src_mask\",\n            target_type=src.dtype\n        )\n\n        src_mask = _canonical_mask(\n            mask=src_mask,\n            mask_name=\"src_mask\",\n            other_type=None,\n            other_name=\"\",\n            target_type=src.dtype,\n            check_other=False,\n        )\n\n        is_fastpath_enabled = get_fastpath_enabled()\n\n        # see Fig. 1 of https://arxiv.org/pdf/2002.04745v1.pdf\n        why_not_sparsity_fast_path = ''\n        if not is_fastpath_enabled:\n            why_not_sparsity_fast_path = \"torch.backends.mha.get_fastpath_enabled() was not True\"\n        elif not src.dim() == 3:\n            why_not_sparsity_fast_path = f\"input not batched; expected src.dim() of 3 but got {src.dim()}\"\n        elif self.training:\n            why_not_sparsity_fast_path = \"training is enabled\"\n        elif not self.self_attn.batch_first:\n            why_not_sparsity_fast_path = \"self_attn.batch_first was not True\"\n        elif self.self_attn.in_proj_bias is None:\n            why_not_sparsity_fast_path = \"self_attn was passed bias=False\"\n        elif not self.self_attn._qkv_same_embed_dim:\n            why_not_sparsity_fast_path = \"self_attn._qkv_same_embed_dim was not True\"\n        elif not self.activation_relu_or_gelu:\n            why_not_sparsity_fast_path = \"activation_relu_or_gelu was not True\"\n        elif not (self.norm1.eps == self.norm2.eps):\n            why_not_sparsity_fast_path = \"norm1.eps is not equal to norm2.eps\"\n        elif src.is_nested and (src_key_padding_mask is not None or src_mask is not None):\n            why_not_sparsity_fast_path = \"neither src_key_padding_mask nor src_mask are not supported with NestedTensor input\"\n        elif self.self_attn.num_heads % 2 == 1:\n            why_not_sparsity_fast_path = \"num_head is odd\"\n        elif torch.is_autocast_enabled():\n            why_not_sparsity_fast_path = \"autocast is enabled\"\n        if not why_not_sparsity_fast_path:\n            tensor_args = (\n                src,\n                self.self_attn.in_proj_weight,\n                self.self_attn.in_proj_bias,\n                self.self_attn.out_proj.weight,\n                self.self_attn.out_proj.bias,\n                self.norm1.weight,\n                self.norm1.bias,\n                self.norm2.weight,\n                self.norm2.bias,\n                self.linear1.weight,\n                self.linear1.bias,\n                self.linear2.weight,\n                self.linear2.bias,\n            )\n\n            # We have to use list comprehensions below because TorchScript does not support\n            # generator expressions.\n            _supported_device_type = [\"cpu\", \"cuda\"]\n            if torch.overrides.has_torch_function(tensor_args):\n                why_not_sparsity_fast_path = \"some Tensor argument has_torch_function\"\n            elif not all((x.device.type in _supported_device_type) for x in tensor_args):\n                why_not_sparsity_fast_path = (\"some Tensor argument's device is neither one of \"\n                                              f\"{_supported_device_type}\")\n            elif torch.is_grad_enabled() and any(x.requires_grad for x in tensor_args):\n                why_not_sparsity_fast_path = (\"grad is enabled and at least one of query or the \"\n                                              \"input/output projection weights or biases requires_grad\")\n            '''\n            if not why_not_sparsity_fast_path:\n                merged_mask, mask_type = self.self_attn.merge_masks(src_mask, src_key_padding_mask, src)\n                return torch._transformer_encoder_layer_fwd(\n                    src,\n                    self.self_attn.embed_dim,\n                    self.self_attn.num_heads,\n                    self.self_attn.in_proj_weight,\n                    self.self_attn.in_proj_bias,\n                    self.self_attn.out_proj.weight,\n                    self.self_attn.out_proj.bias,\n                    self.activation_relu_or_gelu == 2,\n                    self.norm_first,\n                    self.norm1.eps,\n                    self.norm1.weight,\n                    self.norm1.bias,\n                    self.norm2.weight,\n                    self.norm2.bias,\n                    self.linear1.weight,\n                    self.linear1.bias,\n                    self.linear2.weight,\n                    self.linear2.bias,\n                    merged_mask,\n                    mask_type,\n                )\n            '''\n\n        x = src\n        if self.norm_first:\n            x = x + self._sa_block(self.norm1(x), AM, src_mask, src_key_padding_mask, is_causal=is_causal)\n            x = x + self._ff_block(self.norm2(x))\n        else:\n            x = self.norm1(x + self._sa_block(x, AM, src_mask, src_key_padding_mask, is_causal=is_causal))\n            x = self.norm2(x + self._ff_block(x))\n\n        return x\n\n\n    # self-attention block\n    def _sa_block(self, x: Tensor,\n                  AM: Tensor,\n                  attn_mask: Optional[Tensor], key_padding_mask: Optional[Tensor], is_causal: bool = False) -> Tensor:\n        x = self.self_attn(x, x, x,\n                           AM,\n                           attn_mask=attn_mask,\n                           key_padding_mask=key_padding_mask,\n                           need_weights=False, is_causal=is_causal)[0]\n        return self.dropout1(x)\n\n    # feed forward block\n    def _ff_block(self, x: Tensor) -> Tensor:\n        x = self.linear2(self.dropout(self.activation(self.linear1(x))))\n        return self.dropout2(x)","metadata":{"id":"lGorkcw7qeA0"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class myGNN_ENCODER(nn.Module):\n    def __init__(self):\n        super(myGNN_ENCODER, self).__init__()\n\n        self.GNN_0 = nn.Sequential(\n            myGATConv(27,64),\n            myGATConv(64,64),\n            myGATConv(64,64),\n            myGATConv(64,64),\n            myGATConv(64,64)\n        )\n        self.GNN_1 = nn.Sequential(\n            myGATConv(27+64,128),\n            myGATConv(128,128),\n            myGATConv(128,128),\n            myGATConv(128,128),\n            myGATConv(128,128)\n        )\n        self.GNN_2 = nn.Sequential(\n            myGATConv(27+64+128,256),\n            myGATConv(256,256),\n            myGATConv(256,256),\n            myGATConv(256,256),\n            myGATConv(256,256)\n        )\n        self.emb = nn.Linear(475,512).to(device)\n        self.TransformerEncoderLayer = TransformerEncoderLayer(d_model=512, nhead=4, batch_first=True).to(device)#, dim_feedforward=512\n\n    def forward(self,x,AM,mask,edges,edge_attr):\n\n        b0,edges,edge_attr = self.GNN_0((x,edges,edge_attr))\n        b1,edges,edge_attr = self.GNN_1((torch.cat([b0,x],-1),edges,edge_attr))\n        b2,edges,edge_attr = self.GNN_2((torch.cat([b1,b0,x],-1),edges,edge_attr))\n\n        x = self.emb(torch.cat([b2,b1,b0,x],-1))\n\n        x = self.TransformerEncoderLayer(x.view(-1,Lmax,512),AM,src_key_padding_mask=mask.view(-1,Lmax)).view(-1,512)#[~mask]\n\n        return x","metadata":{"id":"J9gsbZ8WqeA0"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class myGNN(nn.Module):\n    def __init__(self):\n        super(myGNN, self).__init__()\n\n        self.ENCODER = myGNN_ENCODER()\n        self.DROPOUT = nn.Dropout(DROPOUT).to(device)\n        self.OUT = nn.Linear(512,2).to(device)\n        self.OUT.weight = nn.Parameter(torch.ones(2,512)/512)\n        self.OUT.bias = nn.Parameter(torch.zeros(2))\n\n    def forward(self,X):\n        x,AM,mask,edges,edge_attr = X\n\n        x = self.ENCODER(x,AM,mask,edges,edge_attr)\n        OUT = torch.zeros(len(x),device=device)\n        OUT[~mask] = torch.softmax(self.OUT(x[~mask]),-1)[:,1]\n        OUT = nn.MaxPool1d(Lmax)(OUT.view(-1,Lmax))\n\n        return OUT","metadata":{"id":"ur9PrtZNqeA1"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(2024)\nSEEDS = []\nfor f in range(20):\n    SEED = str(np.random.randint(9)+1)\n    for _ in range(4):\n        SEED = SEED + str(np.random.randint(10))\n    SEEDS.append(int(SEED))\n\nSEEDS","metadata":{"executionInfo":{"elapsed":3,"status":"ok","timestamp":1720312381744,"user":{"displayName":"Angel Sanchez","userId":"12587302736694913189"},"user_tz":-120},"id":"tKNzz40qqeA1","outputId":"601dae20-88ba-4392-d38f-d4c3a8f8b830"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for f in FOLDS:\n    print(f)\n    seed_everything(2024)\n    model = myGNN()\n    model.ENCODER = torch.load('C:/Users/Angel/kaggle/ENCODER',map_location=torch.device(device))\n    seed_everything(SEEDS[ff])\n    tds = BELKA_DS(train,f,ff)\n    seed_everything(SEEDS[ff])\n    vds = BELKA_DS(train,f,ff,VALID=True)\n\n    tdl = torch.utils.data.DataLoader(\n        tds,\n        collate_fn=my_collate_fn,\n        batch_size=BS\n    )\n    vdl = torch.utils.data.DataLoader(\n        vds,\n        collate_fn=my_collate_fn,\n        batch_size=BS\n    )\n    dls = DataLoaders(tdl,vdl)\n\n    learn = Learner(\n        dls,\n        model,\n        lr=LR,\n        loss_func=nn.BCELoss(),\n        cbs=[\n            ShowGraphCallback()\n        ]\n    )\n\n    learn.fit_one_cycle(EPOCHS)\n\n    torch.save(model,'C:/Users/Angel/kaggle/'+T[6:]+'/models/'+T[6:]+'_V2_'+str(SEEDS[f])+'_'+str(f)+'_'+str(ff))\n\n    for w in model.ENCODER.TransformerEncoderLayer.self_attn.nbh_A:\n        plt.plot(w.cpu().detach())\n    plt.show()\n\n    for w in model.ENCODER.TransformerEncoderLayer.self_attn.nbh_B:\n        plt.plot(w.cpu().detach())\n    plt.show()\n\n    with torch.no_grad():\n        for b in vdl:\n            preds = (model(b[0]) > .5).float()\n\n    print(sklearn.metrics.confusion_matrix(b[1].cpu(),preds.cpu()))\n\n    del model,tds,vds,tdl,vdl,dls,learn\n    gc.collect()","metadata":{"id":"wsHk-VfCqeA1","outputId":"72c238fd-20ff-4e9d-d645-9df0526612b0"},"execution_count":null,"outputs":[]}]}