{"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":"import pandas as pd\nimport torch\nimport torch.nn as nn\nfrom tqdm import tqdm\nimport numpy as np\nimport gc","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:18.997884Z","iopub.execute_input":"2023-06-20T05:20:18.999122Z","iopub.status.idle":"2023-06-20T05:20:22.265742Z","shell.execute_reply.started":"2023-06-20T05:20:18.999079Z","shell.execute_reply":"2023-06-20T05:20:22.264787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nDEVICE","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:22.267575Z","iopub.execute_input":"2023-06-20T05:20:22.268204Z","iopub.status.idle":"2023-06-20T05:20:22.300622Z","shell.execute_reply.started":"2023-06-20T05:20:22.26817Z","shell.execute_reply":"2023-06-20T05:20:22.299632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/c5-dataset/train_data.csv')\n# discard rows with long sequences\ntrain_df['seq_len'] = train_df.seq.map(lambda s: len(s))\nMAX_SEQ_LEN = 1000\ntrain_df = train_df[train_df.seq_len <= MAX_SEQ_LEN].reset_index()\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:22.303027Z","iopub.execute_input":"2023-06-20T05:20:22.303813Z","iopub.status.idle":"2023-06-20T05:20:25.34511Z","shell.execute_reply.started":"2023-06-20T05:20:22.303781Z","shell.execute_reply":"2023-06-20T05:20:25.344117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx_to_amino = list(set(ch for s in train_df.seq for ch in s))\nPAD_IDX = len(idx_to_amino)\nidx_to_amino.append('<pad>')\namino_to_idx = {v: k for k,v in enumerate(idx_to_amino)}\nlen(idx_to_amino)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:25.349084Z","iopub.execute_input":"2023-06-20T05:20:25.349872Z","iopub.status.idle":"2023-06-20T05:20:28.46385Z","shell.execute_reply.started":"2023-06-20T05:20:25.349824Z","shell.execute_reply":"2023-06-20T05:20:28.462796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"idx_to_term = list(set(\n    term for ts in [\n        ts[2:-2].split(\"', '\") \n        for ts in train_df.term\n    ]\n    for term in ts)\n)\nlen(idx_to_term)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:28.465243Z","iopub.execute_input":"2023-06-20T05:20:28.465694Z","iopub.status.idle":"2023-06-20T05:20:29.818673Z","shell.execute_reply.started":"2023-06-20T05:20:28.46566Z","shell.execute_reply":"2023-06-20T05:20:29.817771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del train_df\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:29.820301Z","iopub.execute_input":"2023-06-20T05:20:29.820667Z","iopub.status.idle":"2023-06-20T05:20:29.982876Z","shell.execute_reply.started":"2023-06-20T05:20:29.820633Z","shell.execute_reply":"2023-06-20T05:20:29.981903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# amino embedding + LSTM + linear projection to terms\nclass C5Model(nn.Module):\n    def __init__(self, embedding_dim, hidden_dim, lstm_num_layers):\n        super().__init__()\n        self.embedding = nn.Embedding(num_embeddings=len(idx_to_amino), embedding_dim=embedding_dim)\n        self.lstm = nn.LSTM(input_size=embedding_dim, hidden_size=hidden_dim, num_layers=lstm_num_layers, batch_first=True)\n        self.linear_out = nn.Linear(in_features=hidden_dim, out_features=len(idx_to_term))\n        #self.register_buffer('onehot', torch.eye(MAX_SEQ_LEN).bool(), persistent=False)\n\n    def forward(self, x):\n        x, lens = x\n        x = self.embedding(x)\n        x = self.lstm(x)[0]\n        #x = x[self.onehot[lens - 1][:, :lens.max()]] # get the output at the end of each sequence (ignore padding)\n        x = x[:, -1, :]\n        return self.linear_out(torch.relu(x))\n\n\nmodel = C5Model(embedding_dim=64, hidden_dim=256, lstm_num_layers=2).to(DEVICE)\nmodel.load_state_dict(torch.load('/kaggle/input/c5-model/model_weights_epoch_19.pt', map_location=DEVICE))\nmodel.eval()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-20T05:20:29.984633Z","iopub.execute_input":"2023-06-20T05:20:29.985273Z","iopub.status.idle":"2023-06-20T05:20:35.690792Z","shell.execute_reply.started":"2023-06-20T05:20:29.98524Z","shell.execute_reply":"2023-06-20T05:20:35.689848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer","metadata":{}},{"cell_type":"code","source":"from Bio import SeqIO\n\nsequences = SeqIO.parse('/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta', \"fasta\")\n\nsequence_tensors = [\n    (sequence.id, torch.tensor([amino_to_idx[a] for a in sequence.seq], device=DEVICE))\n    for sequence in sequences\n]","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:20:35.69222Z","iopub.execute_input":"2023-06-20T05:20:35.692544Z","iopub.status.idle":"2023-06-20T05:21:09.772794Z","shell.execute_reply.started":"2023-06-20T05:20:35.692513Z","shell.execute_reply":"2023-06-20T05:21:09.771858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(sequence_tensors)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:21:09.774339Z","iopub.execute_input":"2023-06-20T05:21:09.774739Z","iopub.status.idle":"2023-06-20T05:21:09.782407Z","shell.execute_reply.started":"2023-06-20T05:21:09.774705Z","shell.execute_reply":"2023-06-20T05:21:09.781506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with torch.no_grad():\n    all_term_idxs = []\n    all_probs = []\n    seq_ids = []\n    for seq_id, s in tqdm(sequence_tensors):\n        pred = torch.round(torch.sigmoid(model((s[None, :], None))[0]) * 1000) / 1000\n        top = pred.topk(1500)\n        term_idxs = top.indices[top.values > 0.001]\n        probs = pred[term_idxs]\n        seq_ids.append(seq_id)\n        all_term_idxs.append(term_idxs)\n        all_probs.append(probs)","metadata":{"execution":{"iopub.status.busy":"2023-06-20T05:21:09.785455Z","iopub.execute_input":"2023-06-20T05:21:09.786415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\n\nwith open('submission.tsv', 'w') as tsvfile:\n    writer = csv.writer(tsvfile, delimiter='\\t')\n    for (seq_id, term_idxs, probs) in tqdm(zip(seq_ids, all_term_idxs, all_probs)):\n        writer.writerows([\n            [seq_id, idx_to_term[term_idx], '%.3f' % prob]\n            for term_idx, prob in zip(term_idxs.to('cpu').numpy(), probs.to('cpu').numpy())\n        ])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head -10 submission.tsv","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!tail -10 submission.tsv","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}