| --- |
| library_name: kernels |
| license: apache-2.0 |
| --- |
| |
| # graph-spmm |
|
|
| Sparse-dense matrix products for message passing, loadable through |
| `kernels`: a CSR adjacency applied to a dense feature matrix with sum, mean, |
| or max aggregation. The reference baseline is `torch.sparse.mm` (cuSPARSE), |
| matched to fp32 summation order and beaten up to 2.1x on the sparse, |
| narrow-feature shapes graph networks actually run. Companion to |
| [graph-reduce](https://huggingface.co/kernels/phanerozoic/graph-reduce) for |
| unsorted destinations. |
|
|
| A graph convolution is a sparse adjacency times a dense feature matrix, |
| applied once per layer over a graph that does not change. General sparse |
| libraries pay for generality, and atomic-based scatter alternatives return |
| different bits on every run. This kernel keeps one warp per output row with |
| no atomics anywhere, so message passing over million-node graphs is both |
| faster than cuSPARSE in its home regime and bitwise reproducible, and |
| max-aggregation returns the winning neighbour's index in the same sweep, the |
| piece a max-pool backward needs. |
|
|
| <video src="https://huggingface.co/kernels/phanerozoic/graph-spmm/resolve/main/media/hero.mp4" autoplay loop muted playsinline controls width="100%"> |
| <img src="https://huggingface.co/kernels/phanerozoic/graph-spmm/resolve/main/media/hero.gif" alt="Heat from four sources spreads along the arms of a spiral graph, one spmm application per step"> |
| </video> |
|
|
| *Heat from four sources diffusing through an 8,000-node kNN graph, one |
| `spmm` mean-aggregation per step (log-scale coloring). On 1M nodes and 10M |
| edges at 16 channels the product measures 0.39 ms against cuSPARSE's 0.85 |
| ms, bitwise reproducible with no atomics.* |
|
|
| ## Usage |
|
|
| ```python |
| import torch |
| from kernels import get_kernel |
| |
| tsp = get_kernel("phanerozoic/graph-spmm", version=1, trust_remote_code=True) |
| |
| # build the CSR once for a fixed graph, then apply it per layer |
| adj = tsp.SparseAdj.from_edge_index(edge_index, num_nodes, edge_weight) |
| h = adj @ x # sum aggregation |
| h = adj.reduce(x, "mean") # mean |
| h, arg = adj.reduce(x, "max") # max, plus the winning neighbour |
| ``` |
|
|
| `version` selects the release branch; `trust_remote_code` is required by |
| `kernels` for publishers without the trusted-publisher mark. |
|
|
| ## API |
|
|
| | Symbol | Purpose | |
| |---|---| |
| | `spmm_sum(rowptr, col, val, mat)` | `A @ mat`; `val=None` is an unweighted adjacency | |
| | `spmm_mean(rowptr, col, val, mat)` | row-wise mean; empty rows are 0 | |
| | `spmm_max(rowptr, col, val, mat)` | `(values, argmax)`; empty rows are `(0, -1)` | |
| | `spmm(rowptr, col, val, mat, reduce)` | dispatch by name | |
| | `SparseAdj(rowptr, col, val)` | a CSR adjacency held once and applied many times | |
| | `SparseAdj.from_edge_index(edge_index, num_nodes, val)` | sort a `[2, E]` edge list into CSR | |
|
|
| `rowptr` is `[S+1]` int64, `col` is `[nnz]` int64, `mat` is `[K, C]` fp32. |
|
|
| ## Method |
|
|
| One warp per output row. The row's nonzeros are walked cooperatively with |
| each lane holding a strided slice of the feature columns, so dense loads |
| coalesce across the warp and every output element reduces privately in a |
| register. Nothing is shared between warps, so there are no atomics: the |
| result is bitwise reproducible and independent of scheduling. `spmm_max` |
| resolves the argmax during the same sweep. `val=None` treats every stored |
| entry as 1 and skips the value load. |
|
|
| ## Measured |
|
|
| fp32, median of 30, against `torch.sparse.mm` on a CSR tensor (cuSPARSE): |
|
|
| | nodes | nnz | C | avg degree | `torch.sparse.mm` | this | speedup | |
| |---|---|---|---|---|---|---| |
| | 10,000 | 200,000 | 32 | 20 | 0.028 ms | 0.023 ms | 1.23x | |
| | 10,000 | 200,000 | 128 | 20 | 0.035 ms | 0.032 ms | 1.10x | |
| | 100,000 | 1,000,000 | 32 | 10 | 0.081 ms | 0.038 ms | 2.11x | |
| | 100,000 | 4,000,000 | 64 | 40 | 0.301 ms | 0.245 ms | 1.23x | |
| | 1,000,000 | 10,000,000 | 16 | 10 | 0.894 ms | 0.494 ms | 1.81x | |
|
|
| The margin is largest on many nodes, low degree, modest channel count; at |
| high degree and wide features both are bandwidth-bound on the dense operand |
| and converge. |
|
|
| ## Correctness |
|
|
| - Sum agrees with `torch.sparse.mm` to 1.5e-5 absolute across 10^4 to 10^6 |
| nodes (fp32 summation order); mean to 3.0e-7 against the row-normalized |
| product. |
| - Max and its arg are checked against a per-row reference built from the CSR |
| entries (a dense scatter silently drops duplicate pairs, which a real CSR |
| can contain and this kernel sums). |
| - Unweighted adjacency matches an explicit all-ones value array; empty rows |
| reduce to 0 for sum and mean and `(0, -1)` for max. |
| - Deterministic: repeated products are bitwise identical. |
| - `SparseAdj.from_edge_index` sorts an unsorted edge list into CSR and |
| reproduces the same product. |
|
|
| ## Requirements and limits |
|
|
| - NVIDIA GPU with compute capability 8.0+; fp32 only. |
| - One warp per row: power-law graphs with extreme hubs load-imbalance and |
| want a segmented schedule this kernel does not implement. |
| - Forward only; the backward of a sum aggregation is the transposed product, |
| composed from a second `SparseAdj` on the reversed edge list. |
|
|
| ## References |
|
|
| Merrill and Garland, "Merge-based parallel sparse matrix-vector |
| multiplication" (2016); Bell and Garland, "Implementing sparse matrix-vector |
| multiplication on throughput-oriented processors" (2009). |
|
|
| ## License |
|
|
| Apache-2.0. |
|
|