Skip to content

perf: NeighborGraph attention is 4-9x slower than dense and OOM-prone -- move to block-dense attention via the destination CSR #5861

Description

@wanghan-iapcm

Summary

The NeighborGraph route is dense-competitive for attention-free descriptors, but graph attention is 4-9x slower than the dense route and memory-bound to the point of OOM, on both DPA-1 and DPA-2. The cause is the ragged pair-axis attention design; the proposed fix is block-dense attention through the existing destination CSR.

Measurements (diamond at experimental density, a = 3.567 A, fp64)

Compiled .pt2 AOTI inference, V100-32GB (6-layer DPA-2, g1=128, g2=32, 4 heads, repinit nsel 120 / rcut 6, repformer nsel 40 / rcut 4; ms/eval, 30 CUDA-synced evals):

natoms graph native graph smooth dense native dense smooth graph/dense
64 54.4 68.9 13.7 14.5 4.7x
216 201.4 259.7 29.4 31.4 8.3x
512 725.2 OOM (32 GB) 57.3 62.0

Eager per-op attribution (T4, 64 atoms, torch.no_grad, CUDA-synced): DPA-2 call_graph = 172 ms vs dense call = 41 ms, of which LocalAtten 104 ms + Atten2Map 57 ms + Atten2MultiHeadApply 19 ms — every non-attention graph op (conv, grrg, symmetrization, pair enumeration) is <= 8 ms. With smooth=False the total only drops to 151 ms: the cost is the graph attention structure, not the smooth slot-occupancy softmax (which adds a secondary +14% eager / +20-29% compiled).

Cross-checks:

  • DPA-1 attn_layer=0: graph 5.7 / 18.9 ms vs dense 4.7 / 12.5 ms (64 / 216 atoms) -> 1.2-1.5x, graph is fine without attention.
  • DPA-1 attn_layer=2: graph 176 ms vs dense 20.5 ms at 64 atoms (8.6x), and OOM at 216 atoms on a 15 GB T4 -- a single 5.14 GiB allocation (P = E*nnei ~ 5.5M pairs x attn=128 channels, fp64) in _graph_attention_one_layer.
  • The carry-all graph builder is exonerated: 4.1 / 4.4 / 8.1 / 17.9 ms at 64 / 216 / 512 / 1000 atoms -- near-linear and faster-scaling than the dense route's extend_coord_with_ghosts + build_neighbor_list (2.1 / 8.5 / 41 / 152 ms).

Root cause

The ragged pair-axis attention design (NeighborGraph PR-D, "Approach A"): enumerate P = sum(nnei^2) center pairs, materialize (P, channels) intermediates, and run segment-softmax + index_add over them. This is gather/scatter- and allocation-bound (fp64 atomics included), while dense attention is a batched (nloc, nnei, nnei) cuBLAS GEMM, compute-bound. It explains the super-linear compiled scaling and both OOMs.

Proposed fix: block-dense attention on the graph

Gather per-center compact (nloc, nnei_max, c) blocks through the already-existing destination CSR (destination_order / destination_row_ptr), run the DENSE attention math (batched GEMMs, per-row masks for ragged tails), and scatter back to the edge axis. Raggedness then lives only at the gather/scatter boundary; the P axis disappears entirely. Applies to:

  • DPA-1 se_atten graph attention (_graph_attention_one_layer),
  • DPA-2 Atten2Map / Atten2MultiHeadApply / Atten2EquiVarApply / LocalAtten graph twins.

Numerics note: the smooth-attention semantics (signed phantom count + slot-occupancy denominator) are layout-independent and carry over unchanged; per-op parity tests against the current ragged reference pin the transform.

Cheaper complements, independent of the above:

  • fused attention kernels (the existing DPA-1 CUDA/Triton kernels deliberately cover only the attention-free factorizable path);
  • per-forward caching of the slot-occupancy theta (geometry-only and layer-invariant: computed nlayers x 2 times today, reclaiming most of the +20-29% smooth overhead);
  • an fp32 attention path.

Scope

Follow-up to #5779 (dpa2 graph-native full-stack) and #5715 (PR-D graph attention); performance-only, no semantics change. The dense route and attention-free graph configs are unaffected.

Metadata

Metadata

Assignees

No one assigned

    Type

    No type

    Fields

    No fields configured for issues without a type.

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions