Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions .github/workflows/tests.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
name: Component tests

on:
pull_request:
push:
branches: [main]

permissions:
contents: read

jobs:
tests:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
with:
persist-credentials: false
- uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v7.0.0
with:
python-version: '3.11'
- run: python -m pip install pytest
- run: python -m pytest -q tests
14 changes: 14 additions & 0 deletions GraphLib/device.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
"""Tensor-attribute transfer for the legacy DGL graph interface."""


def send_graph_to_device(g, device):
"""Move node/edge attributes in place without removing them before transfer.

This preserves the legacy helper's graph identity and topology behavior.
If a tensor transfer fails, its source attribute remains available.
"""
for name in g.node_attr_schemes():
g.ndata[name] = g.ndata[name].to(device, non_blocking=True)
for name in g.edge_attr_schemes():
g.edata[name] = g.edata[name].to(device, non_blocking=True)
return g
17 changes: 5 additions & 12 deletions GraphLib/dglutil.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,8 @@
if __package__:
from .device import send_graph_to_device
else: # Legacy scripts import dglutil from GraphLib directly.
from device import send_graph_to_device

import dgl
import dgl.function as fn
import torch as th
Expand Down Expand Up @@ -70,15 +75,3 @@ def merge_graphs(graphs, feature_arrays, label_arrays, training_masks):
labels = np.concatenate(label_arrays)
training_masks = np.concatenate(training_masks)
return g, features, labels, training_masks

def send_graph_to_device(g, device):
# nodes
labels = g.node_attr_schemes()
for l in labels.keys():
g.ndata[l] = g.ndata.pop(l).to(DEVICE, non_blocking=True)

# edges
labels = g.edge_attr_schemes()
for l in labels.keys():
g.edata[l] = g.edata.pop(l).to(DEVICE, non_blocking=True)
return g
15 changes: 14 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,19 @@ Start with the [model code](GraphLib/model.py) to inspect the architecture or th
- [FastAFGCN](https://github.com/lmlearning/FastAFGCN): quantized ONNX inference.
- [AFSubsample](https://github.com/lmlearning/AFSubsample): framework subsampling and analysis.

## License
## Component tests

```bash
python -m pip install pytest
python -m pytest tests
```

The device-transfer helper is isolated in `GraphLib/device.py` and remains available
through `GraphLib.dglutil`. It uses the requested device and keeps the original
attribute if its transfer fails. The graph is modified in place; earlier successful
transfers are not rolled back. These tests use graph/tensor protocol doubles and do
not require DGL or validate GPU execution or model training.

## License

See [LICENSE](LICENSE).
53 changes: 53 additions & 0 deletions tests/test_device.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
import pytest

from GraphLib.device import send_graph_to_device


class Tensor:
def __init__(self, fail=False):
self.fail = fail
self.calls = []

def to(self, device, **kwargs):
self.calls.append((device, kwargs))
if self.fail:
raise RuntimeError("transfer failed")
return ("converted", device)


class Graph:
def __init__(self, nodes=None, edges=None):
self.ndata = nodes or {}
self.edata = edges or {}

def node_attr_schemes(self):
return dict.fromkeys(self.ndata)

def edge_attr_schemes(self):
return dict.fromkeys(self.edata)


@pytest.mark.parametrize("device", ["cpu", "cuda:1"])
def test_transfers_both_attribute_kinds_to_requested_device(device):
node, edge = Tensor(), Tensor()
graph = Graph({"features": node}, {"weights": edge})
assert send_graph_to_device(graph, device) is graph
assert graph.ndata == {"features": ("converted", device)}
assert graph.edata == {"weights": ("converted", device)}
assert node.calls == edge.calls == [(device, {"non_blocking": True})]


@pytest.mark.parametrize("attribute_kind", ["ndata", "edata"])
def test_failed_transfer_keeps_original_attribute(attribute_kind):
tensor = Tensor(fail=True)
graph = Graph()
getattr(graph, attribute_kind)["features"] = tensor
with pytest.raises(RuntimeError, match="transfer failed"):
send_graph_to_device(graph, "cpu")
assert getattr(graph, attribute_kind)["features"] is tensor


def test_empty_graph_is_unchanged():
graph = Graph()
assert send_graph_to_device(g=graph, device="cpu") is graph
assert graph.ndata == graph.edata == {}
Loading