{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":41875,"databundleVersionId":5521661,"sourceType":"competition"},{"sourceId":5499219,"sourceType":"datasetVersion","datasetId":3167603},{"sourceId":7338488,"sourceType":"datasetVersion","datasetId":4260533}],"dockerImageVersionId":30627,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### ESM2","metadata":{}},{"cell_type":"code","source":"from Bio import SeqIO\nimport numpy as np\nimport torch\nimport pandas as pd\n\n\ndef read_fasta(file_path):\n    sequences = []\n    with open(file_path, 'r') as fasta_file:\n        for record in SeqIO.parse(fasta_file, 'fasta'):\n            sequences.append({\n                'id': record.id,\n                'description': record.description,\n                'sequence': str(record.seq)\n            })\n    return sequences\n\n# Replace 'your_sequence.fasta' with the actual file path of your FASTA file\nfile_path = '/kaggle/input/scop-sequneces/scop_fa_represeq_lib_latest.fa.txt'\nfamily_sequences = read_fasta(file_path)\nfile_path = '/kaggle/input/scop-sequneces/scop_sf_represeq_lib_latest.fa.txt'\nsuperfamily_sequences = read_fasta(file_path)\nfamily_df = pd.DataFrame(family_sequences)\ndf_fa = family_df['description'].str.extract(r'FA=(\\d+) FA-PDBID=([A-Z0-9_]+) FA-UNIID=([A-Za-z0-9_]+)')\ndf_fa.columns = ['FA', 'FA-PDBID', 'FA-UNIID']\ndf_fa = pd.concat([df_fa, family_df], axis=1)\nsuperfamily_df = pd.DataFrame(superfamily_sequences)\ndf_sf = superfamily_df['description'].str.extract(r'SF=(\\d+) SF-PDBID=([A-Z0-9_]+) SF-UNIID=([A-Za-z0-9_]+)')\ndf_sf.columns = ['SF', 'SF_PDBID', 'SF_UNIID']\ndf_sf = pd.concat([df_sf, superfamily_df], axis=1)\ndf_sf['sequence']\nsentences_sf = df_sf['sequence'].to_list()\nsentences_f = df_fa['sequence'].to_list()\n","metadata":{"execution":{"iopub.status.busy":"2024-01-04T18:47:38.862285Z","iopub.execute_input":"2024-01-04T18:47:38.862608Z","iopub.status.idle":"2024-01-04T18:47:43.995557Z","shell.execute_reply.started":"2024-01-04T18:47:38.862579Z","shell.execute_reply":"2024-01-04T18:47:43.994734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-01-04T18:59:39.707521Z","iopub.execute_input":"2024-01-04T18:59:39.708537Z","iopub.status.idle":"2024-01-04T18:59:39.739142Z","shell.execute_reply.started":"2024-01-04T18:59:39.708497Z","shell.execute_reply":"2024-01-04T18:59:39.738242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model, alphabet = esm.pretrained.esm2_t36_3B_UR50D()\n# batch_converter = alphabet.get_batch_converter()\n# model.eval()  # disables dropout for deterministic results\n\n# model = model.cuda() # uncomment to run on GPU","metadata":{"execution":{"iopub.status.busy":"2024-01-04T16:12:55.686722Z","iopub.execute_input":"2024-01-04T16:12:55.687362Z","iopub.status.idle":"2024-01-04T16:12:55.690694Z","shell.execute_reply.started":"2024-01-04T16:12:55.687334Z","shell.execute_reply":"2024-01-04T16:12:55.689852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoTokenizer, EsmForMaskedLM\nimport torch\n\ntokenizer = AutoTokenizer.from_pretrained(\"facebook/esm2_t6_8M_UR50D\")\nmodel = EsmForMaskedLM.from_pretrained(\"facebook/esm2_t6_8M_UR50D\", output_hidden_states=True)\n\npad_len = max(df_sf['sequence'].map(len))\n\ndef get_embeddings(sequence):\n    inputs = tokenizer(sequence, padding='max_length', max_length=pad_len, truncation=True, return_tensors=\"pt\")\n\n    with torch.no_grad():\n        logits = model(**inputs)\n\n#     # retrieve index of <mask>\n#     mask_token_index = (inputs.input_ids == tokenizer.mask_token_id)[0].nonzero(as_tuple=True)[0]\n\n#     predicted_token_id = logits[0, mask_token_index].argmax(axis=-1)\n\n#     labels = tokenizer(sequence, return_tensors=\"pt\")[\"input_ids\"]\n#     # mask labels of non-<mask> tokens\n#     outputs = model(**inputs, labels=labels)\n    \n    return logits.hidden_states[-1].mean(dim=1)[0]","metadata":{"execution":{"iopub.status.busy":"2024-01-04T19:18:40.581525Z","iopub.execute_input":"2024-01-04T19:18:40.581959Z","iopub.status.idle":"2024-01-04T19:18:40.966973Z","shell.execute_reply.started":"2024-01-04T19:18:40.581925Z","shell.execute_reply":"2024-01-04T19:18:40.965962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\nscop_sf_res = []\n\nfor i in tqdm(sentences_sf):\n    scop_sf_res.append(get_embeddings(i))\n\nres_sf = torch.cat(scop_sf_res).reshape([len(scop_sf_res), len(scop_sf_res[0])])\n    \ntorch.save(res_sf, f'esm2_scop_sf_{len(scop_sf_res)}.pt')","metadata":{"execution":{"iopub.status.busy":"2024-01-04T19:20:11.815837Z","iopub.execute_input":"2024-01-04T19:20:11.816499Z","iopub.status.idle":"2024-01-04T19:20:31.546823Z","shell.execute_reply.started":"2024-01-04T19:20:11.816464Z","shell.execute_reply":"2024-01-04T19:20:31.545856Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\n\nscop_f_res = []\n\nfor i in tqdm(sentences_f):\n    scop_f_res.append(get_embeddings(i))\n\nres_f = torch.cat(scop_f_res).reshape([len(scop_f_res), len(scop_f_res[0])])    \n\ntorch.save(res_f, f'esm2_scop_f_{len(scop_f_res)}.pt')","metadata":{"execution":{"iopub.status.busy":"2024-01-04T19:20:31.548641Z","iopub.execute_input":"2024-01-04T19:20:31.549027Z","iopub.status.idle":"2024-01-04T19:20:51.322303Z","shell.execute_reply.started":"2024-01-04T19:20:31.548989Z","shell.execute_reply":"2024-01-04T19:20:51.321462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\n\ndef get_plot(vectors, model_name=\"\", fixed=False):\n    # Set the center point (e.g., point 0)\n    center_point = np.zeros_like(vectors[0])\n\n    # Calculate distances from the center point for each vector in the vocabulary\n    distances = [np.linalg.norm(word_vector) for word_vector in vectors]\n    if fixed:\n        # Set x-axis limits\n        plt.xlim(0, 12)\n\n        # Set y-axis limits\n        plt.ylim(0, 0.6)\n\n    # Plot a histogram of distances to visualize the density\n    plt.hist(distances, bins=1000, density=True, alpha=0.75, color='b')\n    plt.title(f'Density of {model_name} Vectors from Center Point')\n    plt.xlabel('Distance from Center Point')\n    plt.ylabel('Density')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-01-04T19:29:14.88143Z","iopub.execute_input":"2024-01-04T19:29:14.882171Z","iopub.status.idle":"2024-01-04T19:29:14.889173Z","shell.execute_reply.started":"2024-01-04T19:29:14.882133Z","shell.execute_reply":"2024-01-04T19:29:14.888177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_plot(res_sf, \"ESM2 Embeddings on SF\")","metadata":{"execution":{"iopub.status.busy":"2024-01-04T19:29:19.499845Z","iopub.execute_input":"2024-01-04T19:29:19.500217Z","iopub.status.idle":"2024-01-04T19:29:21.132982Z","shell.execute_reply.started":"2024-01-04T19:29:19.500186Z","shell.execute_reply":"2024-01-04T19:29:21.131993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_plot(res_f, \"ESM2 Embeddings on F\")","metadata":{"execution":{"iopub.status.busy":"2024-01-04T19:29:25.613584Z","iopub.execute_input":"2024-01-04T19:29:25.614554Z","iopub.status.idle":"2024-01-04T19:29:27.322103Z","shell.execute_reply.started":"2024-01-04T19:29:25.614509Z","shell.execute_reply":"2024-01-04T19:29:27.321226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### bad code below","metadata":{}},{"cell_type":"code","source":"# def process_sequence(sequence, sequence_id, max_allowed_len, model, embed_layer, alphabet, batch_converter):\n#     \"\"\"\n#     Process a sequence and return token representations.\n#     \"\"\"\n#     residual_embeds = []\n\n#     while len(sequence) > 0:\n#         # Prepare data\n#         data = [(sequence_id, sequence[:min(len(sequence), max_allowed_len)])]\n#         _, _, batch_tokens = batch_converter(data)\n#         batch_tokens = batch_tokens.cuda()\n#         tokens_len = (batch_tokens != alphabet.padding_idx).sum(1)\n\n#         # Model inference\n#         with torch.no_grad():\n#             results = model(batch_tokens, repr_layers=[embed_layer], return_contacts=False)\n#         token_representations = results[\"representations\"][embed_layer]\n        \n#         # Adjust token representations and add to list\n#         offset = max_allowed_len // 2 if residual_embeds else 1\n#         token_representations = token_representations[:, offset:(tokens_len-1), :].squeeze(dim=0)\n#         residual_embeds.append(token_representations)\n\n#         # Update sequence\n#         sequence = sequence[max_allowed_len // 2:]\n\n#     return residual_embeds\n\n# # Usage example\n# MAX_ALLOWED_LEN = 1022  # Assuming this is a predefined constant\n# sequence_id = '1'  # Replace with actual sequence ID\n# sequence = sentences_sf[0]  # Replace with actual sequence\n\n# residual_embeddings = process_sequence(sequence, sequence_id, MAX_ALLOWED_LEN, model, 0, alphabet, batch_converter)\n","metadata":{"execution":{"iopub.status.busy":"2024-01-04T16:00:49.155053Z","iopub.execute_input":"2024-01-04T16:00:49.155904Z","iopub.status.idle":"2024-01-04T16:00:49.425266Z","shell.execute_reply.started":"2024-01-04T16:00:49.155865Z","shell.execute_reply":"2024-01-04T16:00:49.424066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# residual_embeddings[0][0]","metadata":{"execution":{"iopub.status.busy":"2024-01-04T16:02:08.696091Z","iopub.execute_input":"2024-01-04T16:02:08.696891Z","iopub.status.idle":"2024-01-04T16:02:08.70579Z","shell.execute_reply.started":"2024-01-04T16:02:08.696851Z","shell.execute_reply":"2024-01-04T16:02:08.704554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # embed_layer = 33\n# embed_layer = 36\n\n# max_allowed_len = 1022\n\n# # List for all protein embeds (mean of residual embeds)\n# protein_embeds = []\n# residual_embeds = []\n\n# id = list_ids[0]\n# # seq = list_seqs[0]\n# seq = df_seqs[df_seqs[\"id\"] == \"O48653\"][\"seq\"].values.tolist()\n# seq = seq[0]\n# len(seq)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # List for all residual embeds for current protein\n# residual_embeds = []\n    \n# # Embeds of first 1022 residuals\n# data = [(id, seq[0:min(len(seq), max_allowed_len)])]\n# batch_labels, batch_strs, batch_tokens = batch_converter(data)\n# batch_tokens = batch_tokens.cuda() \n# tokens_len = (batch_tokens != alphabet.padding_idx).sum(1)\n\n# with torch.no_grad():\n#     results = model(batch_tokens, repr_layers=[embed_layer], return_contacts=False)\n# token_representations = results[\"representations\"][embed_layer]\n# token_representations = token_representations[:, 1:(tokens_len-1), :]\n# token_representations = token_representations.squeeze(dim = 0)\n# residual_embeds.append(token_representations)\n    \n# # Remove first 511 residuals\n# seq = seq[max_allowed_len//2:]\n    \n# while len(seq) > 0:\n#     data = [(id, seq[0:min(len(seq), max_allowed_len)])]\n#     batch_labels, batch_strs, batch_tokens = batch_converter(data)\n#     batch_tokens = batch_tokens.cuda() \n#     tokens_len = (batch_tokens != alphabet.padding_idx).sum(1)\n#     with torch.no_grad():\n#         results = model(batch_tokens, repr_layers=[embed_layer], return_contacts=False)\n#     token_representations = results[\"representations\"][embed_layer]\n#     token_representations = token_representations[:, 1:(tokens_len-1), :]\n#     token_representations = token_representations[:, max_allowed_len//2:, :]\n#     token_representations = token_representations.squeeze(dim = 0)\n#     residual_embeds.append(token_representations)\n                 \n#     seq = seq[max_allowed_len//2:]\n","metadata":{},"execution_count":null,"outputs":[]}]}