Download model/nn/_tp_scatter_base.py from OneScience-Group/NequIP: direct link, hf CLI and curl.
- Browser
- Download file 4.26 kB
-
https://huggingface.co/OneScience-Group/NequIP/resolve/main/model/nn/_tp_scatter_base.py
- Command line
-
hf download hf://OneScience-Group/NequIP/model/nn/_tp_scatter_base.py
-
curl -L -o _tp_scatter_base.py https://huggingface.co/OneScience-Group/NequIP/resolve/main/model/nn/_tp_scatter_base.py
4.26 kB
| # 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 | |
| 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) | |
| 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) | |