{"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 code is mostly from [here](https://www.kaggle.com/code/cdeotte/rapids-svr-cv-0-450-lb-0-44x).\nFor now i do it with the small ESM model as the others take much more compute.","metadata":{}},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os, gc, re, warnings\nfrom Bio import SeqIO\nfrom tqdm import tqdm\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-22T14:48:00.750439Z","iopub.execute_input":"2023-06-22T14:48:00.750882Z","iopub.status.idle":"2023-06-22T14:48:00.901667Z","shell.execute_reply.started":"2023-06-22T14:48:00.750839Z","shell.execute_reply":"2023-06-22T14:48:00.900487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"GENERATE ID AND SEQ LIST FOR TRAIN.\")\ntrain_ids = []\ntrain_sequences = []\nfor record in tqdm(SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\", \"fasta\")):\n    train_ids.append(record.id)\n    train_sequences.append(str(record.seq))\n    \nprint(\"PUT TRAIN INFO IN A DATAFRAME.\")\n# put the info in a dataframe with columns id, sequence, label_1, label_2, ...\ntrain = pd.DataFrame({'id': train_ids, 'sequence': train_sequences})\ndef add_spaces(x):\n    return \" \".join(list(x))\ntrain.sequence = train.sequence.map(add_spaces)\nprint('Train has shape',train.shape)\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:48:00.943019Z","iopub.execute_input":"2023-06-22T14:48:00.943374Z","iopub.status.idle":"2023-06-22T14:48:07.386681Z","shell.execute_reply.started":"2023-06-22T14:48:00.943341Z","shell.execute_reply":"2023-06-22T14:48:07.385758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"GENERATE ID AND SEQ LIST FOR TEST.\")\ntest_ids = []\ntest_sequences = []\nfor record in tqdm(SeqIO.parse(\"/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta\", \"fasta\")):\n    test_ids.append(record.id)\n    test_sequences.append(str(record.seq))\n    \nprint(\"PUT TEST INFO IN A DATAFRAME.\")\n# put the info in a dataframe with columns id, sequence, label_1, label_2, ...\ntest = pd.DataFrame({'id': test_ids, 'sequence': test_sequences})\ndef add_spaces(x):\n    return \" \".join(list(x))\ntest.sequence = test.sequence.map(add_spaces)\nprint('Train has shape',test.shape)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:48:07.388621Z","iopub.execute_input":"2023-06-22T14:48:07.389462Z","iopub.status.idle":"2023-06-22T14:48:15.79188Z","shell.execute_reply.started":"2023-06-22T14:48:07.389436Z","shell.execute_reply":"2023-06-22T14:48:15.790947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Train shape:',train.shape,'Test shape:',test.shape,'Columns:',test.columns)","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:48:15.793437Z","iopub.execute_input":"2023-06-22T14:48:15.794141Z","iopub.status.idle":"2023-06-22T14:48:15.800445Z","shell.execute_reply.started":"2023-06-22T14:48:15.794107Z","shell.execute_reply":"2023-06-22T14:48:15.799511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import AutoModel,AutoTokenizer\nimport torch\nimport torch.nn.functional as F","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:48:15.80296Z","iopub.execute_input":"2023-06-22T14:48:15.803443Z","iopub.status.idle":"2023-06-22T14:48:20.709676Z","shell.execute_reply.started":"2023-06-22T14:48:15.803413Z","shell.execute_reply":"2023-06-22T14:48:20.708648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mean_pooling(model_output, attention_mask):\n    token_embeddings = model_output.last_hidden_state.detach().cpu()\n    input_mask_expanded = (\n        attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()\n    )\n    return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(\n        input_mask_expanded.sum(1), min=1e-9\n    )","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:48:20.711087Z","iopub.execute_input":"2023-06-22T14:48:20.712147Z","iopub.status.idle":"2023-06-22T14:48:20.719621Z","shell.execute_reply.started":"2023-06-22T14:48:20.712113Z","shell.execute_reply":"2023-06-22T14:48:20.718511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 32\n\nclass EmbedDataset(torch.utils.data.Dataset):\n    def __init__(self,df):\n        self.df = df.reset_index(drop=True)\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self,idx):\n        text = self.df.loc[idx,\"sequence\"]\n        tokens = tokenizer(\n                text,\n                None,\n                add_special_tokens=True,\n                padding='max_length',\n                truncation=True,\n                max_length=MAX_LEN,return_tensors=\"pt\")\n        tokens = {k:v.squeeze(0) for k,v in tokens.items()}\n        return tokens\n\nds_tr = EmbedDataset(train)\nembed_dataloader_tr = torch.utils.data.DataLoader(ds_tr,\\\n                        batch_size=BATCH_SIZE,\\\n                        shuffle=False)\nds_te = EmbedDataset(test)\nembed_dataloader_te = torch.utils.data.DataLoader(ds_te,\\\n                        batch_size=BATCH_SIZE,\\\n                        shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:49:02.242327Z","iopub.execute_input":"2023-06-22T14:49:02.242771Z","iopub.status.idle":"2023-06-22T14:49:02.291318Z","shell.execute_reply.started":"2023-06-22T14:49:02.242733Z","shell.execute_reply":"2023-06-22T14:49:02.29031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer = None\nMAX_LEN = 640\n\ndef get_embeddings(MODEL_NM='', MAX=640, BATCH_SIZE=BATCH_SIZE, verbose=True):\n    global tokenizer, MAX_LEN\n    DEVICE=\"cuda\"\n    model = AutoModel.from_pretrained( MODEL_NM )\n    tokenizer = AutoTokenizer.from_pretrained( MODEL_NM )\n    MAX_LEN = MAX\n    \n    model = model.to(DEVICE)\n    model.eval()\n    all_train_text_feats = []\n    for batch in tqdm(embed_dataloader_tr,total=len(embed_dataloader_tr)):\n        input_ids = batch[\"input_ids\"].to(DEVICE)\n        attention_mask = batch[\"attention_mask\"].to(DEVICE)\n        with torch.no_grad():\n            model_output = model(input_ids=input_ids,attention_mask=attention_mask)\n        sentence_embeddings = mean_pooling(model_output, attention_mask.detach().cpu())\n        # Normalize the embeddings\n        sentence_embeddings = F.normalize(sentence_embeddings, p=2, dim=1)\n        sentence_embeddings =  sentence_embeddings.squeeze(0).detach().cpu().numpy()\n        all_train_text_feats.extend(sentence_embeddings)\n    all_train_text_feats = np.array(all_train_text_feats)\n    if verbose:\n        print('Train embeddings shape',all_train_text_feats.shape)\n        \n    te_text_feats = []\n    for batch in tqdm(embed_dataloader_te,total=len(embed_dataloader_te)):\n        input_ids = batch[\"input_ids\"].to(DEVICE)\n        attention_mask = batch[\"attention_mask\"].to(DEVICE)\n        with torch.no_grad():\n            model_output = model(input_ids=input_ids,attention_mask=attention_mask)\n        sentence_embeddings = mean_pooling(model_output, attention_mask.detach().cpu())\n        # Normalize the embeddings\n        sentence_embeddings = F.normalize(sentence_embeddings, p=2, dim=1)\n        sentence_embeddings =  sentence_embeddings.squeeze(0).detach().cpu().numpy()\n        te_text_feats.extend(sentence_embeddings)\n    te_text_feats = np.array(te_text_feats)\n    if verbose:\n        print('Test embeddings shape',te_text_feats.shape)\n        \n    return all_train_text_feats, te_text_feats","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:49:06.138633Z","iopub.execute_input":"2023-06-22T14:49:06.139008Z","iopub.status.idle":"2023-06-22T14:49:06.151618Z","shell.execute_reply.started":"2023-06-22T14:49:06.138979Z","shell.execute_reply":"2023-06-22T14:49:06.150557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_NM = 'zjunlp/OntoProtein'\nLEN = 1028\nall_train_text_feats, te_text_feats = get_embeddings(MODEL_NM, MAX=LEN)\n# Save the numpy arrays in a file with model name and max len indicated in the filename\nnp.save(f'train_{MODEL_NM.split(\"/\")[1]}_{LEN}', all_train_text_feats)\nnp.save(f'test_{MODEL_NM.split(\"/\")[1]}_{LEN}', te_text_feats)","metadata":{"execution":{"iopub.status.busy":"2023-06-22T14:49:10.08386Z","iopub.execute_input":"2023-06-22T14:49:10.084584Z","iopub.status.idle":"2023-06-22T14:49:17.731536Z","shell.execute_reply.started":"2023-06-22T14:49:10.084542Z","shell.execute_reply":"2023-06-22T14:49:17.730106Z"},"trusted":true},"execution_count":null,"outputs":[]}]}