Skip to content

Commit 2560148

Browse files
author
Han Wang
committed
feat(dpmodel): pad_and_guard_angles angle-axis padder
1 parent a1061af commit 2560148

3 files changed

Lines changed: 77 additions & 0 deletions

File tree

deepmd/dpmodel/utils/neighbor_graph/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
NeighborGraph,
3131
frame_id_from_n_node,
3232
node_validity_mask,
33+
pad_and_guard_angles,
3334
pad_and_guard_edges,
3435
)
3536
from .pairs import (
@@ -54,6 +55,7 @@
5455
"from_dense_quartet",
5556
"neighbor_graph_from_ijs",
5657
"node_validity_mask",
58+
"pad_and_guard_angles",
5759
"pad_and_guard_edges",
5860
"segment_max",
5961
"segment_mean",

deepmd/dpmodel/utils/neighbor_graph/graph.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -123,6 +123,57 @@ def pad_and_guard_edges(
123123
return ei, ev, edge_mask
124124

125125

126+
def pad_and_guard_angles(
127+
angle_index: Array,
128+
angle_capacity: int | None = None,
129+
min_angles: int = 2,
130+
pad_value: int = 0,
131+
) -> tuple[Array, Array]:
132+
"""Append padding/guard angles as a contiguous suffix and build angle_mask.
133+
134+
Real angles (``angle_index``) stay at the front (compact layout).
135+
Dummy angles point at edge ``pad_value`` (in-range).
136+
137+
Parameters
138+
----------
139+
angle_index
140+
(2, A_real) ``[edge_a, edge_b]`` edge endpoints of the real angles.
141+
angle_capacity
142+
Target angle-axis length ``A_max``. ``None`` (torch dynamic) appends
143+
exactly ``min_angles`` masked dummy angles so the axis has a known lower
144+
bound and shape-stable guards for export; an int (jax static) pads to
145+
``A_max = angle_capacity`` and raises ``ValueError`` on overflow.
146+
min_angles
147+
Number of dummy angles appended when ``angle_capacity is None``.
148+
pad_value
149+
Edge index the dummy angles point at (must be in range).
150+
151+
Returns
152+
-------
153+
angle_index
154+
(2, target) padded angle endpoints.
155+
angle_mask
156+
(target,) boolean mask, ``True`` for the real-angle prefix.
157+
"""
158+
xp = array_api_compat.array_namespace(angle_index)
159+
dev = array_api_compat.device(angle_index)
160+
a_real = angle_index.shape[1]
161+
if angle_capacity is None:
162+
target = a_real + min_angles
163+
else:
164+
if a_real > angle_capacity:
165+
raise ValueError(
166+
f"angle overflow: {a_real} real angles > angle_capacity {angle_capacity}"
167+
)
168+
target = angle_capacity
169+
n_pad = target - a_real
170+
pad_idx = xp.full((2, n_pad), pad_value, dtype=angle_index.dtype, device=dev)
171+
ai = xp.concat([angle_index, pad_idx], axis=1)
172+
arange = xp.arange(target, dtype=angle_index.dtype, device=dev)
173+
angle_mask = arange < a_real
174+
return ai, angle_mask
175+
176+
126177
def frame_id_from_n_node(n_node: Array, n_total: int | None = None) -> Array:
127178
"""Node->frame map for a flat node axis: ``repeat(arange(nf), n_node)``.
128179
Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,24 @@
1+
import numpy as np
2+
import pytest
3+
4+
from deepmd.dpmodel.utils.neighbor_graph import pad_and_guard_angles
5+
6+
7+
def test_pad_angles_dynamic_appends_min_guard():
8+
ai = np.array([[0, 1], [1, 0]], dtype=np.int64) # 2 real angles
9+
out_ai, out_mask = pad_and_guard_angles(ai, angle_capacity=None, min_angles=2)
10+
assert out_ai.shape == (2, 4) # 2 real + 2 guard
11+
np.testing.assert_array_equal(out_mask, [True, True, False, False])
12+
13+
14+
def test_pad_angles_static_capacity():
15+
ai = np.array([[0, 1], [1, 0]], dtype=np.int64)
16+
out_ai, out_mask = pad_and_guard_angles(ai, angle_capacity=5)
17+
assert out_ai.shape == (2, 5)
18+
assert int(out_mask.sum()) == 2
19+
20+
21+
def test_pad_angles_overflow_raises():
22+
ai = np.zeros((2, 6), dtype=np.int64)
23+
with pytest.raises(ValueError):
24+
pad_and_guard_angles(ai, angle_capacity=4)

0 commit comments

Comments
 (0)