{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from IPython.display import clear_output as clr\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\nfrom transformers import AutoModelForMaskedLM, AutoTokenizer\n\nchemberta = AutoModelForMaskedLM.from_pretrained(\"DeepChem/ChemBERTa-77M-MTR\")\ntokenizer = AutoTokenizer.from_pretrained(\"DeepChem/ChemBERTa-77M-MTR\")\n\nchemberta.eval()\ndef featurize_ChemBERTa(smiles_list, padding=True):\n    embeddings_cls = torch.zeros(len(smiles_list), 600)\n    embeddings_mean = torch.zeros(len(smiles_list), 600)\n\n    with torch.no_grad():\n        for i, smiles in enumerate(tqdm(smiles_list)):\n            encoded_input = tokenizer(smiles, return_tensors=\"pt\",padding=padding,truncation=True)\n            model_output = chemberta(**encoded_input)\n            \n            embedding = model_output[0][::,0,::]\n            embeddings_cls[i] = embedding\n            \n            embedding = torch.mean(model_output[0],1)\n            embeddings_mean[i] = embedding\n            \n    return embeddings_cls.numpy(), embeddings_mean.numpy()\nclr()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-19T07:32:20.085297Z","iopub.execute_input":"2023-09-19T07:32:20.085666Z","iopub.status.idle":"2023-09-19T07:32:37.667517Z","shell.execute_reply.started":"2023-09-19T07:32:20.085636Z","shell.execute_reply":"2023-09-19T07:32:37.666300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fn = '/kaggle/input/open-problems-single-cell-perturbations/de_train.parquet'\ndf_de_train = pd.read_parquet(fn)","metadata":{"execution":{"iopub.status.busy":"2023-09-19T07:32:37.670027Z","iopub.execute_input":"2023-09-19T07:32:37.670450Z","iopub.status.idle":"2023-09-19T07:32:40.638965Z","shell.execute_reply.started":"2023-09-19T07:32:37.670412Z","shell.execute_reply":"2023-09-19T07:32:40.637722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feat_train = df_de_train[['cell_type', 'sm_name', 'SMILES', ]]\n\nsm_name2smiles = {\n        name: smiles \n        for name, smiles \n        in feat_train.drop_duplicates(subset='sm_name').iloc[::,1:].values\n    }","metadata":{"execution":{"iopub.status.busy":"2023-09-19T07:32:40.640264Z","iopub.execute_input":"2023-09-19T07:32:40.640609Z","iopub.status.idle":"2023-09-19T07:32:40.658309Z","shell.execute_reply.started":"2023-09-19T07:32:40.640583Z","shell.execute_reply":"2023-09-19T07:32:40.657138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fn = '/kaggle/input/open-problems-single-cell-perturbations/id_map.csv'\ndf_id_map = pd.read_csv(fn)","metadata":{"execution":{"iopub.status.busy":"2023-09-19T07:32:40.661251Z","iopub.execute_input":"2023-09-19T07:32:40.661580Z","iopub.status.idle":"2023-09-19T07:32:40.673026Z","shell.execute_reply.started":"2023-09-19T07:32:40.661553Z","shell.execute_reply":"2023-09-19T07:32:40.671716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_id_map['SMILES'] = [sm_name2smiles[name] for name in df_id_map.sm_name.values]","metadata":{"execution":{"iopub.status.busy":"2023-09-19T07:32:40.675185Z","iopub.execute_input":"2023-09-19T07:32:40.675794Z","iopub.status.idle":"2023-09-19T07:32:40.683899Z","shell.execute_reply.started":"2023-09-19T07:32:40.675747Z","shell.execute_reply":"2023-09-19T07:32:40.682537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_cls_pad_true, train_mean_pad_true = featurize_ChemBERTa(df_de_train.SMILES)\ntest_cls_pad_true, test_mean_pad_true = featurize_ChemBERTa(df_id_map.SMILES)\n\ntrain_cls_pad_false, train_mean_pad_false = featurize_ChemBERTa(df_de_train.SMILES, padding=False)\ntest_cls_pad_false, test_mean_pad_false = featurize_ChemBERTa(df_id_map.SMILES, padding=False)","metadata":{"execution":{"iopub.status.busy":"2023-09-19T07:32:40.685595Z","iopub.execute_input":"2023-09-19T07:32:40.686034Z","iopub.status.idle":"2023-09-19T07:32:56.594513Z","shell.execute_reply.started":"2023-09-19T07:32:40.685996Z","shell.execute_reply":"2023-09-19T07:32:56.593449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.save('train_sm_name.npy', df_de_train.sm_name.values)\nnp.save('train_ChemBERTa_v2_77MTR_cls_pad_True.npy', train_cls_pad_true)\nnp.save('train_ChemBERTa_v2_77MTR_mean_pad_True.npy', train_mean_pad_true)\nnp.save('train_ChemBERTa_v2_77MTR_cls_pad_False.npy', train_cls_pad_false)\nnp.save('train_ChemBERTa_v2_77MTR_mean_pad_False.npy', train_mean_pad_false)\n\nnp.save('test_sm_name.npy', df_id_map.sm_name.values)\nnp.save('test_ChemBERTa_v2_77MTR_cls_pad_True.npy', test_cls_pad_true)\nnp.save('test_ChemBERTa_v2_77MTR_mean_pad_True.npy', test_mean_pad_true)\nnp.save('test_ChemBERTa_v2_77MTR_cls_pad_False.npy', test_cls_pad_false)\nnp.save('test_ChemBERTa_v2_77MTR_mean_pad_False.npy', test_mean_pad_false)","metadata":{"execution":{"iopub.status.busy":"2023-09-19T07:32:56.596141Z","iopub.execute_input":"2023-09-19T07:32:56.596563Z","iopub.status.idle":"2023-09-19T07:32:56.614379Z","shell.execute_reply.started":"2023-09-19T07:32:56.596532Z","shell.execute_reply":"2023-09-19T07:32:56.613185Z"},"trusted":true},"execution_count":null,"outputs":[]}]}