# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. import torch from e3nn.o3._tensor_product._tensor_product import TensorProduct from .utils import scatter from .model_modifier_utils import replace_submodules, model_modifier class TensorProductScatter(torch.nn.Module): def __init__( self, feature_irreps_in, irreps_edge_attr, irreps_mid, instructions, ) -> None: super().__init__() self.feature_irreps_in = feature_irreps_in self.irreps_edge_attr = irreps_edge_attr self.irreps_mid = irreps_mid self.instructions = instructions self.tp = TensorProduct( feature_irreps_in, irreps_edge_attr, irreps_mid, instructions, shared_weights=False, internal_weights=False, ) self.model_dtype = torch.get_default_dtype() def forward(self, x, edge_attr, edge_weight, edge_dst, edge_src): edge_features = self.tp(x[edge_src], edge_attr, edge_weight) x = scatter(edge_features, edge_dst, dim=0, dim_size=x.size(0)) return x @model_modifier( persistent=False, private=False, unsupported_devices=["cpu"], supported_compile_modes=["torchscript", "aotinductor"], ) @classmethod def enable_OpenEquivariance(cls, model): """ Enable OpenEquivariance tensor product kernel for accelerated NequIP training and inference. For usage instructions, see https://nequip.readthedocs.io/en/latest/guide/accelerations/openequivariance.html """ from ._tp_scatter_oeq import OpenEquivarianceTensorProductScatter from onescience.utils.nequip.internal.dtype import torch_default_dtype from onescience.utils.nequip.internal.versions.torch_versions import _TORCH_GE_2_7 if not _TORCH_GE_2_7: raise RuntimeError("OpenEquivariance requires PyTorch >= 2.7.") _TRAIN_TIME_COMPILE: bool = model.is_compile_graph_model def factory(old): with torch_default_dtype(old.model_dtype): new = OpenEquivarianceTensorProductScatter( feature_irreps_in=old.feature_irreps_in, irreps_edge_attr=old.irreps_edge_attr, irreps_mid=old.irreps_mid, instructions=old.instructions, use_opaque=_TRAIN_TIME_COMPILE, ) # c.f. https://github.com/mir-group/nequip/issues/572 # reuse old.tp to preserve e3nn compiled buffers (_tensor_constant*) # this ensures state dict compatibility whether the modifier is applied or notwa new.tp = old.tp return new return replace_submodules(model, cls, factory) @model_modifier( persistent=False, private=False, unsupported_devices=["cpu"], supported_compile_modes=["torchscript", "aotinductor"], ) @classmethod def enable_CuEquivariance(cls, model): """ [ALPHA SUPPORT] Enable CuEquivariance tensor product kernel for accelerated NequIP inference. For usage instructions, see https://nequip.readthedocs.io/en/latest/guide/accelerations/cuequivariance.html """ from ._tp_scatter_cueq import CuEquivarianceTensorProductScatter from onescience.utils.nequip.internal.dtype import torch_default_dtype def factory(old): with torch_default_dtype(old.model_dtype): new = CuEquivarianceTensorProductScatter( feature_irreps_in=old.feature_irreps_in, irreps_edge_attr=old.irreps_edge_attr, irreps_mid=old.irreps_mid, instructions=old.instructions, ) # c.f. https://github.com/mir-group/nequip/issues/572 # reuse old.tp to preserve e3nn compiled buffers (_tensor_constant*) # this ensures state dict compatibility whether the modifier is applied or not new.tp = old.tp return new return replace_submodules(model, cls, factory)