@@ -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 ),
0 commit comments