{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"}],"dockerImageVersionId":30746,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"See discussion at : https://www.kaggle.com/competitions/leash-BELKA/discussion/521894\n\nThis is code snippet (i.e. not full running code).  \nIt uses functions from my solutiuon at  \nhttps://github.com/hengck23/solution-leash-BELKA  \n\n---\n\nNext, you need the following:  \nGradCAM : \n- https://github.com/jacobgil/pytorch-grad-cam  \n\nXSMILES :\n- https://github.com/Bayer-Group/xsmiles  \n- https://github.com/Bayer-Group/xsmiles-use-cases \n\nPlease follow their documents and tutorials first to familiarize with their API and check for correct installation.\n\n\nprint-dict: \n- https://pypi.org/project/print-dict/\n\nThis print dict for json.dumps.","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"## 1. Helper for generating gradCAM","metadata":{}},{"cell_type":"code","source":"%%script false --no-raise-error\n\n\nfrom pytorch_grad_cam import GradCAM\n\n#https://github.com/jacobgil/pytorch-grad-cam/issues/410\nclass MultiLabelBinaryClassifierOutputTarget:\n    def __init__(self, output_index, category):\n        self.category = category\n        self.output_index = output_index\n\n    def __call__(self, model_output):\n        print('MultiLabelBinaryClassifierOutputTarget(): model_output', model_output.shape)\n        if self.category == 1:\n            sign = 1\n        else:\n            sign = -1\n        return model_output[self.output_index] * sign\n    \n    \n\n## wrap your model to CAM model\n#   - this has to be hand-coded. It is different for different net\n#   - please refer to: https://github.com/hengck23/solution-leash-BELKA/blob/main/src/cnn1d-nonshare-05-mean-layer5-bn/model.py\n#     for the cnn1d net being warp here   \n#   \n# gradCAM only support image(2D) and volume(3D)\n# we convert 1d seq to 2d image with height=1, width=seq length\nclass ToCam(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n    def forward(self, x):\n        B, DIM, L  = x.shape  # B, L, DIM :transformer\n        #x = x.permute(0, 2, 1).reshape(B, DIM, 1, L).contiguous() #transformer\n        x = x.reshape(B, DIM, 1, L)\n        return x\n\nclass FromCam(nn.Module):\n    def __init__(self, ):\n        super().__init__()\n    def forward(self, x):\n        B, DIM, _1_, L = x.shape\n        #x = x.permute(0, 3, 1, 2).reshape(B, L, DIM).contiguous() #transformer\n        x = x.reshape(B, DIM, L)\n        return x\n\n#convert to cam model\nclass CamModel(nn.Module):\n    def __init__(self, net):\n        super().__init__()\n        self.net=net\n        self.to_cam = ToCam()\n        self.from_cam = FromCam()\n\n    def forward(self, cam_input):\n        smiles_token_id = cam_input.reshape(1,-1)\n        \n        net = self.net\n        x = net.embedding(smiles_token_id)\n        x = x.permute(0, 2, 1).contiguous()\n        x = net.token_encoder(x)\n        x = self.to_cam(x)\n        x = self.from_cam(x)\n        last = F.adaptive_avg_pool1d(x, 1).squeeze(-1)\n        bind = net.bind(last)\n\n        probability = torch.sigmoid(bind)\n        return probability\n\nprint('CAM helper functions OK!!!')","metadata":{"execution":{"iopub.status.busy":"2024-07-28T21:12:19.659970Z","iopub.execute_input":"2024-07-28T21:12:19.661226Z","iopub.status.idle":"2024-07-28T21:12:19.671897Z","shell.execute_reply.started":"2024-07-28T21:12:19.661182Z","shell.execute_reply":"2024-07-28T21:12:19.670920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2. Run gradCAM and make heatmap array","metadata":{}},{"cell_type":"code","source":"%%script false --no-raise-error\n \n\nfrom configure import default_cfg\n\ncfg = deepcopy(default_cfg)\ncfg.fold = 2\ncfg.fold_dir = f'{RESULT_DIR}/{cfg.experiment_name}/fold-{cfg.fold}'\ncfg.resume_from.checkpoint = f'{RESULT_DIR}/{cfg.experiment_name}/fold-{cfg.fold}/checkpoint/00430000.pth'\n\n\n\n\n### start here! ################################################\nif 0: #dummy data\n    smiles_token_id=np.random.choice(cfg.VOCAB_SIZE, (1, cfg.MAX_LENGTH))\n    bind = np.random.choice(2, (1, 3))\n    \n    \nif 1:\n    token_id, bind = load_kaggle_data()\n    index=21821828 #97116270 #21821828\n    smiles_token_id=token_id[[index]]\n    bind = bind[[index]]\n\n\nsmiles_token_id=torch.from_numpy(smiles_token_id).long().cuda()\nbind = torch.from_numpy(bind).float().cuda()\nprint('smiles_token_id', smiles_token_id.shape)\n\n\n#----\n# setup net\nnet = Net(cfg)\nif cfg.resume_from.checkpoint is not None:\n    f = torch.load(cfg.resume_from.checkpoint, map_location=lambda storage, loc: storage)\n    state_dict = f['state_dict']\n    print(net.load_state_dict(state_dict, strict=False))  # True\n\nnet.cuda()\nnet.output_type = ['infer']\n\n#PROTEIN_NAME=['BRD4', 'HSA', 'sEH']\n_BRD4_, _HSA_, _sEH_ = 0,1,2\n\n\n#----\n# setup cam\ncam_model = CamModel(net)\ncam_layer = [cam_model.to_cam]\ncam_input = smiles_token_id.reshape(1,1,1,-1)\ncam = GradCAM(model=cam_model, target_layers=cam_layer)\ncam_target = [MultiLabelBinaryClassifierOutputTarget(_sEH_,1)]\n\n# run one input\nheatmap = cam(input_tensor=cam_input, targets=cam_target)\nprint('heatmap')\nprint(heatmap.shape)\nprint(heatmap)\nprint('outputs')\nprint(cam.outputs) #check that ouput is the same as orginal net\nprint('')","metadata":{"execution":{"iopub.status.busy":"2024-07-28T21:12:19.673867Z","iopub.execute_input":"2024-07-28T21:12:19.674683Z","iopub.status.idle":"2024-07-28T21:12:19.691942Z","shell.execute_reply.started":"2024-07-28T21:12:19.674652Z","shell.execute_reply":"2024-07-28T21:12:19.690962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Convert heatmap array to XSMILES dict","metadata":{}},{"cell_type":"code","source":"%%script false --no-raise-error\n\n#convert to XSMILES format\nheatmap = heatmap[0,0].tolist()\nbind = bind.data.cpu().numpy()[0].tolist()\npredict = cam.outputs.data.cpu().numpy()[0].tolist()\n\n#load your SMILE strings\nsmiles_string = load_compressed_bz2_pickle(\n    f'/home/hp/work/2024/kaggle/leash-belka/data/processed/smiles-string/train-replace-c.smiles.bytestring.bz2'\n)\nstring = smiles_string[index].decode()\n\n\nmolecule={\n    'string':string,\n    'sequence' : [ s for s in string ],\n    'methods': [{\n        'name': f'cnn1d for molecule {index}',\n        'scores': heatmap[1:len(string)+1], #remove BOS,EOS and PAD token ....\n    }],\n    'attributes': { \n        'bind': str(bind),\n        'predict': str(predict),\n        'index': index,\n    },\n}\nprint('cnn1d')\nprint(index)\nprint(format_dict(molecule))\n\n'''\n#you should see this printed\n{\n    'string': 'CCN(CCCNc1nc(NCC(C)S(=O)(=O)N2CCN(c3ccccc3)CC2)nc(N[C@@H](CC(=O)NC)Cc2ccc(C#N)cc2)n1)S(C)(=O)=O',\n    'sequence': [\n        'C', 'C', 'N', '(', 'C', 'C', 'C', 'N', 'c', '1', 'n', 'c', '(', 'N', 'C', 'C', '(',\n        'C', ')', 'S', '(', '=', 'O', ')', '(', '=', 'O', ')', 'N', '2', 'C', 'C', 'N', '(',\n        'c', '3', 'c', 'c', 'c', 'c', 'c', '3', ')', 'C', 'C', '2', ')', 'n', 'c', '(', 'N',\n        '[', 'C', '@', '@', 'H', ']', '(', 'C', 'C', '(', '=', 'O', ')', 'N', 'C', ')', 'C',\n        'c', '2', 'c', 'c', 'c', '(', 'C', '#', 'N', ')', 'c', 'c', '2', ')', 'n', '1', ')',\n        'S', '(', 'C', ')', '(', '=', 'O', ')', '=', 'O'\n    ],\n    'methods': [{\n        'name': 'cnn1d for molecule 21821828',\n        'scores': [\n            0.0, 0.0, 0.07322029024362564, 0.2863195836544037, 0.0005969193880446255, 0.0, 0.0,\n            0.0, 0.0164840929210186, 0.1949203759431839, 0.0, 0.0, 0.011598554439842701,\n            0.38833820819854736, 0.09658758342266083, 0.0, 0.033475764095783234,\n            0.3708333969116211, 0.9999998807907104, 0.1468649059534073, 0.0,\n            0.01157574076205492, 0.05313240736722946, 0.029402032494544983, 0.02064199186861515,\n            0.02789340540766716, 0.05652434751391411, 0.08874323219060898, 0.0, 0.0,\n            0.05295601114630699, 0.029905853793025017, 0.0, 0.036882173269987106,\n            0.07905981689691544, 0.06992898136377335, 0.23845051229000092, 0.05124058201909065,\n            0.005834367126226425, 0.0, 0.0, 0.3467688262462616, 0.0, 0.049468256533145905, 0.0,\n            0.11993400007486343, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.05276830121874809,\n            0.2737159729003906, 0.39585331082344055, 0.0915205329656601, 0.0, 0.0, 0.0, 0.0,\n            0.0, 0.0, 0.0, 0.0, 0.0, 0.2772041857242584, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,\n            0.0, 0.16493263840675354, 0.018042370676994324, 0.0, 0.0, 0.03622734174132347,\n            0.04259834066033363, 0.0, 0.0, 0.0, 0.0, 0.0, 0.002740263706073165, 0.0, 0.0, 0.0,\n            0.0, 0.030052833259105682, 0.0, 0.0\n        ]\n    }],\n    'attributes': {\n        'bind': '[0.0, 0.0, 1.0]',\n        'predict': '[5.0749404181260616e-06, 2.471163861628156e-05, 0.9467632174491882]',\n        'index': 21821828\n    }\n}\n\n'''","metadata":{"execution":{"iopub.status.busy":"2024-07-28T21:12:19.693530Z","iopub.execute_input":"2024-07-28T21:12:19.694037Z","iopub.status.idle":"2024-07-28T21:12:19.708943Z","shell.execute_reply.started":"2024-07-28T21:12:19.693996Z","shell.execute_reply":"2024-07-28T21:12:19.707912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Show XSMILES (js script)","metadata":{}},{"cell_type":"code","source":"%%script false --no-raise-error\n\n#copy the printed text above and write the code:\n#\n# if you want to improve contrast, use:\n#  'scores': (np.array([ 0.0, 0.0, 0.07322029024362564, ...])])**7).tolist()\n#\n#\nxmol = {\n    'string': 'CCN(CCCNc1nc(NCC(C)S(=O)(=O)N2CCN(c3ccccc3)CC2)nc(N[C@@H](CC(=O)NC)Cc2ccc(C#N)cc2)n1)S(C)(=O)=O',\n    'sequence': [\n        'C', 'C', 'N', '(', 'C', 'C', 'C', 'N', 'c', '1', 'n', 'c', '(', 'N', 'C', 'C', '(',\n        'C', ')', 'S', '(', '=', 'O', ')', '(', '=', 'O', ')', 'N', '2', 'C', 'C', 'N', '(',\n        'c', '3', 'c', 'c', 'c', 'c', 'c', '3', ')', 'C', 'C', '2', ')', 'n', 'c', '(', 'N',\n        '[', 'C', '@', '@', 'H', ']', '(', 'C', 'C', '(', '=', 'O', ')', 'N', 'C', ')', 'C',\n        'c', '2', 'c', 'c', 'c', '(', 'C', '#', 'N', ')', 'c', 'c', '2', ')', 'n', '1', ')',\n        'S', '(', 'C', ')', '(', '=', 'O', ')', '=', 'O'\n    ],\n    'methods': [{\n        'name': 'cnn1d for molecule 21821828',\n        'scores': [\n            0.0, 0.0, 0.07322029024362564, 0.2863195836544037, 0.0005969193880446255, 0.0, 0.0,\n            0.0, 0.0164840929210186, 0.1949203759431839, 0.0, 0.0, 0.011598554439842701,\n            0.38833820819854736, 0.09658758342266083, 0.0, 0.033475764095783234,\n            0.3708333969116211, 0.9999998807907104, 0.1468649059534073, 0.0,\n            0.01157574076205492, 0.05313240736722946, 0.029402032494544983, 0.02064199186861515,\n            0.02789340540766716, 0.05652434751391411, 0.08874323219060898, 0.0, 0.0,\n            0.05295601114630699, 0.029905853793025017, 0.0, 0.036882173269987106,\n            0.07905981689691544, 0.06992898136377335, 0.23845051229000092, 0.05124058201909065,\n            0.005834367126226425, 0.0, 0.0, 0.3467688262462616, 0.0, 0.049468256533145905, 0.0,\n            0.11993400007486343, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.05276830121874809,\n            0.2737159729003906, 0.39585331082344055, 0.0915205329656601, 0.0, 0.0, 0.0, 0.0,\n            0.0, 0.0, 0.0, 0.0, 0.0, 0.2772041857242584, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,\n            0.0, 0.16493263840675354, 0.018042370676994324, 0.0, 0.0, 0.03622734174132347,\n            0.04259834066033363, 0.0, 0.0, 0.0, 0.0, 0.0, 0.002740263706073165, 0.0, 0.0, 0.0,\n            0.0, 0.030052833259105682, 0.0, 0.0\n        ]\n    }],\n    'attributes': {\n        'bind': '[0.0, 0.0, 1.0]',\n        'predict': '[5.0749404181260616e-06, 2.471163861628156e-05, 0.9467632174491882]',\n        'index': 21821828\n    }\n}\n\nw=xsmiles.XSmilesWidget(molecules=json.dumps([xmol])) # create xsmiles object\\n\",\nw\n    \n#this should display XSMILES heatmap diagram","metadata":{"execution":{"iopub.status.busy":"2024-07-28T21:12:19.711120Z","iopub.execute_input":"2024-07-28T21:12:19.711488Z","iopub.status.idle":"2024-07-28T21:12:19.726931Z","shell.execute_reply.started":"2024-07-28T21:12:19.711453Z","shell.execute_reply":"2024-07-28T21:12:19.725852Z"},"trusted":true},"execution_count":null,"outputs":[]}]}