{"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":"# 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-08-21T11:29:24.96638Z","iopub.execute_input":"2023-08-21T11:29:24.966871Z","iopub.status.idle":"2023-08-21T11:29:24.992272Z","shell.execute_reply.started":"2023-08-21T11:29:24.966835Z","shell.execute_reply":"2023-08-21T11:29:24.990962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_labels = pd.concat([\n#     pd.read_csv('../../temporal/labels/prop_test_leak_no_dup.tsv', sep='\\t'),\n#     pd.read_csv('../../quickgo/prop_quickgo51.tsv', sep='\\t'),\n# ]).drop_duplicates().reset_index(drop=True)\n\nfn1 = '/kaggle/input/cafa5-go-annotations/go_labels/labels/prop_test_leak_no_dup.tsv'\nfn2 = '/kaggle/input/cafa5-go-annotations/prop_quickgo51.tsv'\n\ntest_labels = pd.concat([\n    pd.read_csv(fn1, sep='\\t'),\n    pd.read_csv(fn2, sep='\\t'),\n]).drop_duplicates().reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-21T10:30:21.226924Z","iopub.execute_input":"2023-08-21T10:30:21.227436Z","iopub.status.idle":"2023-08-21T10:30:21.328186Z","shell.execute_reply.started":"2023-08-21T10:30:21.227398Z","shell.execute_reply":"2023-08-21T10:30:21.326824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels","metadata":{"execution":{"iopub.status.busy":"2023-08-21T10:30:26.735876Z","iopub.execute_input":"2023-08-21T10:30:26.736296Z","iopub.status.idle":"2023-08-21T10:30:26.759209Z","shell.execute_reply.started":"2023-08-21T10:30:26.736262Z","shell.execute_reply":"2023-08-21T10:30:26.758027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels['EntryID'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-08-21T10:36:12.483948Z","iopub.execute_input":"2023-08-21T10:36:12.484664Z","iopub.status.idle":"2023-08-21T10:36:12.504958Z","shell.execute_reply.started":"2023-08-21T10:36:12.484618Z","shell.execute_reply":"2023-08-21T10:36:12.503316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head /kaggle/input/protein-subs/sub/submission.tsv","metadata":{"execution":{"iopub.status.busy":"2023-08-21T10:39:40.566263Z","iopub.execute_input":"2023-08-21T10:39:40.566686Z","iopub.status.idle":"2023-08-21T10:39:41.711591Z","shell.execute_reply.started":"2023-08-21T10:39:40.56665Z","shell.execute_reply":"2023-08-21T10:39:41.709891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nfn = '/kaggle/input/protein-subs/sub/submission.tsv'\nds = pd.read_csv(fn, sep = '\\t', header = None, index_col = None)\nds","metadata":{"execution":{"iopub.status.busy":"2023-08-21T10:40:06.109253Z","iopub.execute_input":"2023-08-21T10:40:06.109812Z","iopub.status.idle":"2023-08-21T10:40:40.869743Z","shell.execute_reply.started":"2023-08-21T10:40:06.109745Z","shell.execute_reply":"2023-08-21T10:40:40.868603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = ds[0].isin( set(test_labels['EntryID']) )\nds2 = ds[m].copy()\nds2","metadata":{"execution":{"iopub.status.busy":"2023-08-21T11:54:56.141983Z","iopub.execute_input":"2023-08-21T11:54:56.142455Z","iopub.status.idle":"2023-08-21T11:55:00.526733Z","shell.execute_reply.started":"2023-08-21T11:54:56.142416Z","shell.execute_reply":"2023-08-21T11:55:00.525302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds2.to_csv('ogryzok_predictions.tsv', sep = '\\t', header = None, index = None )","metadata":{"execution":{"iopub.status.busy":"2023-08-21T11:59:05.779461Z","iopub.execute_input":"2023-08-21T11:59:05.77992Z","iopub.status.idle":"2023-08-21T11:59:06.534706Z","shell.execute_reply.started":"2023-08-21T11:59:05.779884Z","shell.execute_reply":"2023-08-21T11:59:06.533399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head ogryzok_predictions.tsv","metadata":{"execution":{"iopub.status.busy":"2023-08-21T11:59:06.536278Z","iopub.execute_input":"2023-08-21T11:59:06.53667Z","iopub.status.idle":"2023-08-21T11:59:07.929524Z","shell.execute_reply.started":"2023-08-21T11:59:06.53662Z","shell.execute_reply":"2023-08-21T11:59:07.927697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Create a dictionary to assign each go term to the roots (CCO, MFO, BPO)\n\nimport re\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport joblib\nimport pickle\nfrom tqdm import tqdm\nfrom Bio import SeqIO\nimport gc\n\ndef extract_go_terms_and_branches(file_path):\n    with open(file_path, 'r') as file:\n        content = file.read()\n        # Match each stanza with [Term] in the OBO file\n        stanzas = re.findall(r'\\[Term\\][\\s\\S]*?(?=\\n\\[|$)', content)\n\n    go_terms_dict = {}\n    for stanza in stanzas:\n        # Extract the GO term ID\n        go_id = re.search(r'^id: (GO:\\d+)', stanza, re.MULTILINE)\n        if go_id:\n            go_id = go_id.group(1)\n\n        # Extract the namespace (branch)\n        namespace = re.search(r'^namespace: (\\w+)', stanza, re.MULTILINE)\n        if namespace:\n            namespace = namespace.group(1)\n\n        if go_id and namespace:\n            # Map the branch abbreviation to the corresponding BPO, CCO, or MFO\n            branch_abbr = {'biological_process': 'BPO', 'cellular_component': 'CCO', 'molecular_function': 'MFO'}\n            go_terms_dict[go_id] = branch_abbr[namespace]\n\n    return go_terms_dict\n\nfile_path = '/kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo'\ngo_terms_dict = extract_go_terms_and_branches(file_path)\n","metadata":{"execution":{"iopub.status.busy":"2023-08-21T11:34:21.518171Z","iopub.execute_input":"2023-08-21T11:34:21.518701Z","iopub.status.idle":"2023-08-21T11:34:26.056665Z","shell.execute_reply.started":"2023-08-21T11:34:21.51866Z","shell.execute_reply":"2023-08-21T11:34:26.054903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"str(go_terms_dict)[:100]","metadata":{"execution":{"iopub.status.busy":"2023-08-21T11:34:29.902525Z","iopub.execute_input":"2023-08-21T11:34:29.90311Z","iopub.status.idle":"2023-08-21T11:34:29.928056Z","shell.execute_reply.started":"2023-08-21T11:34:29.903059Z","shell.execute_reply":"2023-08-21T11:34:29.926021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define a class to manage predictions for proteins.\n# The class keeps track of the highest score for each GO (Gene Ontology) term prediction.\n# Note: This assumes scores are comparable, which might not be the case.\n# A ranking-based selection could be more suitable.\n# Each branch outputs a maximum of 35 predictions for each protein after sorting predictions from highest to lowest.\n# There is an option to add a bonus to the score if the term is predicted by multiple methods.\n\nclass ProteinPredictions:\n    # Initialize an empty dictionary to store the predictions\n    def __init__(self):\n        self.predictions = {}\n\n    # Add a prediction to the storage, with optional bonus\n    # Arguments:\n    #   - protein: Identifier for the protein\n    #   - go_term: GO term that is being predicted\n    #   - score: Confidence score of the prediction\n    #   - branch: Branch of the Gene Ontology (e.g., 'CCO', 'MFO', 'BPO')\n    #   - bonus: Optional bonus to be added to the score\n    def add_prediction(self, protein, go_term, score, branch, bonus=0, hist=0, inc=1):\n        # If the protein is not already in the storage, initialize its structure\n        if protein not in self.predictions:\n            self.predictions[protein] = {'CCO': {}, 'MFO': {}, 'BPO': {}}\n        \n        # Convert the score to a float for comparison and calculation\n        score = float(score)\n\n        # If this GO term has already been predicted for this protein and branch,\n        # add the bonus to the score. Keep the highest score.\n        if go_term in self.predictions[protein][branch]:\n            if self.predictions[protein][branch][go_term] < score:\n                self.predictions[protein][branch][go_term] = (score*inc+self.predictions[protein][branch][go_term]*hist)/(hist+inc)    + bonus\n            else:\n                self.predictions[protein][branch][go_term]  = (score*inc+self.predictions[protein][branch][go_term]*hist)/(hist+inc) + bonus\n        # If this GO term has not been predicted yet, store it with the score\n        else:\n            self.predictions[protein][branch][go_term] = score\n\n        # Ensure that the score does not exceed 1\n        if self.predictions[protein][branch][go_term] > 1:\n            self.predictions[protein][branch][go_term] = 1\n\n    # Export the stored predictions to a file\n    # Arguments:\n    #   - output_file: File name for the exported predictions\n    #   - top: Number of top predictions to export for each protein and branch\n    def get_predictions(self, output_file='submission.tsv', top=40):\n        # Open the output file\n        with open(output_file, 'w') as f:\n            # Iterate through each protein and its branches\n            for protein, branches in self.predictions.items():     \n                # For each branch, sort the GO terms by score in descending order and select the top ones\n                for branch, go_terms in branches.items():\n                    if branch =='CCO': \n                        top = 39\n                    if branch =='MFO': \n                        top = 39\n                    if branch =='BPO': \n                        top = 39     \n                    # Sort go_terms by score in descending order and take the top ones\n                    top_go_terms = sorted(go_terms.items(), key=lambda x: x[1], reverse=True)[:top]\n                    # Write each of the top predictions to the file\n                    for go_term, score in top_go_terms:\n                        f.write(f\"{protein}\\t{go_term}\\t{score:.3f}\\n\")\n","metadata":{"execution":{"iopub.status.busy":"2023-08-21T11:41:52.093819Z","iopub.execute_input":"2023-08-21T11:41:52.094324Z","iopub.status.idle":"2023-08-21T11:41:52.113604Z","shell.execute_reply.started":"2023-08-21T11:41:52.094284Z","shell.execute_reply":"2023-08-21T11:41:52.112553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"protein_predictions = ProteinPredictions()\n","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:02:16.164397Z","iopub.execute_input":"2023-08-21T12:02:16.164928Z","iopub.status.idle":"2023-08-21T12:02:16.170898Z","shell.execute_reply.started":"2023-08-21T12:02:16.164889Z","shell.execute_reply":"2023-08-21T12:02:16.169693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(set(test_labels['EntryID']) )","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:02:17.900569Z","iopub.execute_input":"2023-08-21T12:02:17.901069Z","iopub.status.idle":"2023-08-21T12:02:17.912593Z","shell.execute_reply.started":"2023-08-21T12:02:17.901033Z","shell.execute_reply":"2023-08-21T12:02:17.910983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_selected_proteins = list( set(test_labels['EntryID']) )","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:02:18.75234Z","iopub.execute_input":"2023-08-21T12:02:18.752755Z","iopub.status.idle":"2023-08-21T12:02:18.761168Z","shell.execute_reply.started":"2023-08-21T12:02:18.752722Z","shell.execute_reply":"2023-08-21T12:02:18.759659Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nfor l in tqdm(open('/kaggle/input/quick-go-2022-03-02/quickgo.tsv')):\n    item_list = l.split('\\t')\n    temp_id = item_list[1]\n    go=item_list[2].strip()\n    score = float(1)\n    if temp_id not in list_selected_proteins: continue \n    if go in go_terms_dict:\n        root = go_terms_dict[go]\n        #branch = item_list[3].strip()\n        protein_predictions.add_prediction(temp_id, go, score, root, 0, 0,1)","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:02:19.256938Z","iopub.execute_input":"2023-08-21T12:02:19.257442Z","iopub.status.idle":"2023-08-21T12:03:48.538058Z","shell.execute_reply.started":"2023-08-21T12:02:19.257405Z","shell.execute_reply":"2023-08-21T12:03:48.536521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fn = 'ogryzok_predictions.tsv'\nfor l in tqdm(open( fn )):\n    item_list = l.split('\\t')\n    temp_id = item_list[0]\n    go=item_list[1]\n    score = float(item_list[2].strip())\n    if go in go_terms_dict:\n        root = go_terms_dict[go]\n        #branch = item_list[3].strip()\n        protein_predictions.add_prediction(temp_id, go, score, root, 0, 1, 1)","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:04:20.28274Z","iopub.execute_input":"2023-08-21T12:04:20.283288Z","iopub.status.idle":"2023-08-21T12:04:21.04846Z","shell.execute_reply.started":"2023-08-21T12:04:20.28325Z","shell.execute_reply":"2023-08-21T12:04:21.046967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"protein_predictions.get_predictions()","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:04:48.497186Z","iopub.execute_input":"2023-08-21T12:04:48.497588Z","iopub.status.idle":"2023-08-21T12:04:48.620722Z","shell.execute_reply.started":"2023-08-21T12:04:48.497559Z","shell.execute_reply":"2023-08-21T12:04:48.619534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ndf_pred = pd.read_csv('submission.tsv', index_col = None, header = None, sep = '\\t')\ndf_pred","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:06:01.678897Z","iopub.execute_input":"2023-08-21T12:06:01.679424Z","iopub.status.idle":"2023-08-21T12:06:01.7376Z","shell.execute_reply.started":"2023-08-21T12:06:01.679389Z","shell.execute_reply":"2023-08-21T12:06:01.736097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CAFA5 metric","metadata":{}},{"cell_type":"code","source":"%%time\n\n# Evaluation for CAFA \n# https://github.com/BioComputingUP/CAFA-evaluator\n\nflag_correct_metric_computation_bug_found_by_Anton = True\n# https://www.kaggle.com/competitions/cafa-5-protein-function-prediction/discussion/420241 - Anton Vakhrushev - correcting error in the initial code of the metric computation - please upvote \n\n\nimport numpy as np\nimport pandas as pd\nimport multiprocessing as mp\nimport copy\nimport logging\n\nclass Graph:\n    \"\"\"\n    Ontology class. One ontology == one namespace\n    DAG is the adjacence matrix (sparse) which represent a Directed Acyclic Graph where\n    DAG(i,j) == 1 means that the go term i is_a (or is part_of) j\n    Parents that are in a different namespace are discarded\n    \"\"\"\n    def __init__(self, namespace, terms_dict, ia_dict=None, orphans=False):\n        \"\"\"\n        terms_dict = {term: {name: , namespace: , def: , alt_id: , rel:}}\n        \"\"\"\n        self.namespace = namespace\n        self.dag = []  # [[], ...] terms (rows, axis 0) x parents (columns, axis 1)\n        self.terms_dict = {}  # {term: {index: , name: , namespace: , def: }  used to assign term indexes in the gt\n        self.terms_list = []  # [{id: term, name:, namespace: , def:, adg: [], children: []}, ...]\n        self.idxs = None  # Number of terms\n        self.order = None\n        self.toi = None\n        self.ia = None\n\n        rel_list = []\n        for self.idxs, (term_id, term) in enumerate(terms_dict.items()):\n            rel_list.extend([[term_id, rel, term['namespace']] for rel in term['rel']])\n            self.terms_list.append({'id': term_id, 'name': term['name'], 'namespace': namespace, 'def': term['def'],\n                                 'adj': [], 'children': []})\n            self.terms_dict[term_id] = {'index': self.idxs, 'name': term['name'], 'namespace': namespace, 'def': term['def']}\n            for a_id in term['alt_id']:\n                self.terms_dict[a_id] = copy.copy(self.terms_dict[term_id])\n        self.idxs += 1\n\n        self.dag = np.zeros((self.idxs, self.idxs), dtype='bool')\n\n        # id1 term (row, axis 0), id2 parent (column, axis 1)\n        for id1, id2, ns in rel_list:\n            if self.terms_dict.get(id2):\n                i = self.terms_dict[id1]['index']\n                j = self.terms_dict[id2]['index']\n                self.dag[i, j] = 1\n                self.terms_list[i]['adj'].append(j)\n                self.terms_list[j]['children'].append(i)\n                logging.debug(\"i,j {},{} {},{}\".format(i, j, id1, id2))\n            else:\n                logging.debug('Skipping branch to external namespace: {}'.format(id2))\n        logging.debug(\"dag {}\".format(self.dag))\n        # Topological sorting\n        self.top_sort()\n        logging.debug(\"order sorted {}\".format(self.order))\n\n        if orphans:\n            self.toi = np.arange(self.dag.shape[0])  # All terms, also those without parents\n        else:\n            self.toi = np.nonzero(self.dag.sum(axis=1) > 0)[0]  # Only terms with parents\n        logging.debug(\"toi {}\".format(self.toi))\n\n        if ia_dict is not None:\n            self.set_ia(ia_dict)\n\n        return\n\n    def top_sort(self):\n        \"\"\"\n        Takes a sparse matrix representing a DAG and returns an array with nodes indexes in topological order\n        https://en.wikipedia.org/wiki/Topological_sorting\n        \"\"\"\n        indexes = []\n        visited = 0\n        (rows, cols) = self.dag.shape\n\n        # create a vector containing the in-degree of each node\n        in_degree = self.dag.sum(axis=0)\n        # logging.debug(\"degree {}\".format(in_degree))\n\n        # find the nodes with in-degree 0 (leaves) and add them to the queue\n        queue = np.nonzero(in_degree == 0)[0].tolist()\n        # logging.debug(\"queue {}\".format(queue))\n\n        # for each element of the queue increment visits, add them to the list of ordered nodes\n        # and decrease the in-degree of the neighbor nodes\n        # and add them to the queue if they reach in-degree == 0\n        while queue:\n            visited += 1\n            idx = queue.pop(0)\n            indexes.append(idx)\n            in_degree[idx] -= 1\n            l = self.terms_list[idx]['adj']\n            if len(l) > 0:\n                for j in l:\n                    in_degree[j] -= 1\n                    if in_degree[j] == 0:\n                        queue.append(j)\n\n        # if visited is equal to the number of nodes in the graph then the sorting is complete\n        # otherwise the graph can't be sorted with topological order\n        if visited == rows:\n            self.order = indexes\n        else:\n            raise Exception(\"The sparse matrix doesn't represent an acyclic graph\")\n\n    def set_ia(self, ia_dict):\n        self.ia = np.zeros(self.idxs, dtype='float')\n        for term_id in self.terms_dict:\n            if ia_dict.get(term_id):\n                self.ia[self.terms_dict[term_id]['index']] = ia_dict.get(term_id)\n            else:\n                logging.debug('Missing IA for term: {}'.format(term_id))\n        # Convert inf to zero\n        np.nan_to_num(self.ia, copy=False, nan=0, posinf=0, neginf=0)\n        self.toi = np.nonzero(self.ia > 0)[0]\n\n\nclass Prediction:\n    \"\"\"\n    The score matrix contains the scores given by the predictor for every node of the ontology\n    \"\"\"\n    def __init__(self, ids, matrix, idx, namespace=None):\n        self.ids = ids\n        self.matrix = matrix  # scores\n        self.next_idx = idx\n        # self.n_pred_seq = idx + 1\n        self.namespace = namespace\n\n    def __str__(self):\n        return \"\\n\".join([\"{}\\t{}\\t{}\".format(index, self.matrix[index], self.namespace) for index, _id in enumerate(self.ids)])\n\n\nclass GroundTruth:\n    def __init__(self, ids, matrix, namespace=None):\n        self.ids = ids\n        self.matrix = matrix\n        self.namespace = namespace\n\n\ndef propagate(matrix, ont, order, mode='max'):\n    \"\"\"\n    Update inplace the score matrix (proteins x terms) up to the root taking the max between children and parents\n    \"\"\"\n    if matrix.shape[0] == 0:\n        raise Exception(\"Empty matrix\")\n\n    deepest = np.where(np.sum(matrix[:, order], axis=0) > 0)[0][0]\n    if deepest.size == 0:\n        raise Exception(\"The matrix is empty\")\n\n    # Remove leaves\n    order_ = np.delete(order, [range(0, deepest)])\n\n    for i in order_:\n        # Get direct children\n        children = np.where(ont.dag[:, i] != 0)[0]\n        if children.size > 0:\n            cols = np.concatenate((children, [i]))\n            if mode == 'max':\n                matrix[:, i] = matrix[:, cols].max(axis=1)\n            elif mode == 'fill':\n                rows = np.where(matrix[:, i] == 0)[0]\n                if rows.size:\n                    idx = np.ix_(rows, cols)\n                    if flag_correct_metric_computation_bug_found_by_Anton:\n                        # Corrected way (see https://www.kaggle.com/competitions/cafa-5-protein-function-prediction/discussion/420241 )\n                        matrix[rows, i] = matrix[idx].max(axis=1) #  matrix[idx].max(axis=1)[0] # Correction: https://www.kaggle.com/competitions/cafa-5-protein-function-prediction/discussion/420241\n                    else:\n                        # Old way - not corrected\n                        matrix[rows, i] = matrix[idx].max(axis=1)[0] # Correction: https://www.kaggle.com/competitions/cafa-5-protein-function-prediction/discussion/420241\n                        \n    return\n\n\ndef obo_parser(obo_file, valid_rel=(\"is_a\", \"part_of\")):\n    \"\"\"\n    Parse a OBO file and returns a list of ontologies, one for each namespace.\n    Obsolete terms are excluded as well as external namespaces.\n    \"\"\"\n    term_dict = {}\n    term_id = None\n    namespace = None\n    name = None\n    term_def = None\n    alt_id = []\n    rel = []\n    obsolete = True\n    with open(obo_file) as f:\n        for line in f:\n            line = line.strip().split(\": \")\n            if line and len(line) > 1:\n                k = line[0]\n                v = \": \".join(line[1:])\n                if k == \"id\":\n                    # Populate the dictionary with the previous entry\n                    if term_id is not None and obsolete is False and namespace is not None:\n                        term_dict.setdefault(namespace, {})[term_id] = {'name': name,\n                                                                       'namespace': namespace,\n                                                                       'def': term_def,\n                                                                       'alt_id': alt_id,\n                                                                       'rel': rel}\n                    # Assign current term ID\n                    term_id = v\n\n                    # Reset optional fields\n                    alt_id = []\n                    rel = []\n                    obsolete = False\n                    namespace = None\n\n                elif k == \"alt_id\":\n                    alt_id.append(v)\n                elif k == \"name\":\n                    name = v\n                elif k == \"namespace\" and v != 'external':\n                    namespace = v\n                elif k == \"def\":\n                    term_def = v\n                elif k == 'is_obsolete':\n                    obsolete = True\n                elif k == \"is_a\" and k in valid_rel:\n                    s = v.split('!')[0].strip()\n                    rel.append(s)\n                elif k == \"relationship\" and v.startswith(\"part_of\") and \"part_of\" in valid_rel:\n                    s = v.split()[1].strip()\n                    rel.append(s)\n\n        # Last record\n        if obsolete is False and namespace is not None:\n            term_dict.setdefault(namespace, {})[term_id] = {'name': name,\n                                                          'namespace': namespace,\n                                                          'def': term_def,\n                                                          'alt_id': alt_id,\n                                                          'rel': rel}\n    return term_dict\n\n\ndef gt_parser(gt_file, ontologies):\n    \"\"\"\n    Parse ground truth file. Discard terms not included in the ontology.\n    \"\"\"\n    gt_dict = {}\n    with open(gt_file) as f:\n        for line in f:\n            line = line.strip().split()\n            if line:\n                p_id, term_id = line[:2]\n                for ont in ontologies:\n                    if term_id in ont.terms_dict:\n                        gt_dict.setdefault(ont.namespace, {}).setdefault(p_id, []).append(term_id)\n                        break\n\n    gts = {}\n    for ont in ontologies:\n        if gt_dict.get(ont.namespace):\n            matrix = np.zeros((len(gt_dict[ont.namespace]), ont.idxs), dtype='bool')\n            ids = {}\n            for i, p_id in enumerate(gt_dict[ont.namespace]):\n                ids[p_id] = i\n                for term_id in gt_dict[ont.namespace][p_id]:\n                    matrix[i, ont.terms_dict[term_id]['index']] = 1\n            propagate(matrix, ont, ont.order, mode='max')\n            gts[ont.namespace] = GroundTruth(ids, matrix, ont.namespace)\n\n    return gts\n\n\ndef pred_parser(f, ontologies, gts, prop_mode, max_terms=None):\n    \"\"\"\n    Parse a prediction file and returns a list of prediction objects, one for each namespace.\n    If a predicted is predicted multiple times for the same target, it stores the max.\n    This is the slow step if the input file is huge, ca. 1 minute for 5GB input on SSD disk.\n    \"\"\"\n    ids = {}\n    matrix = {}\n    ns_dict = {}  # {namespace: term}\n    onts = {ont.namespace: ont for ont in ontologies}\n    for ns in gts:\n        matrix[ns] = np.zeros(gts[ns].matrix.shape, dtype='float')\n        ids[ns] = {}\n        for term in onts[ns].terms_dict:\n            ns_dict[term] = ns\n\n    for line in f:\n        p_id, term_id, prob = line\n        ns = ns_dict.get(term_id)\n        if ns in gts and p_id in gts[ns].ids:\n            i = gts[ns].ids[p_id]\n            if max_terms is None or np.count_nonzero(matrix[ns][i]) <= max_terms:\n                j = onts[ns].terms_dict.get(term_id)['index']\n                ids[ns][p_id] = i\n                matrix[ns][i, j] = max(matrix[ns][i, j], float(prob))\n\n    predictions = []\n    for ns in ids:\n        if ids[ns]:\n            propagate(matrix[ns], onts[ns], onts[ns].order, mode=prop_mode)\n            predictions.append(Prediction(ids[ns], matrix[ns], len(ids[ns]), ns))\n\n    if not predictions:\n        raise Exception(\"Empty prediction, check format\")\n\n    return predictions\n\n\ndef ia_parser(file):\n    ia_dict = {}\n    with open(file) as f:\n        for line in f:\n            if line:\n                term, ia = line.strip().split()\n                ia_dict[term] = float(ia)\n    return ia_dict\n\n# Computes the root terms in the dag\ndef get_roots_idx(dag):\n    return np.where(dag.sum(axis=1) == 0)[0]\n\n\n# Computes the leaf terms in the dag\ndef get_leafs_idx(dag):\n    return np.where(dag.sum(axis=0) == 0)[0]\n\n\n# Return a mask for all the predictions (matrix) >= tau\ndef solidify_prediction(pred, tau):\n    return pred >= tau\n\n\n# computes the f metric for each precision and recall in the input arrays\ndef compute_f(pr, rc):\n    n = 2 * pr * rc\n    d = pr + rc\n    return np.divide(n, d, out=np.zeros_like(n, dtype=float), where=d != 0)\n\n\ndef compute_s(ru, mi):\n    return np.sqrt(ru**2 + mi**2)\n    # return np.where(np.isnan(ru), mi, np.sqrt(ru + np.nan_to_num(mi)))\n\nimport time\nfrom scipy.sparse import csr_matrix\n\ndef compute_metrics_(tau_arr, g, pred, toi, n_gt, wn_gt=None, ic_arr=None):\n\n    verbose = 0;         \n\n    if verbose >= 10:\n        t0 = time.time()\n    \n    metrics = np.zeros((len(tau_arr), 7), dtype='float')  # cov, pr, rc, wpr, wrc, ru, mi\n\n    if verbose >= 10:\n        print('type(toi), toi', type(toi), toi )\n    tmp = pred.matrix[:, toi]\n    if verbose >= 10:\n        print('type(tmp), tmp.shape', type(tmp), tmp.shape )\n    p_s = csr_matrix(tmp )\n    ic_arr_toi = ic_arr[toi]\n    if verbose >= 10:\n        print('type(ic_arr_toi), ic_arr_toi.shape', type(ic_arr_toi), ic_arr_toi.shape )\n\n    \n    g_s = csr_matrix( g )\n    \n    if verbose >= 10:\n        print( 'csr_matrix done %.1f'%(time.time( ) - t0 ), 'p_s.shape, g_s.shape:', p_s.shape, g_s.shape )    \n\n\n    \n    for i, tau in enumerate(tau_arr):\n\n#         p = solidify_prediction(pred_matrix_toi, tau)\n\n#         # number of proteins with at least one term predicted with score >= tau\n#         metrics[i, 0] = (p.sum(axis=1) > 0).sum()\n\n#         # Terms subsets\n#         intersection = np.logical_and(p, g)  # TP                                # SLOW PART !!!\n\n#         # Subsets size\n#         n_pred = p.sum(axis=1)\n#         n_intersection = intersection.sum(axis=1)\n\n#         # Precision, recall\n#         metrics[i, 1] = np.divide(n_intersection, n_pred, out=np.zeros_like(n_intersection, dtype='float'),\n#                                   where=n_pred > 0).sum()\n#         metrics[i, 2] = np.divide(n_intersection, n_gt, out=np.zeros_like(n_gt, dtype='float'), where=n_gt > 0).sum()\n\n\n        if verbose >= 100:\n            t0 = time.time()\n            print()\n            print(i, tau, 'Start %.1f'%(time.time( ) - t0 ) )\n        p = p_s > tau # solidify_prediction(p, tau)\n        if verbose >= 100:\n            print(i, tau, 'solidify done %.1f'%(time.time( ) - t0 ), 'p.shape:', p.shape,  )\n#         p_s = csr_matrix(p)\n#         print(i, tau, 'csr_matrix done %.1f'%(time.time( ) - t0 ), 'p.shape:', p.shape,  )\n        \n#         print(p.shape)\n        # number of proteins with at least one term predicted with score >= tau\n        metrics[i, 0] = (p.sum(axis=1) > 0).sum()\n\n        # Terms subsets\n#         intersection = np.logical_and(p, g)  # TP\n        intersection = p.multiply( g_s)  # TP\n\n\n        if ic_arr is not None:\n            \n            # Weighted precision, recall\n            # wn_pred = (p * ic_arr_toi).sum(axis=1) # \n#             wn_pred = np.dot(p , ic_arr_toi)\n            wn_pred = p.dot( ic_arr_toi)\n            # wn_intersection = (intersection *ic_arr_toi).sum(axis=1)\n#             wn_intersection = np.dot( intersection , ic_arr_toi )\n            wn_intersection =  intersection.dot( ic_arr_toi )\n            \n            if verbose >= 100:\n                print(i, tau, 'After w_pred wn_intersection  %.1f'%(time.time( ) - t0 ) )\n            \n            metrics[i, 3] = np.divide(wn_intersection, wn_pred, out=np.zeros( wn_intersection.shape, dtype='float'),\n                                      where=wn_pred > 0).sum()\n            metrics[i, 4] = np.divide(wn_intersection, wn_gt, out=np.zeros(wn_intersection.shape, dtype='float'),\n                                      where=n_gt > 0).sum()\n            if verbose >= 100:\n                print(i, tau, 'After metrics 3,4   %.1f'%(time.time( ) - t0 ) )\n\n#             # Terms subsets\n#             remaining = np.logical_and(np.logical_not(p), g)  # FN --> not predicted but in the ground truth\n#             mis = np.logical_and(p, np.logical_not(g))  # FP --> predicted but not in the ground truth\n\n#             print(i, tau, 'After remining and miss  %.1f'%(time.time( ) - t0 ) )\n            \n#             # Misinformation, remaining uncertainty\n#             metrics[i, 5] = (remaining * ic_arr_toi).sum(axis=1).sum()\n#             metrics[i, 6] = (mis * ic_arr_toi).sum(axis=1).sum()\n\n    return metrics\n\n\ndef compute_metrics(pred, gt, toi, tau_arr, ic_arr=None, n_cpu=0):\n    \"\"\"\n    Takes the prediction and the ground truth and for each threshold in tau_arr\n    calculates the confusion matrix and returns the coverage,\n    precision, recall, remaining uncertainty and misinformation.\n    Toi is the list of terms (indexes) to be considered\n    \"\"\"\n    g = gt.matrix[:, toi]\n    n_gt = g.sum(axis=1)\n    wn_gt = None\n    if ic_arr is not None:\n        wn_gt = (g * ic_arr[toi]).sum(axis=1)\n\n    # Parallelization\n    if n_cpu == 0:\n        n_cpu = mp.cpu_count()\n\n    arg_lists = [[tau_arr, g, pred, toi, n_gt, wn_gt, ic_arr] for tau_arr in np.array_split(tau_arr, n_cpu)]\n    if 0:\n        # Original parallel way (# It does not work on Kaggle)\n        arg_lists = [[tau_arr, g, pred, toi, n_gt, wn_gt, ic_arr] for tau_arr in np.array_split(tau_arr, n_cpu)]\n        with mp.Pool(processes=n_cpu) as pool:\n            metrics = np.concatenate(pool.starmap(compute_metrics_, arg_lists), axis=0)\n    else: \n        # no-parallel: \n        metrics = compute_metrics_(tau_arr, g, pred, toi, n_gt, wn_gt, ic_arr )\n\n    return pd.DataFrame(metrics, columns=[\"cov\", \"pr\", \"rc\", \"wpr\", \"wrc\", \"ru\", \"mi\"])\n\n\ndef evaluate_prediction(prediction, gt, ontologies, tau_arr, normalization='cafa', n_cpu=0):\n    dfs = []\n    for p in prediction:\n        ns = p.namespace\n        ne = np.full(len(tau_arr), gt[ns].matrix.shape[0])\n\n        ont = [o for o in ontologies if o.namespace == ns][0]\n\n        # cov, pr, rc, wpr, wrc, ru, mi\n        metrics = compute_metrics(p, gt[ns], ont.toi, tau_arr, ont.ia, n_cpu)\n\n        for column in [\"pr\", \"rc\", \"wpr\", \"wrc\", \"ru\", \"mi\"]:\n            if normalization == 'gt' or (column in [\"rc\", \"wrc\"] and normalization == 'cafa'):\n                metrics[column] = np.divide(metrics[column], ne, out=np.zeros_like(metrics[column], dtype='float'), where=ne > 0)\n            else:\n                metrics[column] = np.divide(metrics[column], metrics[\"cov\"], out=np.zeros_like(metrics[column], dtype='float'), where=metrics[\"cov\"] > 0)\n\n        metrics['ns'] = [ns] * len(tau_arr)\n        metrics['tau'] = tau_arr\n        metrics['cov'] = np.divide(metrics['cov'], ne, out=np.zeros_like(metrics['cov'], dtype='float'), where=ne > 0)\n        metrics['f'] = compute_f(metrics['pr'], metrics['rc'])\n        metrics['wf'] = compute_f(metrics['wpr'], metrics['wrc'])\n        metrics['s'] = compute_s(metrics['ru'], metrics['mi'])\n\n        dfs.append(metrics)\n\n    return pd.concat(dfs)\n\n# Tau array, used to compute metrics at different score thresholds\nth_step = 0.01\ntau_arr = np.arange(0.01, 1, th_step)\n#Consider terms without parents, e.g. the root(s), in the evaluation\nno_orphans = False\n# Parse and set information accretion (optional)\nia_dict = ia_parser('/kaggle/input/cafa-5-protein-function-prediction/IA.txt')\n\n# Parse the OBO file and creates a different graph for each namespace\nontologies = []\nobo_file = '/kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo'\nfor ns, terms_dict in obo_parser(obo_file).items():\n    ontologies.append(Graph(ns, terms_dict, ia_dict, not no_orphans))","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:15:26.419646Z","iopub.execute_input":"2023-08-21T12:15:26.420887Z","iopub.status.idle":"2023-08-21T12:15:38.755749Z","shell.execute_reply.started":"2023-08-21T12:15:26.420819Z","shell.execute_reply":"2023-08-21T12:15:38.754084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Wrapper to call CAFA5 metric computation\n\nFunction to call CAFA5 metric computation from current notebook environment.\n\nPay attention - that metric computation is very slow and RAM consuming - be careful !\n\nIt is quite technical - no need to go into details - just use as a blacbox.\n\nThanks to Sergei Fironov https://www.kaggle.com/code/sergeifironov/validate-ridge - please upvote his work. We are based on his code.\n\nhttps://www.kaggle.com/competitions/cafa-5-protein-function-prediction/discussion/420241 - Anton Vakhrushev - correcting error in the initial code of the metric computation - please upvote","metadata":{}},{"cell_type":"code","source":"!head /kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:17:59.313435Z","iopub.execute_input":"2023-08-21T12:17:59.31399Z","iopub.status.idle":"2023-08-21T12:18:00.582328Z","shell.execute_reply.started":"2023-08-21T12:17:59.31395Z","shell.execute_reply":"2023-08-21T12:18:00.580683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nimport os.path\n\n######################################################################################3\n###############  Load trainTerms\n######################################################################################3\n\nprint()\n# fn = '/kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv'\n# print(fn)\n# trainTerms = pd.read_csv(fn, sep=\"\\t\")\ntrainTerms = test_labels\nprint(trainTerms.shape)\nprint('trainTerms memory_usage Mb:', trainTerms.memory_usage().sum()/1e6  )\ndisplay(trainTerms.head(3))\n\ndef get_F1_etc_scores_official_CAFA_evaluation( Y_pred, IX, cutoff_threshold_low = 0.01,    make_plots = True ,  verbose = 0 ): \n    '''\n    Computation of F1-weighted scores are called here. \n    Here we prepare Y_pred, Y in format required by functions provided by organizers - see github: https://github.com/BioComputingUP/CAFA-evaluator\n    Y_pred  -  predictions\n    IX  - indices selecting part which correspond to Y_pred in full Y \n    Params:\n    cutoff_threshold_low - predictions lower (strictly) will be dropped (effectively set to zero)\n        (!) higher cutoff_threshold_low will improve  both RAM/speed. For most models 0.1 is Okay, and even 0.18. Consider using 0.1-0.15.   \n    make_plots = True  - create plots of  F1,precision,recall (weighted) depending on threshold  \n    verbose = 0\n    \n    Function uses external variables: \n    trainTerms - training labels provided by orgs: /kaggle/input/cafa-5-protein-function-prediction/Train/train_terms.tsv\n    vec_train_protein_ids - ids of the proteins in the current train - should correspond to \"X\" - features \n    ontologies - data from: /kaggle/input/cafa-5-protein-function-prediction/Train/go-basic.obo\n    tau_arr - array of thresholds \n    \n    Note: pay attention - Y_true is NOT an input argument. (Unusual for metric computation). \n    We will create Y_true here from \"trainTerms\" cutting only those proteins which correspond to Y_pred indexes \"IX\".\n    Thus it is highly important pass here the correct \"IX\" i.e. corresponding to Y_pred, otherwise results will not be correct\n    '''\n\n    t00 = time.time()\n    if verbose >= 100:\n        print('Scoring starts. n_samples:', len(IX) )\n\n    ##########################################################################################\n    # Prepare \"ground truth\" - \"gt\" terms(labels) in required format  \n    ##########################################################################################\n\n    # First save to file, because function \"gt_parser\" works with files as input \n    # Only part corresponding to providex indices IX will be generated \n    t0 = time.time()\n    trainTerms[ trainTerms.EntryID.isin(vec_train_protein_ids[IX]) ].to_csv('valid.tsv', sep='\\t', index=False) # Wall time: 4.11 s  for 28k samples\n    if verbose >= 1000:\n        print('save valid.csv %.1f'%(time.time() - t0 )) \n\n    # Prepare \"gt\" labels \n    t0 = time.time()\n    gt = gt_parser('valid.tsv', ontologies) # Wall time: 1min 22s  for 28k samples\n    if verbose >= 100:\n        print('gt_parser %.1f'%(time.time() - t0 ))\n\n    ##########################################################################################\n    # prepare predicitons as list of triples - (protein, term(label), prediction) \n    ##########################################################################################\n\n    t0 = time.time()\n    vec_train_protein_ids_loc = vec_train_protein_ids[IX]\n    preds = []\n    for i in range(len(vec_train_protein_ids_loc)):\n        for j in range(Y_pred.shape[1]):\n            if Y_pred[i,j] >= cutoff_threshold_low:\n                preds.append((vec_train_protein_ids_loc[i], \n                              labels_to_consider[j],\n                              Y_pred[i,j]                        ))\n    if verbose >= 1000:            \n        print('create preds %.1f'%(time.time() - t0 ))       \n\n    ##########################################################################################\n    # Parse predictions - propagation happens here  \n    ##########################################################################################\n    t0 = time.time()\n    preds = pred_parser(preds, ontologies, gt, prop_mode='fill', max_terms=500) # \n    if verbose >= 1000:            \n        print('pred_parser %.1f'%(time.time() - t0 ), 'len(preds)', len(preds) )            \n\n    gc.collect()\n\n    ##########################################################################################\n    # Main scores calculations happends here: \n    ##########################################################################################\n    # %%time\n    t0 = time.time()\n    df_metrics = evaluate_prediction(preds, gt, ontologies, tau_arr, n_cpu=1) # Wall time: 37.7 s for 28k samples\n    if verbose >= 1000:            \n        print('evaluate_prediction %.1f'%(time.time() - t0 ), 'got df_metrics with shape:', df_metrics.shape )            \n    if verbose >= 10000:            \n        display( df_metrics.head(2) )\n\n        \n    ##########################################################################################\n    # Comptutations finished. Below are optional plots, output preparartions etc.  \n    ##########################################################################################\n    \n    ##########################################################################################\n    ##########################################################################################\n    ##########################################################################################\n    ##########################################################################################\n    ##########################################################################################\n    \n    \n    if verbose >= 100:\n        _t = df_metrics.groupby('ns').agg({'wf':'max'})\n        display( _t )\n        print( _t.mean() ) \n\n    if verbose >= 100:\n        print('F1-scoring finished. %.1f secs passed'%(time.time() - t00 ))\n\n    # %%time\n    if make_plots:\n        try:\n            list_uv = list(df_metrics['ns'].unique() )\n            #print(list_uv)\n            fig = plt.figure(figsize = (20,4))\n            i0 = 0;\n            for  col in  ['wf', 'wpr', 'wrc' ] :\n                i0+=1\n    #             print(i0,col)\n                fig.add_subplot(1,3,i0)\n\n                for uv in list_uv:\n                    mask = df_metrics['ns'] == uv\n                    v = df_metrics[mask][col]\n                    plt.plot(v.values, label = uv)\n                plt.title(col, fontsize  = 20)\n                plt.legend()\n                plt.grid()\n            plt.show()        \n        except:\n            print('Exception in plot')\n    \n    ########################################################################################\n    # Prepare output of scores : \n    ########################################################################################\n    _t = {'cellular_component':'CCO', 'biological_process':'BPO','molecular_function':'MFO'}\n    dict_scores_etc = {}\n    df_s = df_metrics.groupby('ns').agg({'wf':'max'})\n    dict_scores_etc['F1w'] = np.round( df_s.mean().iloc[0], 6) \n    for k in _t:\n        k2 = _t[k]\n        # print(k,dict_scores_etc )\n        if k in  df_s.index:\n            dict_scores_etc['F1 '+ k2 ] = np.round( df_s.loc[k].iat[0], 6) \n        else:\n            dict_scores_etc['F1 '+ k2 ] = 0\n\n    ########################################################################################\n    # Prepare output of thresholds : \n    ########################################################################################\n    for k in _t:\n        k2 = _t[k]\n        m = df_metrics['ns'] == k\n        if m.sum()>0:\n            IX = df_metrics[m]['wf'].argmax()\n            thres_optimal = df_metrics[m]['tau'].iat[IX]\n            dict_scores_etc['thres '+ k2 ] = thres_optimal\n        else:\n            dict_scores_etc['thres '+ k2 ] = 0\n            \n    dict_scores_etc['F-Scores Time'] = np.round( time.time() - t00   ,1)         \n    if verbose >= 100:\n        print('Scores: ', dict_scores_etc)        \n\n    if os.path.isfile('valid.tsv') :\n        os.remove('valid.tsv')\n        \n    return  dict_scores_etc  ","metadata":{"execution":{"iopub.status.busy":"2023-08-21T12:19:21.444562Z","iopub.execute_input":"2023-08-21T12:19:21.44606Z","iopub.status.idle":"2023-08-21T12:19:21.498635Z","shell.execute_reply.started":"2023-08-21T12:19:21.446005Z","shell.execute_reply":"2023-08-21T12:19:21.497037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\nIX  = np.arange(100)\nresult1 = get_F1_etc_scores_official_CAFA_evaluation( Y_pred[IX,:], IX, cutoff_threshold_low = 0.1,    make_plots = True ,  verbose = 1000 )\nprint(result1)","metadata":{},"execution_count":null,"outputs":[]}]}