@@ -148,12 +148,11 @@ def __init__(
148148 self .model = model
149149
150150 # ZBL nuclear repulsion (fixed analytical potential)
151+ self .repuls : NuclearRepulsionPyG | None = None # type: ignore[name-defined]
151152 if calc_repuls :
152153 from matgl .layers ._zbl_pyg import NuclearRepulsionPyG
153154
154155 self .repuls = NuclearRepulsionPyG (float (model .cutoff ))
155- else :
156- self .repuls = None
157156
158157 self .model_config = ModelConfig (
159158 outputs = frozenset ({"energy" , "forces" , "stress" }),
@@ -222,8 +221,8 @@ def from_potential(cls, potential: Potential) -> TensorNetWrapper:
222221
223222 return cls (
224223 model = potential .model ,
225- data_mean = potential .data_mean .clone (),
226- data_std = potential .data_std .clone (),
224+ data_mean = potential .data_mean .clone (), # type: ignore[operator]
225+ data_std = potential .data_std .clone (), # type: ignore[operator]
227226 element_refs = element_refs ,
228227 calc_repuls = getattr (potential , "calc_repuls" , False ),
229228 )
@@ -277,7 +276,7 @@ def adapt_input(self, data: AtomicData | Batch, **kwargs: Any) -> dict[str, Any]
277276 device = data .positions .device
278277 B : int = data .num_graphs
279278
280- node_type = self ._z_to_type [data .atomic_numbers ]
279+ node_type = self ._z_to_type [data .atomic_numbers ] # type: ignore[index]
281280
282281 # nvalchemi (E, 2) -> TensorNet/PyG (2, E)
283282 edge_index = data .neighbor_list .T # [2, E]
0 commit comments