Skip to content

Commit f4b7141

Browse files
author
Han Wang
committed
feat(dpmodel): center_edge_pairs primitive (shared by attention/angles)
Segment-based (global (E,E) boolean deliberately avoided): compact eager form for carry-all graphs + shape-static nonzero-free form for the center-major static layout (jit/export/make_fx traceable). Part of NeighborGraph PR-D; PR-E angles reuse (unordered, no-self).
1 parent f322d78 commit f4b7141

3 files changed

Lines changed: 309 additions & 0 deletions

File tree

deepmd/dpmodel/utils/neighbor_graph/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,9 @@
3232
node_validity_mask,
3333
pad_and_guard_edges,
3434
)
35+
from .pairs import (
36+
center_edge_pairs,
37+
)
3538
from .segment import (
3639
segment_max,
3740
segment_mean,
@@ -44,6 +47,7 @@
4447
"NeighborGraph",
4548
"build_neighbor_graph",
4649
"build_neighbor_graph_ase",
50+
"center_edge_pairs",
4751
"edge_env_mat",
4852
"edge_force_virial",
4953
"frame_id_from_n_node",
Lines changed: 172 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,172 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
"""Pairs of edges sharing a center (``dst``) — the edge-pair axis.
3+
4+
Shared primitive: graph-native attention (NeighborGraph PR-D) uses
5+
``(ordered=True, include_self=True)`` = the full transformer neighbor-pair
6+
square per center; 3-body angles (PR-E) use ``(ordered=False,
7+
include_self=False)``.
8+
9+
Two forms:
10+
11+
- **compact eager** (``static_nnei=None``): segment-based enumeration over the
12+
real edges only — sort edge ids by center, expand each center's Cartesian
13+
square via cumsum offsets. Dynamic ``P = sum(deg**2)``; memory ``O(P)``
14+
(same order as dense attention's ``O(nloc * nnei**2)``). Uses data-dependent
15+
shapes (``nonzero``) so it is EAGER-ONLY.
16+
- **shape-static** (``static_nnei`` set): assumes the center-major static
17+
layout (``E = n_center * static_nnei``, edge ``c * static_nnei + m`` belongs
18+
to center ``c`` — the layout ``from_dense_quartet(compact=False)`` emits).
19+
Pure arange/reshape arithmetic, ``P = n_center * static_nnei**2`` with all
20+
pairs materialized and validity carried by ``pair_mask`` — no data-dependent
21+
ops, so it stays jit/export/make_fx-traceable.
22+
23+
A global ``(E, E)`` same-center boolean is deliberately NOT used: with
24+
``E ~ N * nnei`` it costs ``O(N**2 * nnei**2)`` memory.
25+
"""
26+
27+
from __future__ import annotations
28+
29+
from typing import (
30+
Any,
31+
)
32+
33+
import array_api_compat
34+
35+
from deepmd.dpmodel.array_api import (
36+
Array,
37+
xp_add_at,
38+
)
39+
40+
41+
def center_edge_pairs(
42+
dst: Array,
43+
edge_mask: Array,
44+
n_total: int,
45+
*,
46+
include_self: bool = True,
47+
ordered: bool = True,
48+
static_nnei: int | None = None,
49+
) -> tuple[Array, Array, Array]:
50+
"""Enumerate pairs of edges sharing a center.
51+
52+
Parameters
53+
----------
54+
dst : Array
55+
(E,) int64 center of each edge (``edge_index[1]``).
56+
edge_mask : Array
57+
(E,) bool, real (True) vs padding (False) edges.
58+
n_total : int
59+
Number of centers (bounds ``dst``).
60+
include_self : bool
61+
Keep the ``m == n`` diagonal (transformer self-attention needs it).
62+
ordered : bool
63+
Keep both ``(m, n)`` and ``(n, m)`` (attention: yes, ``q_m . k_n`` is
64+
not symmetric). ``False`` keeps only ``n >= m`` (with
65+
``include_self=False``: ``n > m`` — the angle set).
66+
static_nnei : int | None
67+
``None`` -> compact eager form. Set -> shape-static form assuming the
68+
center-major layout ``E = n_center * static_nnei``.
69+
70+
Returns
71+
-------
72+
query_edge : Array
73+
(P,) int64 edge index of the query (``m``).
74+
key_edge : Array
75+
(P,) int64 edge index of the key (``n``).
76+
pair_mask : Array
77+
(P,) bool; False where either edge is padding or the pair is filtered
78+
by the ``include_self`` / ``ordered`` policy (shape-static form; the
79+
compact form drops such pairs and returns all-True).
80+
"""
81+
xp = array_api_compat.array_namespace(dst)
82+
dev = array_api_compat.device(dst)
83+
if static_nnei is not None:
84+
return _pairs_shape_static(
85+
xp, dev, dst, edge_mask, static_nnei, include_self, ordered
86+
)
87+
return _pairs_compact(xp, dev, dst, edge_mask, n_total, include_self, ordered)
88+
89+
90+
def _pairs_shape_static(
91+
xp: Any,
92+
dev: Any,
93+
dst: Array,
94+
edge_mask: Array,
95+
nn: int,
96+
include_self: bool,
97+
ordered: bool,
98+
) -> tuple[Array, Array, Array]:
99+
e_tot = dst.shape[0]
100+
# (E, nn): every edge queries the nn slots of its own center block
101+
eids = xp.arange(e_tot, dtype=xp.int64, device=dev)
102+
base = (eids // nn) * nn # start of each edge's center block
103+
slots = xp.arange(nn, dtype=xp.int64, device=dev)
104+
q2 = xp.broadcast_to(eids[:, None], (e_tot, nn))
105+
k2 = base[:, None] + slots[None, :]
106+
query_edge = xp.reshape(q2, (-1,))
107+
key_edge = xp.reshape(k2, (-1,))
108+
pair_mask = xp.take(edge_mask, query_edge, axis=0) & xp.take(
109+
edge_mask, key_edge, axis=0
110+
)
111+
if not include_self:
112+
pair_mask = pair_mask & (query_edge != key_edge)
113+
if not ordered:
114+
pair_mask = pair_mask & (key_edge >= query_edge)
115+
return query_edge, key_edge, pair_mask
116+
117+
118+
def _pairs_compact(
119+
xp: Any,
120+
dev: Any,
121+
dst: Array,
122+
edge_mask: Array,
123+
n_total: int,
124+
include_self: bool,
125+
ordered: bool,
126+
) -> tuple[Array, Array, Array]:
127+
empty = (
128+
xp.zeros((0,), dtype=xp.int64, device=dev),
129+
xp.zeros((0,), dtype=xp.int64, device=dev),
130+
xp.zeros((0,), dtype=xp.bool, device=dev),
131+
)
132+
if dst.shape[0] == 0:
133+
return empty
134+
# real edges only, grouped by center (stable sort keeps original order
135+
# within a center — irrelevant for correctness, deterministic for tests)
136+
(real_idx,) = xp.nonzero(edge_mask)
137+
r_tot = real_idx.shape[0]
138+
if r_tot == 0:
139+
return empty
140+
d_real = xp.take(dst, real_idx, axis=0)
141+
order = xp.argsort(d_real, stable=True)
142+
eid = xp.take(real_idx, order, axis=0) # (R,) edge ids, center-grouped
143+
ds = xp.take(d_real, order, axis=0) # (R,) sorted centers
144+
# per-center degree and group start (over the sorted layout)
145+
ones = xp.ones((r_tot,), dtype=xp.int64, device=dev)
146+
counts = xp_add_at(
147+
xp.zeros((n_total,), dtype=xp.int64, device=dev), ds, ones
148+
) # (n_total,)
149+
csum = xp.cumulative_sum(counts)
150+
start = csum - counts # (n_total,) group start per center
151+
deg = xp.take(counts, ds, axis=0) # (R,) degree of each edge's center
152+
# each sorted edge t emits deg[t] pairs; P = sum(deg**2)
153+
query_sorted = xp.repeat(xp.arange(r_tot, dtype=xp.int64, device=dev), deg) # (P,)
154+
# within each query's block, a 0..deg-1 ramp indexes the key group
155+
pair_off = xp.cumulative_sum(deg) - deg # (R,) exclusive prefix of deg
156+
p_tot = query_sorted.shape[0]
157+
ramp = xp.arange(p_tot, dtype=xp.int64, device=dev) - xp.take(
158+
pair_off, query_sorted, axis=0
159+
)
160+
key_sorted = xp.take(start, xp.take(ds, query_sorted, axis=0), axis=0) + ramp
161+
query_edge = xp.take(eid, query_sorted, axis=0)
162+
key_edge = xp.take(eid, key_sorted, axis=0)
163+
keep = xp.ones((p_tot,), dtype=xp.bool, device=dev)
164+
if not include_self:
165+
keep = keep & (query_edge != key_edge)
166+
if not ordered:
167+
keep = keep & (key_edge >= query_edge)
168+
(kept,) = xp.nonzero(keep)
169+
query_edge = xp.take(query_edge, kept, axis=0)
170+
key_edge = xp.take(key_edge, kept, axis=0)
171+
pair_mask = xp.ones((query_edge.shape[0],), dtype=xp.bool, device=dev)
172+
return query_edge, key_edge, pair_mask
Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
# SPDX-License-Identifier: LGPL-3.0-or-later
2+
"""center_edge_pairs: pairs of edges sharing a center (NeighborGraph PR-D/E)."""
3+
4+
import numpy as np
5+
6+
from deepmd.dpmodel.utils.neighbor_graph import (
7+
center_edge_pairs,
8+
)
9+
10+
11+
def _oracle(dst, mask, include_self, ordered):
12+
pairs = []
13+
for m in range(len(dst)):
14+
if not mask[m]:
15+
continue
16+
for n in range(len(dst)):
17+
if not mask[n] or dst[m] != dst[n]:
18+
continue
19+
if not include_self and m == n:
20+
continue
21+
if not ordered and n < m:
22+
continue
23+
pairs.append((m, n))
24+
return set(pairs)
25+
26+
27+
def _got(q, k, pm):
28+
return {(int(q[p]), int(k[p])) for p in range(q.shape[0]) if pm[p]}
29+
30+
31+
class TestCompact:
32+
def test_transformer_all_ordered_with_self(self) -> None:
33+
# 3 edges: dst = [0, 0, 1]; center 0 has edges {0,1}, center 1 has {2}
34+
dst = np.array([0, 0, 1], dtype=np.int64)
35+
edge_mask = np.array([True, True, True])
36+
q, k, pm = center_edge_pairs(dst, edge_mask, 2)
37+
assert _got(q, k, pm) == _oracle([0, 0, 1], [1, 1, 1], True, True)
38+
# center 0: (0,0),(0,1),(1,0),(1,1); center 1: (2,2) => 5 pairs
39+
assert len(_got(q, k, pm)) == 5
40+
41+
def test_unordered_no_self_is_angle_set(self) -> None:
42+
dst = np.array([0, 0, 0], dtype=np.int64)
43+
edge_mask = np.array([True, True, True])
44+
q, k, pm = center_edge_pairs(
45+
dst, edge_mask, 1, include_self=False, ordered=False
46+
)
47+
assert _got(q, k, pm) == {(0, 1), (0, 2), (1, 2)}
48+
49+
def test_ignores_padding_edges(self) -> None:
50+
dst = np.array([0, 0, 0], dtype=np.int64)
51+
edge_mask = np.array([True, True, False]) # 3rd is a guard edge
52+
q, k, pm = center_edge_pairs(dst, edge_mask, 1)
53+
assert _got(q, k, pm) == {(0, 0), (0, 1), (1, 0), (1, 1)}
54+
55+
def test_non_contiguous_center_order(self) -> None:
56+
# edges NOT sorted by center: dst = [1, 0, 1, 0]
57+
dst = np.array([1, 0, 1, 0], dtype=np.int64)
58+
edge_mask = np.array([True, True, True, True])
59+
q, k, pm = center_edge_pairs(dst, edge_mask, 2)
60+
assert _got(q, k, pm) == _oracle([1, 0, 1, 0], [1] * 4, True, True)
61+
62+
def test_empty(self) -> None:
63+
dst = np.zeros((0,), dtype=np.int64)
64+
edge_mask = np.zeros((0,), dtype=bool)
65+
q, k, pm = center_edge_pairs(dst, edge_mask, 3)
66+
assert q.shape[0] == 0 and k.shape[0] == 0 and pm.shape[0] == 0
67+
68+
def test_random_vs_oracle(self) -> None:
69+
rng = np.random.default_rng(7)
70+
dst = rng.integers(0, 5, size=23).astype(np.int64)
71+
edge_mask = rng.random(23) > 0.3
72+
for include_self in (True, False):
73+
for ordered in (True, False):
74+
q, k, pm = center_edge_pairs(
75+
dst, edge_mask, 5, include_self=include_self, ordered=ordered
76+
)
77+
assert _got(q, k, pm) == _oracle(
78+
dst, edge_mask, include_self, ordered
79+
), (include_self, ordered)
80+
81+
def test_torch_matches_numpy(self) -> None:
82+
import torch
83+
84+
dst = np.array([0, 0, 1, 1, 1], dtype=np.int64)
85+
edge_mask = np.array([True, False, True, True, True])
86+
ref = _got(*center_edge_pairs(dst, edge_mask, 2))
87+
q, k, pm = center_edge_pairs(
88+
torch.from_numpy(dst), torch.from_numpy(edge_mask), 2
89+
)
90+
assert _got(q.numpy(), k.numpy(), pm.numpy()) == ref
91+
92+
93+
class TestShapeStatic:
94+
def test_matches_compact(self) -> None:
95+
# center-major static layout: 2 centers x static_nnei=3, edges 2,5 padded
96+
dst = np.array([0, 0, 0, 1, 1, 1], dtype=np.int64)
97+
edge_mask = np.array([True, True, False, True, True, False])
98+
qc, kc, pmc = center_edge_pairs(dst, edge_mask, 2)
99+
qs, ks, pms = center_edge_pairs(dst, edge_mask, 2, static_nnei=3)
100+
assert qs.shape[0] == 2 * 3 * 3 # static P, data-independent
101+
assert _got(qs, ks, pms) == _got(qc, kc, pmc)
102+
103+
def test_flags_and_masking(self) -> None:
104+
dst = np.array([0, 0, 0, 1, 1, 1], dtype=np.int64)
105+
edge_mask = np.array([True, True, True, True, False, False])
106+
for include_self in (True, False):
107+
for ordered in (True, False):
108+
qs, ks, pms = center_edge_pairs(
109+
dst,
110+
edge_mask,
111+
2,
112+
include_self=include_self,
113+
ordered=ordered,
114+
static_nnei=3,
115+
)
116+
assert qs.shape[0] == 2 * 3 * 3 # P static regardless of flags
117+
assert _got(qs, ks, pms) == _oracle(
118+
dst, edge_mask, include_self, ordered
119+
), (include_self, ordered)
120+
121+
def test_torch_matches_numpy(self) -> None:
122+
import torch
123+
124+
dst = np.array([0, 0, 1, 1], dtype=np.int64)
125+
edge_mask = np.array([True, False, True, True])
126+
ref = _got(*center_edge_pairs(dst, edge_mask, 2, static_nnei=2))
127+
q, k, pm = center_edge_pairs(
128+
torch.from_numpy(dst),
129+
torch.from_numpy(edge_mask),
130+
2,
131+
static_nnei=2,
132+
)
133+
assert _got(q.numpy(), k.numpy(), pm.numpy()) == ref

0 commit comments

Comments
 (0)