graph-reduce
Segmented reductions for geometric deep learning, loadable through
kernels. Message passing reduces a per-edge tensor into a per-node one,
and the two access patterns that cover it want different kernels: scatter
(arbitrary output index per row, atomic writes) and segment (rows grouped by
a CSR pointer, each output owned by one block, no atomics). Both ship here.
The reference baselines are torch.index_add_ and torch.segment_reduce,
the latter matched bitwise. Companion to
graph-spmm.
Edge lists are commonly stored sorted by destination, which makes the
atomic-free path available, and that path buys two things at once: it is
faster than both the atomic kernels and torch.index_add_ in the regime
graph networks run, and its summation order is fixed, so a reduction that
has to be reproducible stops being a compromise. scatter_max and
scatter_min also return the arg, which torch's native scatter_reduce
does not.
One million rows into 10,000 outputs, drain rates set by the measured
times: the atomic path at 0.047 ms with a run-to-run spread of 8e-6, the
segment path at 0.016 ms, bitwise identical on every run, 6.1x
torch.index_add_.
Usage
import torch
from kernels import get_kernel
ts = get_kernel("phanerozoic/graph-reduce", version=1, trust_remote_code=True)
# general case: arbitrary destination per edge
out = ts.scatter_sum(messages, dst_index, num_nodes)
mx, argmx = ts.scatter_max(messages, dst_index, num_nodes)
# sorted edge list: build the CSR pointer once, then reduce without atomics
counts = torch.bincount(dst_index, minlength=num_nodes)
indptr = torch.cat([counts.new_zeros(1), counts.cumsum(0)])
out = ts.segment_sum_csr(messages_sorted_by_dst, indptr)
version selects the release branch; trust_remote_code is required by
kernels for publishers without the trusted-publisher mark.
API
| Symbol | Purpose |
|---|---|
scatter_sum(src, index, dim_size) / scatter_add |
sum rows into the output named by index |
scatter_mean(src, index, dim_size) |
mean; empty outputs are 0 |
scatter_max(src, index, dim_size) |
(values, argmax) |
scatter_min(src, index, dim_size) |
(values, argmin) |
scatter(src, index, dim_size, reduce) |
dispatch by name |
segment_sum_csr(src, indptr) |
atomic-free sum over contiguous groups |
segment_mean_csr(src, indptr) |
atomic-free mean |
segment_max_csr(src, indptr) / segment_min_csr |
(values, arg) |
src is [N, ...]; index is [N] int64; indptr is [S+1] int64.
Method
The scatter path is one atomic per element. For max and min the value and
its index resolve inside a single atomic: a float's bit pattern made
monotone as an int32, packed as (mapped_value << 32) | ~index, so the
largest 64-bit word carries both the maximum and the row that produced it,
with ties to the lower row. The segment path assigns one block per output
row; a segment's rows are contiguous and no other block touches that
output, so the reduction runs in registers and the store is a single write.
Measured
fp32 sum of [M, C] into S outputs, median of 50, against
torch.index_add_ (current torch and driver on the same host; atomic
throughput has improved since earlier published rows, and the table reflects
the current stack):
| M | C | S | index_add_ |
scatter | segment | segment vs index_add_ |
|---|---|---|---|---|---|---|
| 1,000,000 | 16 | 10,000 | 0.087 ms | 0.047 ms | 0.016 ms | 5.4x |
| 1,000,000 | 64 | 10,000 | 0.415 ms | 0.338 ms | 0.358 ms | 1.16x |
| 4,000,000 | 16 | 50,000 | 0.482 ms | 0.357 ms | 0.366 ms | 1.32x |
| 1,000,000 | 16 | 100 | 0.250 ms | 0.273 ms | 0.182 ms | 1.37x |
The segment margin is largest on narrow features and moderate output counts; at wide channels all paths are bandwidth-bound and converge.
Correctness
- Sum and mean agree with
index_add_to 1.7e-5 absolute overMto 10^6 (fp32 atomic ordering); max and min aretorch.equaltoindex_reduce'samax/amin. - The arg is checked, not assumed: every returned index is gathered back out of the source, must equal the reported extremum, and must belong to that row's segment; ties resolve to the lower row.
segment_csris bitwise equal totorch.segment_reduce(fixed summation order); empty outputs are 0 for sum/mean and-inf/+infwith arg-1for max/min.- Reproducibility, measured: over repeated runs the CSR path is bitwise identical every time; the atomic path is not, because arrival order is not fixed. Use the CSR path when a reduction has to be reproducible.
Requirements and limits
- NVIDIA GPU with compute capability 8.0+; fp32 only (the packed argmax atomic relies on the fp32 bit layout).
- The reduction is over dimension 0; permute first for other axes.
- The segment path requires input already sorted by output; sorting is worth it when amortized over many reductions on a fixed graph.
- Forward only; backwards compose from gathers and scatters.
References
Merrill and Garland, "Merge-based parallel sparse matrix-vector multiplication" (2016); the packed value-index atomic shared with phanerozoic/dpx-decode.
License
Apache-2.0.
- Downloads last month
- -
- OS
- linux
- Arch
- x86_64
- Kernel Builder
- 19aaa64





