提交

[Graphbolt] Add the preprocess_ondisk_dataset function. (#5991)

Co-authored-by: Hongzhi (Steve), Chen <chenhongzhi.nkcs@gmail.com>
这个提交包含在:
keli-wen
2023-07-18 13:05:23 +08:00
提交者 GitHub
父节点 2746aac3ae
当前提交 d90954b10c
修改 3 个文件,包含 338 行新增7 行删除
@@ -4,10 +4,12 @@ import tempfile
import gb_test_utils as gbt
import numpy as np
import pandas as pd
import pydantic
import pytest
import torch
import yaml
from dgl import graphbolt as gb
@@ -747,3 +749,122 @@ def test_OnDiskDataset_Metadata():
assert dataset.dataset_name == dataset_name
assert dataset.num_classes is None
assert dataset.num_labels is None
def test_OnDiskDataset_preprocess_homogeneous():
"""Test preprocess of OnDiskDataset."""
with tempfile.TemporaryDirectory() as test_dir:
# All metadata fields are specified.
dataset_name = "graphbolt_test"
num_nodes = 4000
num_edges = 20000
num_classes = 10
num_labels = 9
# Generate random edges.
nodes = np.repeat(np.arange(num_nodes), 5)
neighbors = np.random.randint(0, num_nodes, size=(num_edges))
edges = np.stack([nodes, neighbors], axis=1)
# Wrtie into edges/edge.csv
os.makedirs(os.path.join(test_dir, "edges/"), exist_ok=True)
edges = pd.DataFrame(edges, columns=["src", "dst"])
edges.to_csv(
os.path.join(test_dir, "edges/edge.csv"),
index=False,
header=False,
)
# Generate random graph edge-feats.
edge_feats = np.random.rand(num_edges, 5)
os.makedirs(os.path.join(test_dir, "data/"), exist_ok=True)
np.save(os.path.join(test_dir, "data/edge-feat.npy"), edge_feats)
# Generate random node-feats.
node_feats = np.random.rand(num_nodes, 10)
np.save(os.path.join(test_dir, "data/node-feat.npy"), node_feats)
# Generate train/test/valid set.
os.makedirs(os.path.join(test_dir, "set/"), exist_ok=True)
train_pairs = (np.arange(1000), np.arange(1000, 2000))
train_labels = np.random.randint(0, 10, size=1000)
train_data = np.vstack([train_pairs, train_labels]).T
train_path = os.path.join(test_dir, "set/train.npy")
np.save(train_path, train_data)
validation_pairs = (np.arange(1000, 2000), np.arange(2000, 3000))
validation_labels = np.random.randint(0, 10, size=1000)
validation_data = np.vstack([validation_pairs, validation_labels]).T
validation_path = os.path.join(test_dir, "set/validation.npy")
np.save(validation_path, validation_data)
test_pairs = (np.arange(2000, 3000), np.arange(3000, 4000))
test_labels = np.random.randint(0, 10, size=1000)
test_data = np.vstack([test_pairs, test_labels]).T
test_path = os.path.join(test_dir, "set/test.npy")
np.save(test_path, test_data)
yaml_content = f"""
dataset_name: {dataset_name}
num_classes: {num_classes}
num_labels: {num_labels}
graph: # graph structure and required attributes.
nodes:
- num: {num_nodes}
edges:
- format: csv
path: edges/edge.csv
feature_data:
- domain: edge
type: null
name: feat
format: numpy
in_memory: true
path: data/edge-feat.npy
feature_data:
- domain: node
type: null
name: feat
format: numpy
in_memory: false
path: data/node-feat.npy
train_sets:
- - type_name: null
# shape: (num_trains, 3), 3 for (src, dst, label).
format: numpy
path: set/train.npy
validation_sets:
- - type_name: null
format: numpy
path: set/validation.npy
test_sets:
- - type_name: null
format: numpy
path: set/test.npy
"""
yaml_file = os.path.join(test_dir, "test.yaml")
with open(yaml_file, "w") as f:
f.write(yaml_content)
output_file = gb.ondisk_dataset.preprocess_ondisk_dataset(yaml_file)
with open(output_file, "rb") as f:
processed_dataset = yaml.safe_load(f)
assert processed_dataset["dataset_name"] == dataset_name
assert processed_dataset["num_classes"] == num_classes
assert processed_dataset["num_labels"] == num_labels
assert "graph" not in processed_dataset
assert "graph_topology" in processed_dataset
csc_sampling_graph = gb.csc_sampling_graph.load_csc_sampling_graph(
os.path.join(test_dir, processed_dataset["graph_topology"]["path"])
)
assert csc_sampling_graph.num_nodes == num_nodes
assert csc_sampling_graph.num_edges == num_edges
num_samples = 100
fanout = 1
subgraph = csc_sampling_graph.sample_neighbors(
torch.arange(num_samples),
torch.tensor([fanout]),
)
assert len(list(subgraph.node_pairs.values())[0][0]) <= num_samples