Skip to content

Build CSR/CSC with a memory-lean counting sort instead of coo_to_csr - #768

Open
dsaini2-sc wants to merge 4 commits into
mainfrom
dsaini2/compact-csr-topology
Open

dsaini2-sc wants to merge 4 commits into
mainfrom
dsaini2/compact-csr-topology

Conversation

@dsaini2-sc

@dsaini2-sc dsaini2-sc commented Sep 16, 2026 •

Copy link
Copy Markdown
Collaborator

Scope of work done

GiGL's graph build converts each edge type's COO to CSR/CSC with graphlearn_torch.utils.coo_to_csr, which goes through torch_sparse.SparseStorage and holds seven full-size int64 arrays at its peak (measured 7.25x one int64 array at 400M edges). At billion-edge scale that conversion, not the graph itself, is what runs the host out of memory.

  • gigl/utils/csr.py (new): build_csr_from_coo, a two-pass counting sort. Degrees give the output layout up front, the output is allocated once and filled in place, and the input is read in chunks. Peak drops to about 3x one int64 array today, and 2x once the edge index is int32. CompactTopology(Topology) wraps the result. It skips Topology.__init__, which would allocate an arange(num_edges) edge-id array and call coo_to_csr.
  • dist_dataset.py: _build_graph_per_edge_type builds one CompactTopology + Graph per edge type, largest first, freeing each COO before starting the next.

Worth a reviewer's attention:

  • Graphs with edge ids, weights, or features keep GLT's init_graph (_has_per_edge_metadata). Ids and weights would have to be reordered along with the columns. Features are looked up by edge id, which CompactTopology doesn't store, and GLT's compiled sampler would crash (not raise) reading an empty id array.
  • indptr is sized from the node partition book, not max(row) + 1. Seeds are global ids, and the highest-id node a rank owns may have no edge in the compressed direction, which would send the sampler out of bounds.
  • CSR indices stay int64. int32 needs the GLT patch from Apply and verify local patches on the pinned GraphLearn-Torch build #761, but the images CI and jobs run on are still pinned to a pre-Apply and verify local patches on the pinned GraphLearn-Torch build #761 build (7d3182ee). int32 comes in a follow-up once the images are rebuilt.
  • The within-row sort avoids upstream's row * num_cols + col key, which can overflow int64 silently for blocks that span many rows. It sorts by column, then stably by row.
  • Freeing each COO also holds under an edge splitter. build() used to keep a reference to every COO while building the graph on that path.

Where is the documentation for this feature?: docstrings in gigl/utils/csr.py and dist_dataset.py, plus the CHANGELOG entry

Did you add automated tests or write a test plan?

Yes. tests/unit/utils/csr_test.py (29 tests) checks exact parity with coo_to_csr: dense and mostly-empty rows, a row larger than a sort block, blocks cut by the row cap, duplicate edges, large id ranges, CSC, rows split across chunks, empty input, error paths, and the banded (on-disk) scatter against the direct one. tests/unit/distributed/compact_topology_test.py (15 tests) covers CompactTopology's attributes against a real Topology, the metadata gate, indptr sizing, the build path, a test through build() that each COO is freed before the next edge type, and full-fanout sampling parity with GLT's own build for CSR and CSC. distributed_neighborloader_test.py and distributed_dataset_test.py, which spawn real loaders, pass: 51 tests, one pre-existing skip.

Updated Changelog.md? YES

Ready for code review?: YES

graphlearn_torch.utils.coo_to_csr delegates to torch_sparse.SparseStorage,
which holds seven full-size int64 arrays at its peak (measured 7.25x one
int64 array at 400M edges). At billion-edge partitions the conversion, not
the resident graph, is what exhausts the host, and it sits upstream of
every other memory lever.

gigl/utils/csr.py adds build_csr_from_coo, a two-pass counting sort:
degrees give the output layout up front, the output is allocated once and
written in place, and the input is consumed in chunks, so peak is
row + col + indices + 2 x (num_rows + 1) + O(chunk + max_degree).
indices narrows to int32 when the installed wheel accepts it (the CSR
patch from #761) and every column id is verified to fit; indptr stays
int64. indptr is sized from the node partition book rather than
max(row)+1, which can come up short and send the compiled sampler out of
bounds.

dist_dataset._initialize_graph takes that path per edge type, largest
first, when there are no edge ids and no edge weights, and falls back to
GLT's init_graph otherwise. _build_topology_without_edge_ids populates a
Topology directly so the arange edge-id array GLT fabricates for absent
edge ids is never allocated.
Comment thread gigl/distributed/dist_dataset.py
Comment thread gigl/utils/csr.py Outdated

@mkolodner-sc mkolodner-sc left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks Deepak! Did an initial pass here and left some comments

Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/distributed/dist_dataset.py Outdated
Comment thread gigl/distributed/dist_dataset.py Outdated
Comment thread gigl/distributed/dist_dataset.py Outdated
Comment thread gigl/distributed/dist_dataset.py Outdated
Comment thread gigl/distributed/dist_dataset.py
Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/distributed/dist_dataset.py Outdated
Comment thread gigl/distributed/dist_dataset.py Outdated
Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/utils/csr.py Outdated
Comment thread gigl/utils/csr.py Outdated
@dsaini2-sc

Copy link
Copy Markdown
Collaborator Author

Thanks Deepak! Did an initial pass here and left some comments

Thanks @mkolodner-sc and @kmontemayor2-sc and @zfan3-sc -- I have worked on your feedback to make the PR leaner by dropping the workarounds -- we now have our own CompactTopology instead of setting GLT's private fields; and no int32 probe. This makes the PR leaner by ~350lines.
One thing to flag: instead of assuming int32, I've kept the CSR column ids int64. The images CI and jobs run on are still pinned to a build from before #761, so their GLT doesn't have the patch yet. I will do a followup PR with int32 once those images are rebuilt.
Specifically, thanks @zfan3-sc for the dangling reference catch, that was a real leak and I have fixed it in the new commits.
Other replies inline!

@zfan3-sc zfan3-sc left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

lgtm on my feedbacks; thanks for the work

@kmontemayor2-sc kmontemayor2-sc left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks Deepak! LGTM provided we can address these comments :)

Comment on lines +788 to 797
logger.info(
"Edges carry ids, weights, or features; using GLT's init_graph for the topology"
)
self.init_graph(
edge_index=edge_index,
edge_ids=edge_ids,
graph_mode="CPU",
directed=True,
edge_weights=edge_weights,
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

nit. can we put this in else?

so the dict values are what matter.
"""
if isinstance(edge_ids, Mapping):
has_edge_ids = any(ids is not None for ids in edge_ids.values())

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is not None check lets the hash partitioner's placeholder through. Trace:

  1. dist_partitioner.py:1314-1317: an edge type with no local edges sets had_zero_edges and
    partitioned_edge_index = torch.empty((2, 0)).
  2. dist_partitioner.py:1343-1344 (no features) and :1444 (features): partitioned_edge_ids = torch.empty(0), not None.
  3. dist_partitioner.py:1494-1498: that goes into GraphPartitionData.edge_ids.
  4. Here, any(ids is not None ...) is True for that one type, so every edge type takes
    self.init_graph at :791.
  5. GLT Topology.init -> coo_to_csr -> row.max() on the empty COO: RuntimeError: max():
    Expected reduction dim to be specified for input.numel() == 0. (Verified against the installed
    GLT, with and without edge_ids.)

So the comment at :868 ("Also run for an edge type with no edges on this rank, which both
partitioners produce") is right about the lean path but the gate never lets it get there. main
crashes the same way, so not a regression, but the optimization silently switches itself off on
exactly the sparse hetero graphs it's for.

Suggest ids is not None and ids.numel() > 0 here, or have the partitioner emit None at those two
lines so the docstring's "dict of None" contract actually holds. Either way, add the placeholder
case to PerEdgeMetadataTest.

Comment on lines +859 to +863
num_nodes = get_total_ids(
node_partition_book[node_type]
if isinstance(node_partition_book, Mapping)
else node_partition_book
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This changes what indptr costs per rank. GLT sized it max(row) + 1, which under range
partitioning is roughly this rank's range end. get_total_ids(RangePartitionBook) returns
partition_bounds[-1], the global count.

So for a RangePartitionBook rank 0 of 10 goes from ~N/10 to N. At 1e9 nodes that's 8 GiB per edge
type on every rank. With a handful of edge types on a sparse graph that eats a real share of the
saving.

The docstring's reason (seeds are global ids, highest local node may have no edges) is right for
hash partitioning. For a range book, everything this rank is asked about is < partition_bounds[rank + 1], so sizing to that is enough. GLT's C++ samplers also bounds-check (v < row_count returns 0 neighbors), so an out-of-range seed degrades rather than crashes either
way.

Suggest: size to the rank's range end for RangePartitionBook, global for tensor books. If you'd
rather keep it global, put the trade-off in the docstring so the next person doesn't "fix" it.

Is there a reason we can't use the local space here?

if isinstance(partitioned_edge_index, MutableMapping):
partitioned_edge_index.clear()
self.graph = self._build_graph_per_edge_type(
streaming_input, node_partition_book=node_partition_book

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ty fails on this branch (6 diagnostics, main passes), which fails make unit_test_py since it
depends on type_check. CI is currently skipped on the PR so it hasn't surfaced.

error[invalid-argument-type] gigl/distributed/dist_dataset.py:768 streaming_input: ty does not narrow edge_index via the assert
error[unresolved-attribute] gigl/distributed/dist_dataset.py:772 self.graph.keys() on ... | None
error[unresolved-attribute] tests/unit/distributed/compact_topology_test.py:154 dataset.graph.topo
error[unresolved-attribute] tests/unit/distributed/compact_topology_test.py:173 dataset.graph.keys()
error[not-subscriptable] tests/unit/distributed/compact_topology_test.py:174 dataset.graph[_EDGE_TYPE]
error[invalid-argument-type] tests/unit/distributed/compact_topology_test.py:282 layout: str -> Literal["CSR", "CSC"]

Binding a typed local (hetero_edge_index: dict[EdgeType, torch.Tensor] = edge_index) instead of
the bare assert isinstance clears the first; assert isinstance(self.graph, dict) before .keys()
clears the second; the tests need the same narrowing plus Literal on the parameterized layout.

Comment thread gigl/utils/csr.py
DEFAULT_BAND_BYTES = 4 * 2**30


def _compute_degrees_from_coo_rows(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

btw can we use torch.bincount here instead?

Comment thread gigl/utils/csr.py
del destination, col_chunk, unique_rows, counts


def _scatter_whole(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Robot review:

_scatter_whole is _scatter_in_bands with one band covering every row. The per-chunk in_band mask
is a bool array and cheap next to the argsort in _place_chunk.

Merging them (band_bytes=None -> single band) removes ~25 lines, the is_disk_backed dispatch at
:319, and the two wraps= spy tests at csr_test.py:384-435 that assert which private function ran.
Fewer paths to keep in sync, and the disk path gets exercised by every test instead of two.

Comment on lines +112 to +115
def _dataset(self) -> DistDataset:
dataset = DistDataset.__new__(DistDataset)
dataset.edge_dir = "in"
return dataset

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we just do DistDataset(rank=0, world_size=1, edge_dir=...) instead?

Comment thread gigl/utils/csr.py
col: torch.Tensor,
num_rows: int,
chunk_size: int = DEFAULT_CHUNK_SIZE,
sort_within_row: bool = True,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is this ever false? Do we forsee a reason for it to be false?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants