{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":41875,"databundleVersionId":5521661,"sourceType":"competition"},{"sourceId":6204903,"sourceType":"datasetVersion","datasetId":3562732}],"dockerImageVersionId":30476,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# DeepGoZero - Introduction\nWe use DeepGOZero to predict functions for proteins using the GO axioms. \n\nSpecifically:\n* We first use the geometric ontology embedding method EL Embeddings (Kulmanov et al., 2019) to generate a space in which GO classes are n-balls in an n-dimensional space and the location and size of the n-balls are constrained by the GO axioms. The ontology axioms are only used during the training phase of DeepGOZero to generate the space constrained by GO axioms.\n* We then use a neural network to project proteins into the same space in which we embedded the GO classes and predict functions for proteins by their proximity and relation to GO classes.\n- Specifically, DeepGOZero uses as input the InterPro domain annotations of a protein where we represent the annotations as a binary vector.\n- The binary vector of InterPro domain annotations is processed by MLP layers to generate an embedding vector of the same size of an embedding vector for GO classes. \n* Then, DeepGOZero jointly minimizes prediction loss for protein functions and the ELEmbeddings loss that impose constraints on the classes. \n","metadata":{}},{"cell_type":"code","source":"# This notebook attempts to explain all of Gerasevas DeepGoZero code (evaluated on original CAFA4)\n# For the benefit of documentation and to explain our modifications later on\n\nimport numpy as np \nimport pandas as pd \nfrom IPython.display import clear_output\nfrom tqdm import tqdm\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:38.159387Z","iopub.execute_input":"2024-05-24T12:44:38.159682Z","iopub.status.idle":"2024-05-24T12:44:38.20512Z","shell.execute_reply.started":"2024-05-24T12:44:38.159656Z","shell.execute_reply":"2024-05-24T12:44:38.204147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# deep graph lib\n!pip install dgl==1.1.1","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:38.206798Z","iopub.execute_input":"2024-05-24T12:44:38.207073Z","iopub.status.idle":"2024-05-24T12:44:53.76184Z","shell.execute_reply.started":"2024-05-24T12:44:38.207049Z","shell.execute_reply":"2024-05-24T12:44:53.76079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch as th\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch import optim\nfrom torch.optim.lr_scheduler import MultiStepLR\nfrom sklearn.metrics import roc_curve, auc, matthews_corrcoef\nimport copy\nfrom torch.utils.data import DataLoader, IterableDataset, TensorDataset\nfrom itertools import cycle\nimport math\nfrom dgl.nn import GraphConv, GATConv\nimport dgl\nfrom collections import deque, Counter","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:53.763213Z","iopub.execute_input":"2024-05-24T12:44:53.763517Z","iopub.status.idle":"2024-05-24T12:44:57.534106Z","shell.execute_reply.started":"2024-05-24T12:44:53.763487Z","shell.execute_reply":"2024-05-24T12:44:57.53314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Hyperparameter Definition\n","metadata":{}},{"cell_type":"code","source":"#Parameters\ndata_root='/kaggle/input/deepgozero-data/data'\nont='bp' # Ontology class definition, because DGZ has models for each class. mf, bp, or cc\ndevice='cuda:0' \nbatch_size=37 # batch size currently not used for inference, would be used training only\nepochs=25 # Training epochs currently not used for inference, would be used training only\nload=False # If False, attempt retraining retrain. Note not relevant for CAFA5, keep True for inference\n# Path definitions for various files\ngo_file = f'{data_root}/go.norm'  # GO file. contains 10490 GO annotations.\nmodel_file = '/kaggle/working/deepgozero_zero_10.th' #f'{data_root}/{ont}/deepgozero_zero_10.th' # Keep this if youre inferring from saved model.\nterms_file = f'{data_root}/{ont}/terms_zero_10.pkl'\nout_file = '/kaggle/working/predictions_deepgozero_zero_10.pkl'","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:57.536203Z","iopub.execute_input":"2024-05-24T12:44:57.536932Z","iopub.status.idle":"2024-05-24T12:44:57.54352Z","shell.execute_reply.started":"2024-05-24T12:44:57.536903Z","shell.execute_reply":"2024-05-24T12:44:57.542568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DeepGoZero Repository ","metadata":{}},{"cell_type":"markdown","source":"We use code from https://github.com/bio-ontology-research-group/deepgozero. Instead of cloning the repo we copy all necessary code into the notebook and run it.\n\n","metadata":{}},{"cell_type":"code","source":"# from https://github.com/bio-ontology-research-group/deepgozero/blob/main/utils.py\nAALETTER = [\n    'A', 'R', 'N', 'D', 'C', 'Q', 'E', 'G', 'H', 'I',\n    'L', 'K', 'M', 'F', 'P', 'S', 'T', 'W', 'Y', 'V']\nAANUM = len(AALETTER)\nAAINDEX = dict()\nfor i in range(len(AALETTER)):\n    AAINDEX[AALETTER[i]] = i + 1\nINVALID_ACIDS = set(['U', 'O', 'B', 'Z', 'J', 'X', '*'])\nMAXLEN = 2000\nNGRAMS = {}\nfor i in range(20):\n    for j in range(20):\n        for k in range(20):\n            ngram = AALETTER[i] + AALETTER[j] + AALETTER[k]\n            index = 400 * i + 20 * j + k + 1\n            NGRAMS[ngram] = index\n\ndef is_ok(seq):\n    for c in seq:\n        if c in INVALID_ACIDS:\n            return False\n    return True\n\ndef to_ngrams(seq):\n    l = min(MAXLEN, len(seq) - 3)\n    ngrams = np.zeros((l,), dtype=np.int32)\n    for i in range(l):\n        ngrams[i] = NGRAMS.get(seq[i: i + 3], 0)\n    return ngrams\n\ndef to_tokens(seq):\n    tokens = np.zeros((MAXLEN, ), dtype=np.float32)\n    l = min(MAXLEN, len(seq))\n    for i in range(l):\n        tokens[i] = AAINDEX.get(seq[i], 0)\n    return tokens\n\ndef to_onehot(seq, start=0):\n    onehot = np.zeros((21, MAXLEN), dtype=np.float32)\n    l = min(MAXLEN, len(seq))\n    for i in range(start, start + l):\n        onehot[AAINDEX.get(seq[i - start], 0), i] = 1\n    onehot[0, 0:start] = 1\n    onehot[0, start + l:] = 1\n    return onehot","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:57.547617Z","iopub.execute_input":"2024-05-24T12:44:57.547952Z","iopub.status.idle":"2024-05-24T12:44:57.570238Z","shell.execute_reply.started":"2024-05-24T12:44:57.547927Z","shell.execute_reply":"2024-05-24T12:44:57.569451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from https://github.com/bio-ontology-research-group/deepgozero/blob/main/utils.py\n# Ontology definition and methods.\n\nclass Ontology(object):\n    # Loads data, optinonally with relationships\n    def __init__(self, filename='data/go.obo', with_rels=False):\n        self.ont = self.load(filename, with_rels)\n        self.ic = None\n        self.ic_norm = 0.0\n    # Check that term exists in ontology\n    def has_term(self, term_id):\n        return term_id in self.ont\n    # Check and retrieve term\n    def get_term(self, term_id):\n        if self.has_term(term_id):\n            return self.ont[term_id]\n        return None\n    \n    # Calculates the information content (IC)\n    # of each term based on the annotations provided. \n    # IC measures the specificity of a term within an ontology.\n    def calculate_ic(self, annots):\n        cnt = Counter()\n        for x in annots:\n            cnt.update(x)\n        self.ic = {}\n        for go_id, n in cnt.items():\n            parents = self.get_parents(go_id)\n            if len(parents) == 0:\n                min_n = n\n            else:\n                min_n = min([cnt[x] for x in parents])\n\n            self.ic[go_id] = math.log(min_n / n, 2)\n            self.ic_norm = max(self.ic_norm, self.ic[go_id])\n    \n    # Retrieves the IC value for a given GO term ID.   \n    def get_ic(self, go_id):\n        if self.ic is None:\n            raise Exception('Not yet calculated')\n        if go_id not in self.ic:\n            return 0.0\n        return self.ic[go_id]\n    \n    # Retrieves the normalized IC value for a given GO term ID\n    def get_norm_ic(self, go_id):\n        return self.get_ic(go_id) / self.ic_norm\n\n    # Loads ontology data from a file into memory. \n    # It parses the data and organizes it into a dictionary structure \n    # representing the ontology.\n    def load(self, filename, with_rels):\n        ont = dict()\n        obj = None\n        with open(filename, 'r') as f:\n            for line in f:\n                line = line.strip()\n                if not line:\n                    continue\n                if line == '[Term]':\n                    if obj is not None:\n                        ont[obj['id']] = obj\n                    obj = dict()\n                    obj['is_a'] = list()\n                    obj['part_of'] = list()\n                    obj['regulates'] = list()\n                    obj['alt_ids'] = list()\n                    obj['is_obsolete'] = False\n                    continue\n                elif line == '[Typedef]':\n                    if obj is not None:\n                        ont[obj['id']] = obj\n                    obj = None\n                else:\n                    if obj is None:\n                        continue\n                    l = line.split(\": \")\n                    if l[0] == 'id':\n                        obj['id'] = l[1]\n                    elif l[0] == 'alt_id':\n                        obj['alt_ids'].append(l[1])\n                    elif l[0] == 'namespace':\n                        obj['namespace'] = l[1]\n                    elif l[0] == 'is_a':\n                        obj['is_a'].append(l[1].split(' ! ')[0])\n                    elif with_rels and l[0] == 'relationship':\n                        it = l[1].split()\n                        # add all types of relationships\n                        obj['is_a'].append(it[1])\n                    elif l[0] == 'name':\n                        obj['name'] = l[1]\n                    elif l[0] == 'is_obsolete' and l[1] == 'true':\n                        obj['is_obsolete'] = True\n            if obj is not None:\n                ont[obj['id']] = obj\n        for term_id in list(ont.keys()):\n            for t_id in ont[term_id]['alt_ids']:\n                ont[t_id] = ont[term_id]\n            if ont[term_id]['is_obsolete']:\n                del ont[term_id]\n        for term_id, val in ont.items():\n            if 'children' not in val:\n                val['children'] = set()\n            for p_id in val['is_a']:\n                if p_id in ont:\n                    if 'children' not in ont[p_id]:\n                        ont[p_id]['children'] = set()\n                    ont[p_id]['children'].add(term_id)\n     \n        return ont\n    \n    # retrieves all ancestors (immediate and distant) of a given term.\n    def get_anchestors(self, term_id):\n        if term_id not in self.ont:\n            return set()\n        term_set = set()\n        q = deque()\n        q.append(term_id)\n        while(len(q) > 0):\n            t_id = q.popleft()\n            if t_id not in term_set:\n                term_set.add(t_id)\n                for parent_id in self.ont[t_id]['is_a']:\n                    if parent_id in self.ont:\n                        q.append(parent_id)\n        return term_set\n    \n    # Retrieves all terms that are propagated up the ontology hierarchy from the given set of terms.\n    def get_prop_terms(self, terms):\n        prop_terms = set()\n\n        for term_id in terms:\n            prop_terms |= self.get_anchestors(term_id)\n        return prop_terms\n\n    #  Retrieves immediate parent terms of a given term.\n    def get_parents(self, term_id):\n        if term_id not in self.ont:\n            return set()\n        term_set = set()\n        for parent_id in self.ont[term_id]['is_a']:\n            if parent_id in self.ont:\n                term_set.add(parent_id)\n        return term_set\n\n    # Retrieves all terms belonging to a specified namespace.\n    def get_namespace_terms(self, namespace):\n        terms = set()\n        for go_id, obj in self.ont.items():\n            if obj['namespace'] == namespace:\n                terms.add(go_id)\n        return terms\n    \n    # Retrieves the namespace of a given term.\n    def get_namespace(self, term_id):\n        return self.ont[term_id]['namespace']\n    \n    # Retrieves all terms in the subgraph rooted at the given term.   \n    def get_term_set(self, term_id):\n        if term_id not in self.ont:\n            return set()\n        term_set = set()\n        q = deque()\n        q.append(term_id)\n        while len(q) > 0:\n            t_id = q.popleft()\n            if t_id not in term_set:\n                term_set.add(t_id)\n                for ch_id in self.ont[t_id]['children']:\n                    q.append(ch_id)\n        return term_set\n","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2024-05-24T12:44:57.57283Z","iopub.execute_input":"2024-05-24T12:44:57.573178Z","iopub.status.idle":"2024-05-24T12:44:57.603836Z","shell.execute_reply.started":"2024-05-24T12:44:57.573146Z","shell.execute_reply":"2024-05-24T12:44:57.602907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from https://github.com/bio-ontology-research-group/deepgozero/blob/main/deepgozero.py\n    \ndef compute_roc(labels, preds):\n    # Compute ROC curve and ROC area for each class\n    fpr, tpr, _ = roc_curve(labels.flatten(), preds.flatten())\n    roc_auc = auc(fpr, tpr)\n\n    return roc_auc\n\n# Loads normal forms (NF) from a Gene Ontology (GO) file.\n# Normal forms are alternative representations of ontology relationships used in reasoning and inference tasks.\n\ndef load_normal_forms(go_file, terms_dict):\n    nf1 = []\n    nf2 = []\n    nf3 = []\n    nf4 = []\n    relations = {}\n    zclasses = {}\n    \n    def get_index(go_id):\n        if go_id in terms_dict:\n            index = terms_dict[go_id]\n        elif go_id in zclasses:\n            index = zclasses[go_id]\n        else:\n            zclasses[go_id] = len(terms_dict) + len(zclasses)\n            index = zclasses[go_id]\n        return index\n\n    def get_rel_index(rel_id):\n        if rel_id not in relations:\n            relations[rel_id] = len(relations)\n        return relations[rel_id]\n                \n    with open(go_file) as f:\n        for line in f:\n            line = line.strip().replace('_', ':')\n            if line.find('SubClassOf') == -1:\n                continue\n            left, right = line.split(' SubClassOf ')\n            # C SubClassOf D\n            if len(left) == 10 and len(right) == 10:\n                go1, go2 = left, right\n                nf1.append((get_index(go1), get_index(go2)))\n            elif left.find('and') != -1: # C and D SubClassOf E\n                go1, go2 = left.split(' and ')\n                go3 = right\n                nf2.append((get_index(go1), get_index(go2), get_index(go3)))\n            elif left.find('some') != -1:  # R some C SubClassOf D\n                rel, go1 = left.split(' some ')\n                go2 = right\n                nf3.append((get_rel_index(rel), get_index(go1), get_index(go2)))\n            elif right.find('some') != -1: # C SubClassOf R some D\n                go1 = left\n                rel, go2 = right.split(' some ')\n                nf4.append((get_index(go1), get_rel_index(rel), get_index(go2)))\n    return nf1, nf2, nf3, nf4, relations, zclasses \n\n\n# Construct dicitionaries from input data and labels\n\ndef load_data(data_root, ont, terms_file):\n    terms_df = pd.read_pickle(terms_file)\n    terms = terms_df['gos'].values.flatten()\n    terms_dict = {v: i for i, v in enumerate(terms)}\n    print('Terms', len(terms))\n    \n    ipr_df = pd.read_pickle(f'{data_root}/{ont}/interpros.pkl')\n    iprs = ipr_df['interpros'].values\n    iprs_dict = {v:k for k, v in enumerate(iprs)}\n    return iprs_dict, terms_dict\n\n# Format data and labels into tensors\n\ndef get_data(df, iprs_dict, terms_dict):\n    data = th.zeros((len(df), len(iprs_dict)), dtype=th.float32)\n    labels = th.zeros((len(df), len(terms_dict)), dtype=th.float32)\n    for i, row in enumerate(df.itertuples()):\n        for ipr in row.interpros:\n            if ipr in iprs_dict:\n                data[i, iprs_dict[ipr]] = 1\n        for go_id in row.prop_annotations: # prop_annotations for full model\n            if go_id in terms_dict:\n                g_id = terms_dict[go_id]\n                labels[i, g_id] = 1\n    return data, labels\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-05-24T12:44:57.605195Z","iopub.execute_input":"2024-05-24T12:44:57.605482Z","iopub.status.idle":"2024-05-24T12:44:57.628123Z","shell.execute_reply.started":"2024-05-24T12:44:57.605459Z","shell.execute_reply":"2024-05-24T12:44:57.627057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from https://github.com/bio-ontology-research-group/deepgozero/blob/main/deepgozero.py\n# Model Structure\nclass Residual(nn.Module):\n\n    def __init__(self, fn):\n        super().__init__()\n        self.fn = fn\n\n    def forward(self, x):\n        return x + self.fn(x)\n    \n        \nclass MLPBlock(nn.Module):\n\n    def __init__(self, in_features, out_features, bias=True, layer_norm=True, dropout=0.1, activation=nn.ReLU):\n        super().__init__()\n        self.linear = nn.Linear(in_features, out_features, bias)\n        self.activation = activation()\n        self.layer_norm = nn.BatchNorm1d(out_features, track_running_stats=False) if layer_norm else None\n        self.dropout = nn.Dropout(dropout) if dropout else None\n\n    def forward(self, x):\n        x = self.activation(self.linear(x))\n        if self.layer_norm:\n            x = self.layer_norm(x)\n        if self.dropout:\n            x = self.dropout(x)\n        return x\n\n# Model structure, \n\nclass DGELModel(nn.Module):\n\n    def __init__(self, nb_iprs, nb_gos, nb_zero_gos, nb_rels, device, hidden_dim=1024, embed_dim=1024, margin=0.1):\n        super().__init__()\n        self.nb_gos = nb_gos\n        self.nb_zero_gos = nb_zero_gos\n        input_length = nb_iprs\n        net = []\n        net.append(MLPBlock(input_length, hidden_dim))\n        net.append(Residual(MLPBlock(hidden_dim, hidden_dim)))\n        self.net = nn.Sequential(*net)\n\n        # ELEmbeddings\n        self.embed_dim = embed_dim\n        #  CGPT : hasFuncIndex represents whether a GO term has a function index?\n        self.hasFuncIndex = th.LongTensor([nb_rels]).to(device)\n        # This line is key to creating an embedding constrained by Go Size  (as per origianl dgz explanation)\n        self.go_embed = nn.Embedding(nb_gos + nb_zero_gos, embed_dim)\n        self.go_norm = nn.BatchNorm1d(embed_dim)\n        k = math.sqrt(1 / embed_dim)\n        nn.init.uniform_(self.go_embed.weight, -k, k)\n        self.go_rad = nn.Embedding(nb_gos + nb_zero_gos, 1)\n        nn.init.uniform_(self.go_rad.weight, -k, k)\n        # self.go_embed.weight.requires_grad = False\n        # self.go_rad.weight.requires_grad = False\n        \n        self.rel_embed = nn.Embedding(nb_rels + 1, embed_dim)\n        nn.init.uniform_(self.rel_embed.weight, -k, k)\n        self.all_gos = th.arange(self.nb_gos).to(device)\n        self.margin = margin\n\n    # Forward method\n    # Features are from FastTensorDataLoader, which is originally from get_data\n\n    def forward(self, features):\n        x = self.net(features)\n        # Use the ELEmbedding to embed GO terms according to our go-constrained embedding layer.\n        go_embed = self.go_embed(self.all_gos)\n        hasFunc = self.rel_embed(self.hasFuncIndex)\n        hasFuncGO = go_embed + hasFunc\n        go_rad = th.abs(self.go_rad(self.all_gos).view(1, -1))\n        x = th.matmul(x, hasFuncGO.T) + go_rad\n        logits = th.sigmoid(x)\n        return logits\n\n    def predict_zero(self, features, data):\n        x = self.net(features)\n        go_embed = self.go_embed(data)\n        hasFunc = self.rel_embed(self.hasFuncIndex)\n        hasFuncGO = go_embed + hasFunc\n        go_rad = th.abs(self.go_rad(data).view(1, -1))\n        x = th.matmul(x, hasFuncGO.T) + go_rad\n        logits = th.sigmoid(x)\n        return logits\n\n\n    def el_loss(self, go_normal_forms):\n        nf1, nf2, nf3, nf4 = go_normal_forms\n        nf1_loss = self.nf1_loss(nf1)\n        nf2_loss = self.nf2_loss(nf2)\n        nf3_loss = self.nf3_loss(nf3)\n        nf4_loss = self.nf4_loss(nf4)\n        # print()\n        # print(nf1_loss.detach().item(),\n        #       nf2_loss.detach().item(),\n        #       nf3_loss.detach().item(),\n        #       nf4_loss.detach().item())\n        return nf1_loss + nf3_loss + nf4_loss + nf2_loss\n\n    def class_dist(self, data):\n        c = self.go_norm(self.go_embed(data[:, 0]))\n        d = self.go_norm(self.go_embed(data[:, 1]))\n        rc = th.abs(self.go_rad(data[:, 0]))\n        rd = th.abs(self.go_rad(data[:, 1]))\n        dist = th.linalg.norm(c - d, dim=1, keepdim=True) + rc - rd\n        return dist\n        \n    def nf1_loss(self, data):\n        pos_dist = self.class_dist(data)\n        loss = th.mean(th.relu(pos_dist - self.margin))\n        return loss\n\n    def nf2_loss(self, data):\n        c = self.go_norm(self.go_embed(data[:, 0]))\n        d = self.go_norm(self.go_embed(data[:, 1]))\n        e = self.go_norm(self.go_embed(data[:, 2]))\n        rc = th.abs(self.go_rad(data[:, 0]))\n        rd = th.abs(self.go_rad(data[:, 1]))\n        re = th.abs(self.go_rad(data[:, 2]))\n        \n        sr = rc + rd\n        dst = th.linalg.norm(c - d, dim=1, keepdim=True)\n        dst2 = th.linalg.norm(e - c, dim=1, keepdim=True)\n        dst3 = th.linalg.norm(e - d, dim=1, keepdim=True)\n        loss = th.mean(th.relu(dst - sr - self.margin)\n                    + th.relu(dst2 - rc - self.margin)\n                    + th.relu(dst3 - rd - self.margin))\n\n        return loss\n\n    def nf3_loss(self, data):\n        # R some C subClassOf D\n        n = data.shape[0]\n        # rS = self.rel_space(data[:, 0])\n        # rS = rS.reshape(-1, self.embed_dim, self.embed_dim)\n        rE = self.rel_embed(data[:, 0])\n        c = self.go_norm(self.go_embed(data[:, 1]))\n        d = self.go_norm(self.go_embed(data[:, 2]))\n        # c = th.matmul(c, rS).reshape(n, -1)\n        # d = th.matmul(d, rS).reshape(n, -1)\n        rc = th.abs(self.go_rad(data[:, 1]))\n        rd = th.abs(self.go_rad(data[:, 2]))\n        \n        rSomeC = c + rE\n        euc = th.linalg.norm(rSomeC - d, dim=1, keepdim=True)\n        loss = th.mean(th.relu(euc + rc - rd - self.margin))\n        return loss\n\n\n    def nf4_loss(self, data):\n        # C subClassOf R some D\n        n = data.shape[0]\n        c = self.go_norm(self.go_embed(data[:, 0]))\n        rE = self.rel_embed(data[:, 1])\n        d = self.go_norm(self.go_embed(data[:, 2]))\n        \n        rc = th.abs(self.go_rad(data[:, 1]))\n        rd = th.abs(self.go_rad(data[:, 2]))\n        sr = rc + rd\n        # c should intersect with d + r\n        rSomeD = d + rE\n        dst = th.linalg.norm(c - rSomeD, dim=1, keepdim=True)\n        loss = th.mean(th.relu(dst - sr - self.margin))\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:57.62961Z","iopub.execute_input":"2024-05-24T12:44:57.629884Z","iopub.status.idle":"2024-05-24T12:44:57.663223Z","shell.execute_reply.started":"2024-05-24T12:44:57.629861Z","shell.execute_reply":"2024-05-24T12:44:57.662355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from https://github.com/bio-ontology-research-group/deepgozero/blob/main/torch_utils.py\n\nimport torch\n\nclass FastTensorDataLoader:\n\n    def __init__(self, *tensors, batch_size=32, shuffle=False):\n\n        assert all(t.shape[0] == tensors[0].shape[0] for t in tensors)\n        self.tensors = tensors\n\n        self.dataset_len = self.tensors[0].shape[0]\n        self.batch_size = batch_size\n        self.shuffle = shuffle\n\n        # Calculate # batches\n        n_batches, remainder = divmod(self.dataset_len, self.batch_size)\n        if remainder > 0:\n            n_batches += 1\n        self.n_batches = n_batches\n    def __iter__(self):\n        if self.shuffle:\n            r = torch.randperm(self.dataset_len)\n            self.tensors = [t[r] for t in self.tensors]\n        self.i = 0\n        return self\n\n    def __next__(self):\n        if self.i >= self.dataset_len:\n            raise StopIteration\n        batch = tuple(t[self.i:self.i+self.batch_size] for t in self.tensors)\n        self.i += self.batch_size\n        return batch\n\n    def __len__(self):\n        return self.n_batches","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:57.666427Z","iopub.execute_input":"2024-05-24T12:44:57.66674Z","iopub.status.idle":"2024-05-24T12:44:57.678565Z","shell.execute_reply.started":"2024-05-24T12:44:57.666717Z","shell.execute_reply":"2024-05-24T12:44:57.677746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Run DeepGoZero","metadata":{}},{"cell_type":"code","source":"# from https://github.com/bio-ontology-research-group/deepgozero/blob/main/deepgozero.py\n\nloss_func = nn.BCELoss()\niprs_dict, terms_dict = load_data(data_root, ont, terms_file)\nn_terms = len(terms_dict)\nn_iprs = len(iprs_dict)\n    \nnf1, nf2, nf3, nf4, relations, zero_classes = load_normal_forms(go_file, terms_dict)\nn_rels = len(relations)\nn_zeros = len(zero_classes)\n    \nnormal_forms = nf1, nf2, nf3, nf4\nnf1 = th.LongTensor(nf1).to(device)\nnf2 = th.LongTensor(nf2).to(device)\nnf3 = th.LongTensor(nf3).to(device)\nnf4 = th.LongTensor(nf4).to(device)\nnormal_forms = nf1, nf2, nf3, nf4\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:44:57.679751Z","iopub.execute_input":"2024-05-24T12:44:57.680324Z","iopub.status.idle":"2024-05-24T12:45:01.190695Z","shell.execute_reply.started":"2024-05-24T12:44:57.68026Z","shell.execute_reply":"2024-05-24T12:45:01.18976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('Number of terms:',n_terms)\nprint('Number of interpros:',n_iprs)","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:01.192344Z","iopub.execute_input":"2024-05-24T12:45:01.192633Z","iopub.status.idle":"2024-05-24T12:45:01.197137Z","shell.execute_reply.started":"2024-05-24T12:45:01.192608Z","shell.execute_reply":"2024-05-24T12:45:01.196322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net = DGELModel(n_iprs, n_terms, n_zeros, n_rels, device).to(device)\nprint(net)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:01.198284Z","iopub.execute_input":"2024-05-24T12:45:01.198618Z","iopub.status.idle":"2024-05-24T12:45:02.371184Z","shell.execute_reply.started":"2024-05-24T12:45:01.198593Z","shell.execute_reply":"2024-05-24T12:45:02.370263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load=False\nif not load:\n    train_df = pd.read_pickle(f'{data_root}/{ont}/train_data.pkl')\n    train_data = get_data(train_df, iprs_dict, terms_dict)\n    print(train_data[0].shape)\n    train_loader = FastTensorDataLoader(\n            *train_data, batch_size=batch_size, shuffle=True)\n    #del train_df,train_data  # temporary deletion for inspection\n\n    valid_df = pd.read_pickle(f'{data_root}/{ont}/valid_data.pkl')\n    valid_data = get_data(valid_df, iprs_dict, terms_dict)\n    print(valid_data[0].shape)\n    valid_loader = FastTensorDataLoader(\n            *valid_data, batch_size=batch_size, shuffle=False) \n\n    #del valid_df, valid_data # temporary deletion for inspecetion","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:02.372426Z","iopub.execute_input":"2024-05-24T12:45:02.372757Z","iopub.status.idle":"2024-05-24T12:45:25.299786Z","shell.execute_reply.started":"2024-05-24T12:45:02.372726Z","shell.execute_reply":"2024-05-24T12:45:25.298801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.columns","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:25.310092Z","iopub.execute_input":"2024-05-24T12:45:25.310932Z","iopub.status.idle":"2024-05-24T12:45:25.318608Z","shell.execute_reply.started":"2024-05-24T12:45:25.310874Z","shell.execute_reply":"2024-05-24T12:45:25.317524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:25.31995Z","iopub.execute_input":"2024-05-24T12:45:25.320333Z","iopub.status.idle":"2024-05-24T12:45:25.4328Z","shell.execute_reply.started":"2024-05-24T12:45:25.3203Z","shell.execute_reply":"2024-05-24T12:45:25.431818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:25.433981Z","iopub.execute_input":"2024-05-24T12:45:25.43423Z","iopub.status.idle":"2024-05-24T12:45:25.46755Z","shell.execute_reply.started":"2024-05-24T12:45:25.434208Z","shell.execute_reply":"2024-05-24T12:45:25.466619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:25.468565Z","iopub.execute_input":"2024-05-24T12:45:25.468848Z","iopub.status.idle":"2024-05-24T12:45:25.836466Z","shell.execute_reply.started":"2024-05-24T12:45:25.468824Z","shell.execute_reply":"2024-05-24T12:45:25.835448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not load:\n    optimizer = th.optim.Adam(net.parameters(), lr=5e-4)\n    scheduler = MultiStepLR(optimizer, milestones=[5, 20], gamma=0.1)\n    best_loss = 10000.0\n\n    print('Training the model')\n    for epoch in range(epochs):\n        net.train()\n        train_loss = 0\n        train_elloss = 0\n        lmbda = 0.1\n        train_steps = len(train_loader)\n        for batch_features, batch_labels in train_loader:\n            batch_features = batch_features.to(device)\n            batch_labels = batch_labels.to(device)\n            logits = net(batch_features)\n            loss = F.binary_cross_entropy(logits, batch_labels)\n            el_loss = net.el_loss(normal_forms)\n            total_loss = loss + el_loss\n            train_loss += loss.detach().item()\n            train_elloss = el_loss.detach().item()\n            optimizer.zero_grad()\n            total_loss.backward()\n            optimizer.step()\n                    \n        train_loss /= train_steps\n                \n        net.eval()\n        with th.no_grad():\n            valid_steps = len(valid_loader)\n            valid_loss = 0\n            preds = []\n            valid_labels = []\n            for batch_features, batch_labels in valid_loader:\n                batch_features = batch_features.to(device)\n                batch_labels = batch_labels.to(device)\n                logits = net(batch_features)\n                batch_loss = F.binary_cross_entropy(logits, batch_labels)\n                valid_loss += batch_loss.detach().item()\n                preds = np.append(preds, logits.detach().cpu().numpy())\n                valid_labels = np.append(valid_labels, batch_labels.detach().cpu().numpy())\n            valid_loss /= valid_steps\n            roc_auc = compute_roc(valid_labels, preds)\n            print(f'Epoch {epoch}: Loss - {train_loss}, EL Loss: {train_elloss}, Valid loss - {valid_loss}, AUC - {roc_auc}')\n    \n            print('EL Loss', train_elloss)\n            if valid_loss < best_loss:\n                best_loss = valid_loss\n                print('Saving model')\n                th.save(net.state_dict(), model_file)\n\n            scheduler.step()\n            \n","metadata":{"execution":{"iopub.status.busy":"2024-05-24T12:45:25.838121Z","iopub.execute_input":"2024-05-24T12:45:25.838757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_pickle(f'{data_root}/{ont}/test_data.pkl')\ntest_data = get_data(test_df, iprs_dict, terms_dict)\n\ntest_loader = FastTensorDataLoader(\n        *test_data, batch_size=batch_size, shuffle=False)\n    \ntest_labels = test_data[1].detach().cpu().numpy()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"go = Ontology(f'{data_root}/go.obo', with_rels=True)\n\n# Loading best model\nprint('Loading the best model')\nnet.load_state_dict(th.load(model_file, map_location=device))\nnet.eval()\n\nwith th.no_grad():\n    test_steps = int(math.ceil(len(test_labels) / batch_size))\n    test_loss = 0\n    preds = []\n    for batch_features, batch_labels in tqdm(test_loader,total=len(test_loader)):\n        batch_features = batch_features.to(device)\n        batch_labels = batch_labels.to(device)\n        logits = net(batch_features)\n        batch_loss = F.binary_cross_entropy(logits, batch_labels)\n        test_loss += batch_loss.detach().cpu().item()\n        preds = np.append(preds, logits.detach().cpu().numpy())\n    test_loss /= test_steps\n    preds = preds.reshape(-1, n_terms)\n    roc_auc = compute_roc(test_labels, preds)\n    print(f'Test Loss - {test_loss}, AUC - {roc_auc}')\n\n        \npreds = list(preds)\n# Propagate scores using ontology structure\nfor i, scores in tqdm(enumerate(preds), total=len(preds)):\n    prop_annots = {}\n    for go_id, j in terms_dict.items():\n        score = scores[j]\n        for sup_go in go.get_anchestors(go_id):\n            if sup_go in prop_annots:\n                prop_annots[sup_go] = max(prop_annots[sup_go], score)\n            else:\n                prop_annots[sup_go] = score\n    for go_id, score in prop_annots.items():\n        if go_id in terms_dict:\n            scores[terms_dict[go_id]] = score\n\ntest_df['preds'] = preds\n\ntest_df.to_pickle(out_file)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head(10)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}