{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-20T18:55:29.81856Z","iopub.execute_input":"2023-04-20T18:55:29.818981Z","iopub.status.idle":"2023-04-20T18:55:29.838855Z","shell.execute_reply.started":"2023-04-20T18:55:29.818937Z","shell.execute_reply":"2023-04-20T18:55:29.837777Z"},"trusted":true},"execution_count":2,"outputs":[{"name":"stdout","text":"/kaggle/input/cafa-5-protein-function-prediction/sample_submission.tsv\n/kaggle/input/cafa-5-protein-function-prediction/IA.txt\n/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta\n/kaggle/input/cafa-5-protein-function-prediction/Test (Targets)/testsuperset-taxon-list.tsv\n/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\n/kaggle/input/cafa-5-protein-function-prediction/Train/train_sequences.fasta\n/kaggle/input/cafa-5-protein-function-prediction/Train/train_taxonomy.tsv\n/kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo\n","output_type":"stream"}]},{"cell_type":"code","source":"!pip install tensorflow_recommenders -q -U","metadata":{"execution":{"iopub.status.busy":"2023-04-20T18:55:29.840382Z","iopub.execute_input":"2023-04-20T18:55:29.840985Z","iopub.status.idle":"2023-04-20T18:55:41.729793Z","shell.execute_reply.started":"2023-04-20T18:55:29.840947Z","shell.execute_reply":"2023-04-20T18:55:41.72821Z"},"trusted":true},"execution_count":3,"outputs":[{"name":"stdout","text":"\u001b[33mWARNING: Running pip as the 'root' user can result in broken permissions and conflicting behaviour with the system package manager. It is recommended to use a virtual environment instead: https://pip.pypa.io/warnings/venv\u001b[0m\u001b[33m\n\u001b[0m","output_type":"stream"}]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\n    \nfrom sklearn.feature_extraction.text import CountVectorizer\n# import scipy.sparse as sp\nfrom numpy import log2\nfrom tensorflow.keras.utils import Sequence\nfrom scipy.spatial.distance import cosine\n\nimport tensorflow as tf\nfrom keras.layers import Input, Embedding, Dot, Flatten, Concatenate, Dense\nfrom keras.models import Model\nfrom keras.optimizers import Adam\nfrom sklearn.model_selection import train_test_split\n# from sentence_transformers import SentenceTransformer\nimport os\nimport pprint\nimport tempfile\nfrom typing import Dict, Text\nimport tensorflow_recommenders as tfrs\ntf.get_logger().setLevel('FATAL') #ERROR\n\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras import mixed_precision\n# mixed_precision.set_global_policy('mixed_float16')\n\n# %matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:20:44.691149Z","iopub.execute_input":"2023-04-20T19:20:44.691537Z","iopub.status.idle":"2023-04-20T19:20:53.641418Z","shell.execute_reply.started":"2023-04-20T19:20:44.691495Z","shell.execute_reply":"2023-04-20T19:20:53.640221Z"},"trusted":true},"execution_count":4,"outputs":[]},{"cell_type":"code","source":"def get_indexed_annot_preds(prot_index_num):\n    _, titles = index(tf.constant([str(prot_index_num)]))\n    # print(f\"Recommendations for {prot_index_num}, {user_ids[prot_index_num]},: {titles[0, :10]}\")\n    print(f\"Recommendations for {prot_index_num}, {user_ids[prot_index_num]}\")\n    titles = list(titles[0, :].numpy())\n    titles = [i.decode('utf8') for i in titles]\n    # titles = list(titles.numpy().astype(str)) # doesn't map\n    # annotations that protein already ahs: \n    existing_annots = data.loc[data[\"user_idx\"]==prot_index_num][\"item_id\"].unique()\n\n    # print(\"existing_annots\",existing_annots)\n    print(\"Novel predictions (not) existing annots):\",[i for i in titles if i not in existing_annots])\n    print(\"predictions present in existing annots:\",[i for i in titles if i in existing_annots])","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:31:18.652496Z","iopub.execute_input":"2023-04-20T19:31:18.653322Z","iopub.status.idle":"2023-04-20T19:31:18.662437Z","shell.execute_reply.started":"2023-04-20T19:31:18.653279Z","shell.execute_reply":"2023-04-20T19:31:18.661091Z"},"trusted":true},"execution_count":5,"outputs":[]},{"cell_type":"code","source":"print(\"Num GPUs Available: \", len(tf.config.list_physical_devices(\"GPU\")))","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:31:46.887131Z","iopub.execute_input":"2023-04-20T19:31:46.887531Z","iopub.status.idle":"2023-04-20T19:31:46.902002Z","shell.execute_reply.started":"2023-04-20T19:31:46.887494Z","shell.execute_reply":"2023-04-20T19:31:46.900479Z"},"trusted":true},"execution_count":6,"outputs":[{"name":"stdout","text":"Num GPUs Available:  0\n","output_type":"stream"}]},{"cell_type":"code","source":"# USE_PRETRAINED_DL = False\nFAST_RUN = False#True\n\nGET_k_TOP_PREDS = 300\n\nif FAST_RUN:\n    GET_k_TOP_PREDS = 50","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:32:26.54731Z","iopub.execute_input":"2023-04-20T19:32:26.547778Z","iopub.status.idle":"2023-04-20T19:32:26.553981Z","shell.execute_reply.started":"2023-04-20T19:32:26.547737Z","shell.execute_reply":"2023-04-20T19:32:26.552339Z"},"trusted":true},"execution_count":8,"outputs":[]},{"cell_type":"code","source":"data_file =\"/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\"","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:32:44.241653Z","iopub.execute_input":"2023-04-20T19:32:44.242044Z","iopub.status.idle":"2023-04-20T19:32:44.247069Z","shell.execute_reply.started":"2023-04-20T19:32:44.242009Z","shell.execute_reply":"2023-04-20T19:32:44.24575Z"},"trusted":true},"execution_count":9,"outputs":[]},{"cell_type":"code","source":"data = pd.read_csv(data_file,sep=\"\\t\",usecols=[\"EntryID\",\"term\"]).drop_duplicates()\nprint(data.shape[0])\ndata.columns = ['user_id','item_id'] \nif FAST_RUN:\n    data = data.sample(360_000)\n    it_cnts = data['item_id'].value_counts()\n    data = data[data['item_id'].map(it_cnts) > 100]\n    data = data[data['user_id'].map(data['user_id'].value_counts()) >= 2]\n    print(data.shape[0],\"# rows after fast run filter\")\n\n    ## keep sample of mosre popular annotations for speedup\nit_cnts = data['item_id'].value_counts()\ndata = data[data['item_id'].map(it_cnts) > 20]\nprint(data.nunique())\nprint(data[\"user_id\"].value_counts().describe().round(1))\nprint(data[\"item_id\"].value_counts())\ndisplay(data)","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:36:08.150047Z","iopub.execute_input":"2023-04-20T19:36:08.150532Z","iopub.status.idle":"2023-04-20T19:36:16.431098Z","shell.execute_reply.started":"2023-04-20T19:36:08.15048Z","shell.execute_reply":"2023-04-20T19:36:16.429933Z"},"trusted":true},"execution_count":10,"outputs":[{"name":"stdout","text":"5363863\nuser_id    142246\nitem_id     10399\ndtype: int64\ncount    142246.0\nmean         36.9\nstd          41.1\nmin           1.0\n25%          10.0\n50%          23.0\n75%          49.0\nmax         746.0\nName: user_id, dtype: float64\nGO:0005575    92912\nGO:0008150    92210\nGO:0110165    91286\nGO:0003674    78637\nGO:0005622    70785\n              ...  \nGO:0061909       21\nGO:0045196       21\nGO:0045611       21\nGO:1902224       21\nGO:0071111       21\nName: item_id, Length: 10399, dtype: int64\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"            user_id     item_id\n0        A0A009IHW8  GO:0008152\n1        A0A009IHW8  GO:0034655\n2        A0A009IHW8  GO:0072523\n3        A0A009IHW8  GO:0044270\n4        A0A009IHW8  GO:0006753\n...             ...         ...\n5363858      X5L565  GO:0050649\n5363859      X5L565  GO:0016491\n5363860      X5M5N0  GO:0005515\n5363861      X5M5N0  GO:0005488\n5363862      X5M5N0  GO:0003674\n\n[5243101 rows x 2 columns]","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>user_id</th>\n      <th>item_id</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>A0A009IHW8</td>\n      <td>GO:0008152</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>A0A009IHW8</td>\n      <td>GO:0034655</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>A0A009IHW8</td>\n      <td>GO:0072523</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>A0A009IHW8</td>\n      <td>GO:0044270</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>A0A009IHW8</td>\n      <td>GO:0006753</td>\n    </tr>\n    <tr>\n      <th>...</th>\n      <td>...</td>\n      <td>...</td>\n    </tr>\n    <tr>\n      <th>5363858</th>\n      <td>X5L565</td>\n      <td>GO:0050649</td>\n    </tr>\n    <tr>\n      <th>5363859</th>\n      <td>X5L565</td>\n      <td>GO:0016491</td>\n    </tr>\n    <tr>\n      <th>5363860</th>\n      <td>X5M5N0</td>\n      <td>GO:0005515</td>\n    </tr>\n    <tr>\n      <th>5363861</th>\n      <td>X5M5N0</td>\n      <td>GO:0005488</td>\n    </tr>\n    <tr>\n      <th>5363862</th>\n      <td>X5M5N0</td>\n      <td>GO:0003674</td>\n    </tr>\n  </tbody>\n</table>\n<p>5243101 rows × 2 columns</p>\n</div>"},"metadata":{}}]},{"cell_type":"code","source":"user_ids = data['user_id'].unique()\nitem_ids = data['item_id'].unique()\n\nuser_idx_map = {user_id: idx for idx, user_id in enumerate(user_ids)}\nitem_idx_map = {item_id: idx for idx, item_id in enumerate(item_ids)}","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:37:11.113689Z","iopub.execute_input":"2023-04-20T19:37:11.114189Z","iopub.status.idle":"2023-04-20T19:37:12.009985Z","shell.execute_reply.started":"2023-04-20T19:37:11.114136Z","shell.execute_reply":"2023-04-20T19:37:12.008669Z"},"trusted":true},"execution_count":11,"outputs":[]},{"cell_type":"code","source":"data[\"item_id\"] = data[\"item_id\"].astype(str)\ndata[\"user_id\"] = data[\"user_id\"].astype(str)","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:37:32.603819Z","iopub.execute_input":"2023-04-20T19:37:32.604254Z","iopub.status.idle":"2023-04-20T19:37:33.11765Z","shell.execute_reply.started":"2023-04-20T19:37:32.604212Z","shell.execute_reply":"2023-04-20T19:37:33.115889Z"},"trusted":true},"execution_count":12,"outputs":[]},{"cell_type":"code","source":"data['user_idx'] = data['user_id'].map(user_idx_map)\ndata['item_idx'] = data['item_id'].map(item_idx_map)\n\n## add:\nuser_idxs = data['user_idx'].unique()\nitem_idxs = data['item_idx'].unique()\n\nunique_user_ids = data['user_id'].unique()\nunique_movie_titles = data['item_id'].unique()","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:39:05.261756Z","iopub.execute_input":"2023-04-20T19:39:05.263205Z","iopub.status.idle":"2023-04-20T19:39:07.158389Z","shell.execute_reply.started":"2023-04-20T19:39:05.263134Z","shell.execute_reply":"2023-04-20T19:39:07.157184Z"},"trusted":true},"execution_count":13,"outputs":[]},{"cell_type":"markdown","source":"## CF Model\n#### Adapting TFRS","metadata":{}},{"cell_type":"code","source":"movies_df = data[[\"item_id\",\"item_idx\"]].drop_duplicates().reset_index(drop=True)\n\nmovies_df","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:40:55.574272Z","iopub.execute_input":"2023-04-20T19:40:55.57469Z","iopub.status.idle":"2023-04-20T19:40:56.494488Z","shell.execute_reply.started":"2023-04-20T19:40:55.574654Z","shell.execute_reply":"2023-04-20T19:40:56.493266Z"},"trusted":true},"execution_count":14,"outputs":[{"execution_count":14,"output_type":"execute_result","data":{"text/plain":"          item_id  item_idx\n0      GO:0008152         0\n1      GO:0034655         1\n2      GO:0072523         2\n3      GO:0044270         3\n4      GO:0006753         4\n...           ...       ...\n10394  GO:0052621     10394\n10395  GO:0001016     10395\n10396  GO:0015193     10396\n10397  GO:0004703     10397\n10398  GO:0071111     10398\n\n[10399 rows x 2 columns]","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>item_id</th>\n      <th>item_idx</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>GO:0008152</td>\n      <td>0</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>GO:0034655</td>\n      <td>1</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>GO:0072523</td>\n      <td>2</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>GO:0044270</td>\n      <td>3</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>GO:0006753</td>\n      <td>4</td>\n    </tr>\n    <tr>\n      <th>...</th>\n      <td>...</td>\n      <td>...</td>\n    </tr>\n    <tr>\n      <th>10394</th>\n      <td>GO:0052621</td>\n      <td>10394</td>\n    </tr>\n    <tr>\n      <th>10395</th>\n      <td>GO:0001016</td>\n      <td>10395</td>\n    </tr>\n    <tr>\n      <th>10396</th>\n      <td>GO:0015193</td>\n      <td>10396</td>\n    </tr>\n    <tr>\n      <th>10397</th>\n      <td>GO:0004703</td>\n      <td>10397</td>\n    </tr>\n    <tr>\n      <th>10398</th>\n      <td>GO:0071111</td>\n      <td>10398</td>\n    </tr>\n  </tbody>\n</table>\n<p>10399 rows × 2 columns</p>\n</div>"},"metadata":{}}]},{"cell_type":"markdown","source":"## Basic Implicit Recommendation Retrieval Model\n\n## Dot Product. In Batch Negatives.","metadata":{}},{"cell_type":"code","source":"# Variables\nseed = 42\ntest_percentage = 15\ntrain_percentage = 100-test_percentage\nembedding_dimension =  128#32 \nmetrics_batchsize = 256#16\ntrain_batchsize = 256#1024#128\ntest_batchsize = 256#1024#64\nlearning_rate = 0.1#0.004 # 0.5 ?\nepochs = 10#3\nindex_batchsize = 200\n\n\nviews_df =  data[[\"user_id\",\"item_id\"]]\nusers_df = data[\"user_id\"].unique()\ncontents_df = views_df['item_id'].unique()\n\nviews_ds = tf.data.Dataset.from_tensor_slices(dict(views_df))\ncontents_ds = tf.data.Dataset.from_tensor_slices(contents_df)","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:43:35.9672Z","iopub.execute_input":"2023-04-20T19:43:35.967903Z","iopub.status.idle":"2023-04-20T19:43:37.7038Z","shell.execute_reply.started":"2023-04-20T19:43:35.967858Z","shell.execute_reply":"2023-04-20T19:43:37.702438Z"},"trusted":true},"execution_count":15,"outputs":[]},{"cell_type":"code","source":"view_size = len(views_df)\ntrain_size = round(view_size/100*train_percentage)\ntest_size = view_size-train_size\n\ntf.random.set_seed(seed)\n\nviews_ds_shuffled = views_ds.shuffle(len(views_df), seed=seed, reshuffle_each_iteration=False)\ntrain = views_ds_shuffled.take(train_size)\ntest = views_ds_shuffled.skip(train_size).take(test_size)","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:47:15.041878Z","iopub.execute_input":"2023-04-20T19:47:15.042326Z","iopub.status.idle":"2023-04-20T19:47:15.105126Z","shell.execute_reply.started":"2023-04-20T19:47:15.042286Z","shell.execute_reply":"2023-04-20T19:47:15.103655Z"},"trusted":true},"execution_count":16,"outputs":[]},{"cell_type":"code","source":"user_model = tf.keras.Sequential([\ntf.keras.layers.StringLookup(vocabulary=users_df, mask_token=None),\ntf.keras.layers.Embedding(input_dim=len(users_df) + 1, output_dim=embedding_dimension)\n])\n\ncontent_model = tf.keras.Sequential([\ntf.keras.layers.StringLookup(\n    vocabulary=unique_movie_titles, mask_token=None),\ntf.keras.layers.Embedding(input_dim=len(contents_df) + 1, output_dim=embedding_dimension)])\n    \ncandidates=contents_ds.batch(metrics_batchsize).map(content_model)\n\nmetrics = tfrs.metrics.FactorizedTopK(\n  candidates=candidates,\n    ks = (1, 10, 100 # warning - slow\n          )\n)\n\ntask = tfrs.tasks.Retrieval(\n  metrics=metrics\n#   ,remove_accidental_hits = True # ValueError: When accidental hit removal is enabled, candidate ids must be supplied.\n)","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:52:03.254508Z","iopub.execute_input":"2023-04-20T19:52:03.254965Z","iopub.status.idle":"2023-04-20T19:52:03.491174Z","shell.execute_reply.started":"2023-04-20T19:52:03.254924Z","shell.execute_reply":"2023-04-20T19:52:03.489762Z"},"trusted":true},"execution_count":17,"outputs":[]},{"cell_type":"code","source":"class ContentModel(tfrs.Model):\n\n  def __init__(self, user_model, content_model):\n    super().__init__()\n    self.content_model: tf.keras.Model = content_model\n    self.user_model: tf.keras.Model = user_model\n    self.task: tf.keras.layers.Layer = task\n\n  def compute_loss(self, features: Dict[Text, tf.Tensor], training=False) -> tf.Tensor:\n    content_embeddings = self.content_model(features[\"item_id\"])\n    user_embeddings = self.user_model(features[\"user_id\"])\n    return self.task(query_embeddings=user_embeddings,\n                        candidate_embeddings=content_embeddings,\n                    # candidate_ids=contents_df\n                     ## speed up train:\n                     compute_metrics=not training, # speed up training , lose outputs\n                    )","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:56:10.838002Z","iopub.execute_input":"2023-04-20T19:56:10.838452Z","iopub.status.idle":"2023-04-20T19:56:10.848788Z","shell.execute_reply.started":"2023-04-20T19:56:10.838413Z","shell.execute_reply":"2023-04-20T19:56:10.847213Z"},"trusted":true},"execution_count":18,"outputs":[]},{"cell_type":"code","source":"%%time\nmodel = ContentModel(user_model, content_model)\n\nmodel.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=learning_rate))\n# model.compile(optimizer=tf.keras.optimizers.Adagrad(learning_rate=0.1))\n\ncached_train = train.shuffle(view_size).batch(train_batchsize).cache()\ncached_test = test.batch(8192).cache()","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:57:07.564676Z","iopub.execute_input":"2023-04-20T19:57:07.565141Z","iopub.status.idle":"2023-04-20T19:57:07.602953Z","shell.execute_reply.started":"2023-04-20T19:57:07.565082Z","shell.execute_reply":"2023-04-20T19:57:07.601459Z"},"trusted":true},"execution_count":19,"outputs":[{"name":"stdout","text":"CPU times: user 28.5 ms, sys: 706 µs, total: 29.2 ms\nWall time: 30.2 ms\n","output_type":"stream"}]},{"cell_type":"code","source":"%%time\ncallback = tf.keras.callbacks.EarlyStopping(patience=2)\n\nmodel.fit(cached_train, epochs=epochs, \n          validation_data = cached_test\n          , callbacks=[callback],\n          )\nprint(\"test\")\nmodel.evaluate(cached_test, return_dict=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-20T19:57:15.977201Z","iopub.execute_input":"2023-04-20T19:57:15.97763Z"},"trusted":true},"execution_count":null,"outputs":[{"name":"stdout","text":"Epoch 1/10\n 4594/17409 [======>.......................] - ETA: 1:08:46 - factorized_top_k/top_1_categorical_accuracy: 0.0000e+00 - factorized_top_k/top_10_categorical_accuracy: 0.0000e+00 - factorized_top_k/top_100_categorical_accuracy: 0.0000e+00 - loss: 27610.4670 - regularization_loss: 0.0000e+00 - total_loss: 27610.4670","output_type":"stream"}]},{"cell_type":"code","source":"%%time\n# Create a model that takes in raw query features, and\nindex = tfrs.layers.factorized_top_k.BruteForce(query_model=model.user_model,k=GET_k_TOP_PREDS)\n# recommends movies out of the entire (TRAIN)dataset.\nindex.index_from_dataset(\ntf.data.Dataset.zip((contents_ds.batch(GET_k_TOP_PREDS), contents_ds.batch(GET_k_TOP_PREDS).map(model.content_model)))\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_indexed_annot_preds(1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## View Some Predictions:","metadata":{}},{"cell_type":"code","source":"get_indexed_annot_preds(1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##### To get submissions - use index over test set. After adding data there and handling the cold start :)\n* I leave that as an exercise to the reader , or the upcoming paper I scratched this code out from!","metadata":{}},{"cell_type":"code","source":"## This should actually be on test set seqs(submission seems to include too many sequences)\ndata_sub = pd.read_csv(\"/kaggle/input/cafa-5-protein-function-prediction/sample_submission.tsv\",sep=\"\\t\",header=None).iloc[:,0:1].drop_duplicates()\nsub_ds = tf.data.Dataset.from_tensor_slices(data_sub)\ndisplay(data_sub)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### BruteForce Predictions Retrieval:","metadata":{}},{"cell_type":"code","source":"%%time\nlist_id = []\npred_labels = []\npred_score = []\n\nscores, title = index.predict([data_sub.values])\n\ntitle = title.flatten()\ntitle = [i.decode(\"utf8\") for i in title]\nlist_id = []\nfor row in data_sub.itertuples(index = False):\n    list_id.extend(list(row) * GET_k_TOP_PREDS)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submit = pd.DataFrame({\"EntryID\": list_id, \"term\": title, \"Score\": scores.flatten()}).round(3)\ndisplay(df_submit)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submit.to_csv(\"submission.tsv\", header = False, index = False, sep = \"\\t\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}