Skip to content

Commit b20653c

Browse files
sirmarcelclaude
andauthored
Rename species to atomic_numbers (#19)
we had species for historical reasons, now cleanly renamed to atomic_numbers everywhere --------- Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
1 parent b3dc553 commit b20653c

5 files changed

Lines changed: 33 additions & 20 deletions

File tree

src/petjax/model.py

Lines changed: 21 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ def __call__(
4747
R_ij,
4848
centers,
4949
neighbors,
50-
species,
50+
atomic_numbers,
5151
reverse,
5252
pair_mask,
5353
atom_mask,
@@ -65,7 +65,16 @@ def __call__(
6565
max_atomic_number=self.max_atomic_number,
6666
attention_temperature=self.attention_temperature,
6767
name="backbone",
68-
)(R_ij, centers, neighbors, species, reverse, pair_mask, atom_mask, pair_cutoffs)
68+
)(
69+
R_ij,
70+
centers,
71+
neighbors,
72+
atomic_numbers,
73+
reverse,
74+
pair_mask,
75+
atom_mask,
76+
pair_cutoffs,
77+
)
6978

7079
predictions = Energy(d_head=self.d_head, name="energy_head")(
7180
node, edge, cutoffs, pair_mask, atom_mask
@@ -87,7 +96,9 @@ def __call__(
8796
force_scale = self.param(
8897
"force_scale", nn.initializers.ones, (self.max_atomic_number + 1,)
8998
)
90-
out["forces"] = forces * force_scale[species][:, None] * atom_mask[:, None]
99+
out["forces"] = (
100+
forces * force_scale[atomic_numbers][:, None] * atom_mask[:, None]
101+
)
91102
if self.direct_stress:
92103
stress = DirectStress(d_head=self.d_head, name="stress_head")(
93104
node, edge, cutoffs, pair_mask, atom_mask
@@ -121,7 +132,7 @@ def __call__(
121132
R_ij,
122133
centers,
123134
neighbors,
124-
species,
135+
atomic_numbers,
125136
reverse,
126137
pair_mask,
127138
atom_mask,
@@ -130,7 +141,7 @@ def __call__(
130141
d_pet = self.d_pet
131142
d_node = self.d_node
132143
P = R_ij.shape[0]
133-
N = species.shape[0]
144+
N = atomic_numbers.shape[0]
134145
n = P // N
135146

136147
r_ij = safe_norm(R_ij, axis=-1)
@@ -149,11 +160,11 @@ def __call__(
149160
# Initial edge features; species embeddings index by atomic number Z
150161
# (table size max_atomic_number + 1, row 0 unused).
151162
edge_embed = nn.Embed(self.max_atomic_number + 1, d_pet, name="edge_embedder")
152-
messages = edge_embed(species)[neighbors] * pair_mask[..., None]
163+
messages = edge_embed(atomic_numbers)[neighbors] * pair_mask[..., None]
153164

154165
# Node embedding (feedforward: persists across layers)
155166
node_embed = nn.Embed(self.max_atomic_number + 1, d_node, name="node_embedders_0")
156-
node = node_embed(species)[:, None, :] * atom_mask[:, None, None]
167+
node = node_embed(atomic_numbers)[:, None, :] * atom_mask[:, None, None]
157168

158169
for layer_idx in range(self.num_gnn_layers):
159170
# Geometric features
@@ -179,7 +190,9 @@ def __call__(
179190
d_pet,
180191
name=f"gnn_layers_{layer_idx}_neighbor_embed",
181192
)
182-
neighbor_feats = neighbor_embed(species)[neighbors] * pair_mask[..., None]
193+
neighbor_feats = (
194+
neighbor_embed(atomic_numbers)[neighbors] * pair_mask[..., None]
195+
)
183196
tokens_flat = masked(
184197
MLP(
185198
(d_pet, d_pet),

src/petjax/select.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -56,7 +56,7 @@ def truncate(
5656
structure["others"],
5757
structure["reverse"],
5858
structure["pair_mask"],
59-
structure["species"],
59+
structure["atomic_numbers"],
6060
structure["atom_mask"],
6161
structure["k_sel_sizer"].shape[-1],
6262
num_neighbors_adaptive,
@@ -73,7 +73,7 @@ def truncate_edges(
7373
others,
7474
reverse,
7575
pair_mask,
76-
species,
76+
atomic_numbers,
7777
atom_mask,
7878
k_sel,
7979
num_neighbors_adaptive,
@@ -97,7 +97,7 @@ def truncate_edges(
9797
couples pairs only through their center/other atoms. Requires ``centers``
9898
non-decreasing; ``k_sel`` must be a static int under jit.
9999
"""
100-
N = species.shape[0]
100+
N = atomic_numbers.shape[0]
101101

102102
pair_cutoffs, selected = _select_edges(
103103
R_ij,
@@ -121,7 +121,7 @@ def truncate_edges(
121121
"R_ij": R_ij[sel_to_pair],
122122
"centers": centers[sel_to_pair],
123123
"neighbors": others[sel_to_pair],
124-
"species": species,
124+
"atomic_numbers": atomic_numbers,
125125
"reverse": slot[reverse[sel_to_pair]],
126126
"pair_mask": pair_mask_sel,
127127
"atom_mask": atom_mask,
@@ -130,7 +130,7 @@ def truncate_edges(
130130
return truncated, overflow
131131

132132

133-
def pack_edges(R_ij, centers, others, reverse, pair_mask, species, atom_mask, k):
133+
def pack_edges(R_ij, centers, others, reverse, pair_mask, atomic_numbers, atom_mask, k):
134134
"""Fixed-width pack on a flat NL with precomputed displacements — no
135135
selection: every unmasked pair goes into the rectangular ``[N * k]``
136136
layout. Returns the truncated dict keyed like ``truncate_edges``'s, with
@@ -141,15 +141,15 @@ def pack_edges(R_ij, centers, others, reverse, pair_mask, species, atom_mask, k)
141141
Same layout requirements as ``truncate_edges``: ``centers``
142142
non-decreasing, ``k`` a static int under jit.
143143
"""
144-
N = species.shape[0]
144+
N = atomic_numbers.shape[0]
145145
slot, sel_to_pair, pair_mask_sel, overflow = _pack_selected_to_flat(
146146
pair_mask, centers, N, k
147147
)
148148
truncated = {
149149
"R_ij": R_ij[sel_to_pair],
150150
"centers": centers[sel_to_pair],
151151
"neighbors": others[sel_to_pair],
152-
"species": species,
152+
"atomic_numbers": atomic_numbers,
153153
"reverse": slot[reverse[sel_to_pair]],
154154
"pair_mask": pair_mask_sel,
155155
"atom_mask": atom_mask,

src/petjax/structure.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -103,8 +103,8 @@ def to_structure(
103103
pair_mask[:n_pair_raw] = True
104104
reverse[:n_pair_raw] = reverse_sparse.astype(int_dtype)
105105

106-
species = np.zeros(N_padded, dtype=int_dtype)
107-
species[:n_atoms] = atoms.get_atomic_numbers() # embeddings index by atomic number
106+
atomic_numbers = np.zeros(N_padded, dtype=int_dtype)
107+
atomic_numbers[:n_atoms] = atoms.get_atomic_numbers()
108108

109109
atom_mask = np.zeros(N_padded, dtype=bool)
110110
atom_mask[:n_atoms] = True
@@ -119,7 +119,7 @@ def to_structure(
119119
return {
120120
"positions": positions,
121121
"cell": cell,
122-
"species": species,
122+
"atomic_numbers": atomic_numbers,
123123
"atom_mask": atom_mask,
124124
"centers": centers,
125125
"others": others,

tests/test_calculator.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -252,7 +252,7 @@ def test_pair_cutoffs_none_uses_static_cutoff(model_data, mini_xyz):
252252
structure["others"],
253253
structure["reverse"],
254254
structure["pair_mask"],
255-
structure["species"],
255+
structure["atomic_numbers"],
256256
structure["atom_mask"],
257257
k_sel,
258258
config["num_neighbors_adaptive"],

tests/test_select.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@ def _pack(structure, k):
3232
structure["others"],
3333
structure["reverse"],
3434
structure["pair_mask"],
35-
structure["species"],
35+
structure["atomic_numbers"],
3636
structure["atom_mask"],
3737
k,
3838
)

0 commit comments

Comments
 (0)