| """
|
| Module: load_data.py
|
| Description:
|
| This module provides functions for loading heterogeneous networks for drug repositioning.
|
| It supports two datasets: 'Bdataset' and 'Kdataset'. The functions create DGL heterographs
|
| from CSV files containing various interactions and associations, and also assign initial node features.
|
| """
|
|
|
| import dgl
|
| import torch as th
|
| import numpy as np
|
| import pandas as pd
|
| import os
|
| from collections import defaultdict
|
|
|
|
|
| def load(dataset):
|
| """
|
| Load the heterogeneous network for a given dataset.
|
|
|
| Parameters:
|
| dataset (str): The dataset identifier. Options are 'Bdataset' or 'Kdataset'.
|
|
|
| Returns:
|
| dgl.DGLHeteroGraph: The constructed heterogeneous graph.
|
| """
|
| if dataset == "Bdataset":
|
| return load_Bdataset()
|
| if dataset == "Kdataset":
|
| return load_Kdataset()
|
| if dataset == "KGdataset":
|
| return load_KGdataset()
|
| if dataset == 'KGdataset_tiny':
|
| return load_KGdataset_tiny()
|
| raise ValueError("Unsupported dataset. Please choose 'Bdataset', 'Kdataset', or 'KGdataset'.")
|
|
|
|
|
| def _load_node_indices(node_path, node_files):
|
| node_indices = {}
|
| node_count = {}
|
| for node_type, filename in node_files.items():
|
| filepath = os.path.join(node_path, filename)
|
| if os.path.exists(filepath):
|
| df = pd.read_csv(filepath, usecols=['Inter_ID'], low_memory=False)
|
| node_indices[node_type] = pd.Index(df['Inter_ID'])
|
| node_count[node_type] = len(df)
|
| print(f"Loaded {node_type}: {len(df)} nodes")
|
| else:
|
| print(f"Warning: Node file {filepath} not found")
|
| node_indices[node_type] = pd.Index([])
|
| node_count[node_type] = 0
|
| return node_indices, node_count
|
|
|
|
|
| def _build_heterograph(node_path, edge_path, node_files, edge_files, feature_dim=128):
|
|
|
| node_indices, node_count = _load_node_indices(node_path, node_files)
|
|
|
| graph_data = {}
|
|
|
| for edge_type, filename in edge_files.items():
|
| filepath = os.path.join(edge_path, filename)
|
| if not os.path.exists(filepath):
|
| print(f"Warning: Edge file {filepath} not found")
|
| continue
|
|
|
| try:
|
| df = pd.read_csv(filepath, usecols=['x_id', 'y_id'], low_memory=False)
|
|
|
|
|
| src_type, dst_type = edge_type.split('_', 1)
|
| print(f"Loading {edge_type}: {len(df)} edges")
|
| print(f" Source type: {src_type}, Target type: {dst_type}")
|
|
|
|
|
| src_idx_arr = node_indices[src_type].get_indexer(df['x_id'])
|
| dst_idx_arr = node_indices[dst_type].get_indexer(df['y_id'])
|
|
|
|
|
| valid_mask = (src_idx_arr != -1) & (dst_idx_arr != -1)
|
| if not np.any(valid_mask):
|
| print(" No valid edges found")
|
| continue
|
|
|
| s = th.as_tensor(src_idx_arr[valid_mask], dtype=th.long)
|
| d = th.as_tensor(dst_idx_arr[valid_mask], dtype=th.long)
|
|
|
|
|
| graph_data[(src_type, edge_type, dst_type)] = (s, d)
|
|
|
|
|
| if src_type != dst_type:
|
| rev = f"{dst_type}_{src_type}"
|
| graph_data[(dst_type, rev, src_type)] = (d, s)
|
|
|
| print(f" Added {s.numel()} edges")
|
|
|
| except Exception as e:
|
| print(f"Error loading {edge_type}: {e}")
|
|
|
|
|
| g = dgl.heterograph(graph_data)
|
| print(f"Created graph with {g.num_nodes()} total nodes and {g.num_edges()} total edges")
|
|
|
|
|
| for ntype in g.ntypes:
|
| print(f"{ntype}: {g.num_nodes(ntype)} nodes")
|
| for etype in g.etypes:
|
| print(f"{etype}: {g.num_edges(etype)} edges")
|
|
|
|
|
| for ntype in g.ntypes:
|
| num_nodes = g.num_nodes(ntype)
|
| g.nodes[ntype].data['h'] = th.randn(num_nodes, feature_dim, dtype=th.float32)
|
|
|
| return g
|
|
|
|
|
| def load_KGdataset_tiny():
|
| """
|
| Load the heterogeneous network for the tiny knowledge graph dataset (drug, protein, disease).
|
| """
|
| node_path = "/vast/yg3191/AIVS/kg/node"
|
| edge_path = "/vast/yg3191/AIVS/kg/edge"
|
|
|
| node_files = {
|
| 'drug': 'drug.csv',
|
| 'protein': 'protein.csv',
|
| 'disease': 'disease.csv',
|
| }
|
|
|
| edge_files = {
|
| 'drug_drug': 'drug_drug.csv',
|
| 'drug_protein': 'drug_protein.csv',
|
| 'protein_protein': 'protein_protein.csv',
|
| 'protein_disease': 'protein_disease.csv',
|
| 'disease_disease': 'disease_disease.csv',
|
| 'drug_disease': 'drug_disease_indication.csv',
|
| }
|
|
|
| return _build_heterograph(node_path, edge_path, node_files, edge_files, feature_dim=128)
|
|
|
|
|
| def load_KGdataset():
|
| """
|
| Load the heterogeneous network for the full knowledge graph dataset.
|
| """
|
| node_path = "/vast/yg3191/AIVS/kg/node"
|
| edge_path = "/vast/yg3191/AIVS/kg/edge"
|
|
|
| node_files = {
|
| 'drug': 'drug.csv',
|
| 'disease': 'disease.csv',
|
| 'protein': 'protein.csv',
|
| 'bioprocess': 'bioprocess.csv',
|
| 'cellcomp': 'cellcomp.csv',
|
| 'molfunc': 'molfunc.csv',
|
| 'pathway': 'pathway.csv',
|
| 'phenotype': 'phenotype.csv',
|
| 'exposure': 'exposure.csv',
|
| 'effect': 'effect.csv',
|
| }
|
|
|
| edge_files = {
|
| 'drug_drug': 'drug_drug.csv',
|
| 'drug_effect': 'drug_effect.csv',
|
| 'drug_protein': 'drug_protein.csv',
|
| 'drug_disease': 'drug_disease_indication.csv',
|
| 'protein_protein': 'protein_protein.csv',
|
| 'protein_bioprocess': 'protein_bioprocess.csv',
|
| 'protein_cellcomp': 'protein_cellcomp.csv',
|
| 'protein_molfunc': 'protein_molfunc.csv',
|
| 'protein_pathway': 'protein_pathway.csv',
|
| 'protein_disease': 'protein_disease.csv',
|
| 'disease_disease': 'disease_disease.csv',
|
| 'disease_phenotype': 'disease_phenotype_positive.csv',
|
| 'disease_exposure': 'disease_exposure.csv',
|
| 'bioprocess_bioprocess': 'bioprocess_bioprocess.csv',
|
| 'cellcomp_cellcomp': 'cellcomp_cellcomp.csv',
|
| 'molfunc_molfunc': 'molfunc_molfunc.csv',
|
| 'pathway_pathway': 'pathway_pathway.csv',
|
| 'phenotype_phenotype': 'phenotype_phenotype.csv',
|
| }
|
|
|
| return _build_heterograph(node_path, edge_path, node_files, edge_files, feature_dim=128)
|
|
|
|
|
| def load_Kdataset():
|
| """
|
| Load the heterogeneous network for the 'Kdataset'.
|
|
|
| Returns:
|
| dgl.DGLHeteroGraph: The constructed heterogeneous graph for Kdataset.
|
| """
|
|
|
| drug_drug = pd.read_csv("./dataset/Kdataset/drug_drug_baseline.csv", header=None).values
|
| drug_sim = drug_drug.copy()
|
| for i in range(len(drug_drug)):
|
| sorted_idx = np.argpartition(drug_drug[i], 15)
|
| drug_drug[i, sorted_idx[-15:]] = 1
|
| drug_drug_df = pd.DataFrame(np.array(np.where(drug_drug == 1)).T, columns=["Drug1", "Drug2"])
|
|
|
|
|
| protein_protein = pd.read_csv("./dataset/Kdataset/interactions/protein_protein.csv")
|
| gene_gene = pd.read_csv("./dataset/Kdataset/interactions/gene_gene.csv")
|
| pathway_pathway = pd.read_csv("./dataset/Kdataset/interactions/pathway_pathway.csv")
|
| disease_disease = pd.read_csv("./dataset/Kdataset/disease_disease_baseline.csv", header=None).values
|
| disease_sim = disease_disease.copy()
|
| for i in range(len(disease_disease)):
|
| sorted_idx = np.argpartition(disease_disease[i], 15)
|
| disease_disease[i, sorted_idx[-15:]] = 1
|
| disease_disease_df = pd.DataFrame(np.array(np.where(disease_disease == 1)).T, columns=["Disease1", "Disease2"])
|
|
|
| drug_protein = pd.read_csv("./dataset/Kdataset/associations/drug_protein.csv")
|
| protein_gene = pd.read_csv("./dataset/Kdataset/associations/protein_gene.csv")
|
| gene_pathway = pd.read_csv("./dataset/Kdataset/associations/gene_pathway.csv")
|
| pathway_disease = pd.read_csv("./dataset/Kdataset/associations/pathway_disease.csv")
|
| drug_disease = pd.read_csv("./dataset/Kdataset/associations/Kdataset.csv")
|
|
|
|
|
| graph_data = {
|
| ("drug", "drug_drug", "drug"): (
|
| th.tensor(drug_drug_df["Drug1"].values),
|
| th.tensor(drug_drug_df["Drug2"].values),
|
| ),
|
| ("drug", "drug_protein", "protein"): (
|
| th.tensor(drug_protein["Drug"].values),
|
| th.tensor(drug_protein["Protein"].values),
|
| ),
|
| ("protein", "protein_drug", "drug"): (
|
| th.tensor(drug_protein["Protein"].values),
|
| th.tensor(drug_protein["Drug"].values),
|
| ),
|
| ("protein", "protein_protein", "protein"): (
|
| th.tensor(protein_protein["Protein1"].values),
|
| th.tensor(protein_protein["Protein2"].values),
|
| ),
|
| ("protein", "protein_gene", "gene"): (
|
| th.tensor(protein_gene["Protein"].values),
|
| th.tensor(protein_gene["Gene"].values),
|
| ),
|
| ("gene", "gene_protein", "protein"): (
|
| th.tensor(protein_gene["Gene"].values),
|
| th.tensor(protein_gene["Protein"].values),
|
| ),
|
| ("gene", "gene_gene", "gene"): (
|
| th.tensor(gene_gene["Gene1"].values),
|
| th.tensor(gene_gene["Gene2"].values),
|
| ),
|
| ("gene", "gene_pathway", "pathway"): (
|
| th.tensor(gene_pathway["Gene"].values),
|
| th.tensor(gene_pathway["Pathway"].values),
|
| ),
|
| ("pathway", "pathway_gene", "gene"): (
|
| th.tensor(gene_pathway["Pathway"].values),
|
| th.tensor(gene_pathway["Gene"].values),
|
| ),
|
| ("pathway", "pathway_pathway", "pathway"): (
|
| th.tensor(pathway_pathway["Pathway1"].values),
|
| th.tensor(pathway_pathway["Pathway2"].values),
|
| ),
|
| ("pathway", "pathway_disease", "disease"): (
|
| th.tensor(pathway_disease["Pathway"].values),
|
| th.tensor(pathway_disease["Disease"].values),
|
| ),
|
| ("disease", "disease_pathway", "pathway"): (
|
| th.tensor(pathway_disease["Disease"].values),
|
| th.tensor(pathway_disease["Pathway"].values),
|
| ),
|
| ("disease", "disease_disease", "disease"): (
|
| th.tensor(disease_disease_df["Disease1"].values),
|
| th.tensor(disease_disease_df["Disease2"].values),
|
| ),
|
| ("drug", "drug_disease", "disease"): (
|
| th.tensor(drug_disease["Drug"].values),
|
| th.tensor(drug_disease["Disease"].values),
|
| ),
|
| ("disease", "disease_drug", "drug"): (
|
| th.tensor(drug_disease["Disease"].values),
|
| th.tensor(drug_disease["Drug"].values),
|
| ),
|
| }
|
| g = dgl.heterograph(graph_data)
|
|
|
|
|
| drug_feature = np.hstack((drug_sim, np.zeros((g.num_nodes("drug"), g.num_nodes("disease")))))
|
| dis_feature = np.hstack((np.zeros((g.num_nodes("disease"), g.num_nodes("drug"))), disease_sim))
|
| g.nodes["drug"].data["h"] = th.from_numpy(drug_feature).to(th.float32)
|
| g.nodes["disease"].data["h"] = th.from_numpy(dis_feature).to(th.float32)
|
| g.nodes["protein"].data["h"] = th.zeros((g.num_nodes("protein"), drug_feature.shape[1])).to(th.float32)
|
| g.nodes["gene"].data["h"] = th.zeros((g.num_nodes("gene"), drug_feature.shape[1])).to(th.float32)
|
| g.nodes["pathway"].data["h"] = th.zeros((g.num_nodes("pathway"), drug_feature.shape[1])).to(th.float32)
|
| return g
|
|
|
|
|
| def load_Bdataset():
|
| """
|
| Load the heterogeneous network for the 'Bdataset'.
|
|
|
| Returns:
|
| dgl.DGLHeteroGraph: The constructed heterogeneous graph for Bdataset.
|
| """
|
|
|
| drug_drug = pd.read_csv("./dataset/Bdataset/drug_drug_baseline.csv", header=None).values
|
| drug_sim = drug_drug.copy()
|
| for i in range(len(drug_drug)):
|
| sorted_idx = np.argpartition(drug_drug[i], 15)
|
| drug_drug[i, sorted_idx[-15:]] = 1
|
| drug_drug_df = pd.DataFrame(np.array(np.where(drug_drug == 1)).T, columns=["Drug1", "Drug2"])
|
|
|
| protein_protein = pd.read_csv("./dataset/Bdataset/interactions/protein_protein.csv")
|
| disease_disease = pd.read_csv("./dataset/Bdataset/disease_disease_baseline.csv", header=None).values
|
| disease_sim = disease_disease.copy()
|
| for i in range(len(disease_disease)):
|
| sorted_idx = np.argpartition(disease_disease[i], 15)
|
| disease_disease[i, sorted_idx[-15:]] = 1
|
| disease_disease_df = pd.DataFrame(np.array(np.where(disease_disease == 1)).T, columns=["Disease1", "Disease2"])
|
| drug_protein = pd.read_csv("./dataset/Bdataset/associations/drug_protein.csv")
|
| drug_disease = pd.read_csv("./dataset/Bdataset/associations/Bdataset.csv")
|
|
|
|
|
| graph_data = {
|
| ("drug", "drug_drug", "drug"): (
|
| th.tensor(drug_drug_df["Drug1"].values),
|
| th.tensor(drug_drug_df["Drug2"].values),
|
| ),
|
| ("drug", "drug_protein", "protein"): (
|
| th.tensor(drug_protein["Drug"].values),
|
| th.tensor(drug_protein["Protein"].values),
|
| ),
|
| ("protein", "protein_drug", "drug"): (
|
| th.tensor(drug_protein["Protein"].values),
|
| th.tensor(drug_protein["Drug"].values),
|
| ),
|
| ("protein", "protein_protein", "protein"): (
|
| th.tensor(protein_protein["Protein1"].values),
|
| th.tensor(protein_protein["Protein2"].values),
|
| ),
|
| ("disease", "disease_disease", "disease"): (
|
| th.tensor(disease_disease_df["Disease1"].values),
|
| th.tensor(disease_disease_df["Disease2"].values),
|
| ),
|
| ("drug", "drug_disease", "disease"): (
|
| th.tensor(drug_disease["Drug"].values),
|
| th.tensor(drug_disease["Disease"].values),
|
| ),
|
| ("disease", "disease_drug", "drug"): (
|
| th.tensor(drug_disease["Disease"].values),
|
| th.tensor(drug_disease["Drug"].values),
|
| ),
|
| }
|
| g = dgl.heterograph(graph_data)
|
|
|
|
|
| drug_feature = np.hstack((drug_sim, np.zeros((g.num_nodes("drug"), g.num_nodes("disease")))))
|
| dis_feature = np.hstack((np.zeros((g.num_nodes("disease"), g.num_nodes("drug"))), disease_sim))
|
| g.nodes["drug"].data["h"] = th.from_numpy(drug_feature).to(th.float32)
|
| g.nodes["disease"].data["h"] = th.from_numpy(dis_feature).to(th.float32)
|
| g.nodes["protein"].data["h"] = th.zeros((g.num_nodes("protein"), g.num_nodes("protein"))).to(th.float32)
|
| return g
|
|
|
|
|
| def remove_graph(g, test_id):
|
| """
|
| Remove drug-disease association edges that belong to the test set from the graph.
|
|
|
| Parameters:
|
| g (dgl.DGLHeteroGraph): The heterogeneous graph.
|
| test_id (numpy.ndarray): Array of shape (n, 2) where each row is [drug_index, disease_index].
|
|
|
| Returns:
|
| dgl.DGLHeteroGraph: The graph with test edges removed.
|
| """
|
| test_drug_id = test_id[:, 0]
|
| test_dis_id = test_id[:, 1]
|
|
|
| edges_id = g.edge_ids(
|
| th.tensor(test_drug_id),
|
| th.tensor(test_dis_id),
|
| etype=("drug", "drug_disease", "disease"),
|
| )
|
| g = dgl.remove_edges(g, edges_id, etype=("drug", "drug_disease", "disease"))
|
|
|
| edges_id = g.edge_ids(
|
| th.tensor(test_dis_id),
|
| th.tensor(test_drug_id),
|
| etype=("disease", "disease_drug", "drug"),
|
| )
|
| g = dgl.remove_edges(g, edges_id, etype=("disease", "disease_drug", "drug"))
|
| return g |