{"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":"markdown","source":"This notebook demonstrates how to use [ESM-2](https://github.com/facebookresearch/esm) transformer protein language model to extract embeddings from provided protein sequences. The following code calculates the embeddings for each protein sequence in the training and test FASTA files and saves them as individual `.pt` files. \n\nThe resulting embedding can be loaded using the following code:\n```\nimport torch\n\nembedding = torch.load('[EntryID].pt')\nembedding = embedding['mean_representations'][33].numpy()\n```\n\nComputing the embeddings and subsequently reading in the `.pt` files can take a while. The resulting numpy arrays can be found [here](https://www.kaggle.com/datasets/viktorfairuschin/cafa-5-ems-2-embeddings-numpy).\n\n**Note** that the test FASTA file contains duplicate entries. For this reason, this notebook uses cleaned FASTA files, which can be found [here](https://www.kaggle.com/datasets/viktorfairuschin/cafa-5-fasta-files).","metadata":{}},{"cell_type":"code","source":"!pip install -q fair-esm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-27T09:13:24.411907Z","iopub.execute_input":"2023-04-27T09:13:24.412358Z","iopub.status.idle":"2023-04-27T09:13:35.568596Z","shell.execute_reply.started":"2023-04-27T09:13:24.412317Z","shell.execute_reply":"2023-04-27T09:13:35.567393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pathlib\nimport torch\n\nfrom esm import FastaBatchedDataset, pretrained","metadata":{"execution":{"iopub.status.busy":"2023-04-27T09:13:35.571633Z","iopub.execute_input":"2023-04-27T09:13:35.57313Z","iopub.status.idle":"2023-04-27T09:13:37.671738Z","shell.execute_reply.started":"2023-04-27T09:13:35.573081Z","shell.execute_reply":"2023-04-27T09:13:37.67074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def extract_embeddings(model_name, fasta_file, output_dir, tokens_per_batch=4096, seq_length=1022,repr_layers=[33]):\n    \n    model, alphabet = pretrained.load_model_and_alphabet(model_name)\n    model.eval()\n\n    if torch.cuda.is_available():\n        model = model.cuda()\n        \n    dataset = FastaBatchedDataset.from_file(fasta_file)\n    batches = dataset.get_batch_indices(tokens_per_batch, extra_toks_per_seq=1)\n\n    data_loader = torch.utils.data.DataLoader(\n        dataset, \n        collate_fn=alphabet.get_batch_converter(seq_length), \n        batch_sampler=batches\n    )\n\n    output_dir.mkdir(parents=True, exist_ok=True)\n    \n    with torch.no_grad():\n        for batch_idx, (labels, strs, toks) in enumerate(data_loader):\n\n            print(f'Processing batch {batch_idx + 1} of {len(batches)}')\n\n            if torch.cuda.is_available():\n                toks = toks.to(device=\"cuda\", non_blocking=True)\n\n            out = model(toks, repr_layers=repr_layers, return_contacts=False)\n\n            logits = out[\"logits\"].to(device=\"cpu\")\n            representations = {layer: t.to(device=\"cpu\") for layer, t in out[\"representations\"].items()}\n            \n            for i, label in enumerate(labels):\n                entry_id = label.split()[0]\n                \n                filename = output_dir / f\"{entry_id}.pt\"\n                truncate_len = min(seq_length, len(strs[i]))\n\n                result = {\"entry_id\": entry_id}\n                result[\"mean_representations\"] = {\n                        layer: t[i, 1 : truncate_len + 1].mean(0).clone()\n                        for layer, t in representations.items()\n                    }\n\n                torch.save(result, filename)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T09:17:36.365089Z","iopub.execute_input":"2023-04-27T09:17:36.365675Z","iopub.status.idle":"2023-04-27T09:17:36.377982Z","shell.execute_reply.started":"2023-04-27T09:17:36.365639Z","shell.execute_reply":"2023-04-27T09:17:36.376831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Process train file","metadata":{}},{"cell_type":"code","source":"model_name = 'esm2_t33_650M_UR50D'\nfasta_file = pathlib.Path('/kaggle/input/cafa-5-fasta-files/train_sequences.fasta')\noutput_dir = pathlib.Path('train_embeddings')\n\nextract_embeddings(model_name, fasta_file, output_dir)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T10:01:47.293773Z","iopub.execute_input":"2023-04-27T10:01:47.294539Z","iopub.status.idle":"2023-04-27T10:02:05.683327Z","shell.execute_reply.started":"2023-04-27T10:01:47.2945Z","shell.execute_reply":"2023-04-27T10:02:05.681804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Process test file","metadata":{}},{"cell_type":"code","source":"model_name = 'esm2_t33_650M_UR50D'\nfasta_file = pathlib.Path('/kaggle/input/cafa-5-fasta-files/test_sequences.fasta')\noutput_dir = pathlib.Path('test_embeddings')\n\nextract_embeddings(model_name, fasta_file, output_dir)","metadata":{"execution":{"iopub.status.busy":"2023-04-27T10:02:08.194801Z","iopub.execute_input":"2023-04-27T10:02:08.196037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}