{"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 os\nimport gc\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\n\nimport numpy as np\nimport pandas as pd\n\nfrom tqdm import tqdm\ntqdm.pandas()\n\n# annoy for approximate nearest neighbors\nfrom annoy import AnnoyIndex\n\nimport gc","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2023-04-28T16:03:45.837787Z","iopub.status.busy":"2023-04-28T16:03:45.837001Z","iopub.status.idle":"2023-04-28T16:03:47.300914Z","shell.execute_reply":"2023-04-28T16:03:47.298764Z","shell.execute_reply.started":"2023-04-28T16:03:45.83773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Assigning labels","metadata":{}},{"cell_type":"code","source":"DATA_DIR = '/kaggle/input/cafa-5-protein-function-prediction'\nMAX_LABELS = 1500","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:04:43.579185Z","iopub.status.busy":"2023-04-28T16:04:43.578707Z","iopub.status.idle":"2023-04-28T16:04:43.584613Z","shell.execute_reply":"2023-04-28T16:04:43.583109Z","shell.execute_reply.started":"2023-04-28T16:04:43.579147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_terms = pd.read_csv(os.path.join(DATA_DIR, 'Train', 'train_terms.tsv'), sep='\\t')\n\nterms = train_terms.groupby(['aspect', 'term'])['term'].count().reset_index(name='frequency')\nprint(terms.groupby('aspect')['term'].nunique())","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:04:43.874787Z","iopub.status.busy":"2023-04-28T16:04:43.873464Z","iopub.status.idle":"2023-04-28T16:04:49.401292Z","shell.execute_reply":"2023-04-28T16:04:49.39981Z","shell.execute_reply.started":"2023-04-28T16:04:43.874735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fractions = (terms.groupby('aspect')['term'].nunique() / terms['term'].nunique() * MAX_LABELS).apply(round)\nprint(fractions)\n\nselected_terms = set()\nfor aspect, number in fractions.items():\n    selection = terms.loc[(terms.aspect == aspect)]\n    selection = selection.nlargest(number, columns='frequency', keep='first')\n    selected_terms.update(selection.term.to_list())","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:04:57.935016Z","iopub.status.busy":"2023-04-28T16:04:57.93457Z","iopub.status.idle":"2023-04-28T16:04:57.993991Z","shell.execute_reply":"2023-04-28T16:04:57.992379Z","shell.execute_reply.started":"2023-04-28T16:04:57.934967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(selected_terms)","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:05:20.943919Z","iopub.status.busy":"2023-04-28T16:05:20.943437Z","iopub.status.idle":"2023-04-28T16:05:20.950926Z","shell.execute_reply":"2023-04-28T16:05:20.949432Z","shell.execute_reply.started":"2023-04-28T16:05:20.943874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def assign_labels(annotations, selected_terms=selected_terms):\n    \n    intersection = selected_terms.intersection(annotations)\n    labels = np.isin(np.array(list(selected_terms)), np.array(list(intersection)))\n    \n    return list(labels.astype('int'))\n\nannotations = train_terms.groupby('EntryID')['term'].apply(set)\nlabels = annotations.progress_apply(assign_labels)\n\nlabels.head()","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:05:32.700048Z","iopub.status.busy":"2023-04-28T16:05:32.699546Z","iopub.status.idle":"2023-04-28T16:05:53.676203Z","shell.execute_reply":"2023-04-28T16:05:53.674855Z","shell.execute_reply.started":"2023-04-28T16:05:32.700007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading train embeddings","metadata":{}},{"cell_type":"code","source":"labels.head()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ids = np.load('/kaggle/input/t5embeds/train_ids.npy')\n\nx_train = np.load('/kaggle/input/t5embeds/train_embeds.npy')\ny_train = np.array(labels[train_ids].to_list())","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:05:57.935332Z","iopub.status.busy":"2023-04-28T16:05:57.934869Z","iopub.status.idle":"2023-04-28T16:06:05.530673Z","shell.execute_reply":"2023-04-28T16:06:05.529264Z","shell.execute_reply.started":"2023-04-28T16:05:57.935291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"x_train, x_valid, y_train, y_valid = train_test_split(x_train, y_train, shuffle=True, random_state=42)","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:06:21.503028Z","iopub.status.busy":"2023-04-28T16:06:21.502546Z","iopub.status.idle":"2023-04-28T16:06:21.973194Z","shell.execute_reply":"2023-04-28T16:06:21.971979Z","shell.execute_reply.started":"2023-04-28T16:06:21.502985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# build the Annoy index for approximate nearest neighbors\nindex = AnnoyIndex(x_train.shape[1], 'angular')\nfor i, vector in enumerate(x_train):\n    index.add_item(i, vector)\nindex.build(100)","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:06:52.901265Z","iopub.status.busy":"2023-04-28T16:06:52.900838Z","iopub.status.idle":"2023-04-28T16:06:55.658736Z","shell.execute_reply":"2023-04-28T16:06:55.656918Z","shell.execute_reply.started":"2023-04-28T16:06:52.901229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"K = 8\n# find the K nearest neighbors for each vector in the validation set\n# return the indices and distances of the neighbors\nidxs = []\ndists = []\nfor vector in tqdm(x_valid):\n    idx, dist = index.get_nns_by_vector(vector, K, include_distances=True)\n    idxs.append(idx)\n    dists.append(dist)\n# convert the indices and distances to numpy arrays\nidxs = np.array(idxs)\ndists = np.array(dists)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Predict the probability of each label for each vector in the validation set\ny_hat = []\nfor i in tqdm(range(len(x_valid))):\n    y_hat_i = np.zeros(y_valid.shape[1])\n    for j in range(K):\n        y_hat_i += (1 - dists[i, j]) / K * y_train[idxs[i, j]]\n    y_hat.append(y_hat_i)\n# convert the predictions to a numpy array\ny_hat = np.array(y_hat)\n# clip the predictions to be between 0 and 1\ny_hat = np.clip(y_hat, 0, 1)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scores = pd.DataFrame(columns=list(selected_terms), index=['roc_auc'])\n\nfor i, term in enumerate(selected_terms):\n    score = roc_auc_score(y_valid[:, i], y_hat[:, i])\n    scores[term] = score\n\nscores.mean(axis=1)","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:07:04.817735Z","iopub.status.busy":"2023-04-28T16:07:04.817323Z","iopub.status.idle":"2023-04-28T16:07:06.388253Z","shell.execute_reply":"2023-04-28T16:07:06.387071Z","shell.execute_reply.started":"2023-04-28T16:07:04.8177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"test_ids = np.load('/kaggle/input/t5embeds/test_ids.npy')\nx_test = np.load('/kaggle/input/t5embeds/test_embeds.npy')","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:07:31.612771Z","iopub.status.busy":"2023-04-28T16:07:31.612338Z","iopub.status.idle":"2023-04-28T16:07:37.874243Z","shell.execute_reply":"2023-04-28T16:07:37.87302Z","shell.execute_reply.started":"2023-04-28T16:07:31.612734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# calculate the nearest neighbors for each vector in the test set\nidxs = []\ndists = []\nfor vector in tqdm(x_test):\n    idx, dist = index.get_nns_by_vector(vector, K, include_distances=True)\n    idxs.append(idx)\n    dists.append(dist)\n# convert the indices and distances to numpy arrays\nidxs = np.array(idxs)\ndists = np.array(dists)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Predict the probability of each label for each vector in the test set\npredictions = []\nfor i in tqdm(range(len(x_test))):\n    y_hat_i = np.zeros(y_train.shape[1])\n    for j in range(K):\n        y_hat_i += (1 - dists[i, j]) / K * y_train[idxs[i, j]]\n    predictions.append(y_hat_i)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# free up memory\ndel x_train, x_valid, y_train, y_valid, x_test, index, idxs, dists, y_hat\ngc.collect()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# convert the predictions to a numpy array\npredictions = np.array(predictions)\n# clip the predictions to be between 0 and 1\npredictions = np.clip(predictions, 0, 1)\n\nchunk_size = 10_000\nchunks = [range(i, min(i + chunk_size, len(predictions))) for i in range(0, len(predictions), chunk_size)]\n\nfinal_sub = pd.DataFrame()  # Create an empty DataFrame to hold the final result\n\nprint(f\"processing {len(chunks)} chunks of {chunk_size} predictions each\")\n\nfor chunk in chunks:\n    print(f\"processing chunk {chunk}\")\n    sub = pd.DataFrame(data=predictions[chunk], columns=list(selected_terms), index=test_ids[chunk])\n    sub = sub.T.unstack().reset_index(name='prediction')\n    sub = sub.loc[sub['prediction'] > 0]\n    final_sub = pd.concat([final_sub, sub])  # Concatenate current chunk DataFrame to the final DataFrame\n\nfinal_sub.head()","metadata":{"execution":{"iopub.execute_input":"2023-04-28T16:07:53.412245Z","iopub.status.busy":"2023-04-28T16:07:53.410645Z","iopub.status.idle":"2023-04-28T16:07:57.487826Z","shell.execute_reply":"2023-04-28T16:07:57.486338Z","shell.execute_reply.started":"2023-04-28T16:07:53.412189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_sub.to_csv('submission.tsv', sep='\\t', index=False, header=False)","metadata":{},"execution_count":null,"outputs":[]}]}