AIVS / REDDA /load_data.py
yg3191's picture
Upload folder using huggingface_hub
132149b verified
Raw
History Blame Contribute Delete
16.4 kB
"""
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):
# Load node indices (vectorized, no Python loops over rows)
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)
# Infer src/dst types by splitting on first underscore
# Works for e.g. disease_phenotype, protein_bioprocess, bioprocess_bioprocess, etc.
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}")
# Vectorized mapping: get index positions or -1 when missing
src_idx_arr = node_indices[src_type].get_indexer(df['x_id'])
dst_idx_arr = node_indices[dst_type].get_indexer(df['y_id'])
# Keep only rows where both src and dst are valid (not -1)
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)
# Forward edges
graph_data[(src_type, edge_type, dst_type)] = (s, d)
# Reverse edges when not self-loop
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}")
# Build graph
g = dgl.heterograph(graph_data)
print(f"Created graph with {g.num_nodes()} total nodes and {g.num_edges()} total edges")
# Node/edge stats
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")
# Initialize node features (random as placeholder)
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.
"""
# Load and process drug-drug similarity data
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"])
# Load additional interaction data
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")
# Build the graph using the interaction data
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)
# Prepare node features by concatenating similarity matrices and zero padding as needed
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.
"""
# Load and process drug-drug similarity data
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")
# Build the graph using the interaction data
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)
# Prepare node features with appropriate zero padding
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]
# Remove edges for the ('drug', 'drug_disease', 'disease') relation
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"))
# Remove the reciprocal edges for the ('disease', 'disease_drug', 'drug') relation
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