Download model/nn/interaction_block.py from OneScience-Group/NequIP: direct link, hf CLI and curl.
- Browser
- Download file 7.66 kB
-
https://huggingface.co/OneScience-Group/NequIP/resolve/main/model/nn/interaction_block.py
- Command line
-
hf download hf://OneScience-Group/NequIP/model/nn/interaction_block.py
-
curl -L -o interaction_block.py https://huggingface.co/OneScience-Group/NequIP/resolve/main/model/nn/interaction_block.py
7.66 kB
| # This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. | |
| """Interaction Block""" | |
| import torch | |
| from e3nn.o3._irreps import Irreps | |
| from e3nn.o3._linear import Linear | |
| from e3nn.o3._tensor_product._sub import FullyConnectedTensorProduct | |
| from onescience.datapipes.materials.nequip import AtomicDataDict | |
| from ._graph_mixin import GraphModuleMixin | |
| from .mlp import ScalarMLPFunction | |
| from ._ghost_exchange_base import NoOpGhostExchangeModule | |
| from ._tp_scatter_base import TensorProductScatter | |
| from .norm import AvgNumNeighborsNorm | |
| from typing import Optional, Sequence, Union, Dict | |
| class InteractionBlock(GraphModuleMixin, torch.nn.Module): | |
| use_sc: bool | |
| def __init__( | |
| self, | |
| irreps_in, | |
| irreps_out, | |
| radial_mlp_depth: int = 1, | |
| radial_mlp_width: int = 8, | |
| use_sc: bool = True, | |
| is_first_layer: bool = False, | |
| type_names: Optional[Sequence[str]] = None, | |
| avg_num_neighbors: Optional[Union[float, Dict[str, float]]] = None, | |
| ) -> None: | |
| """InteractionBlock. | |
| Args: | |
| irreps_in: input irreps | |
| irreps_out: output irreps | |
| radial_mlp_depth (int): number of radial layers | |
| radial_mlp_width (int): number of hidden neurons in radial function | |
| use_sc (bool): use self-connection or not | |
| is_first_layer (bool): whether to use first layer (default ``False``) | |
| avg_num_neighbors (float/Dict[str, float]): global (float) or per-type (dict) average number of neighbors | |
| type_names (List[str]): list of type names | |
| """ | |
| super().__init__() | |
| self._init_irreps( | |
| irreps_in=irreps_in, | |
| required_irreps_in=[ | |
| AtomicDataDict.EDGE_EMBEDDING_KEY, | |
| AtomicDataDict.EDGE_ATTRS_KEY, | |
| AtomicDataDict.NODE_FEATURES_KEY, | |
| AtomicDataDict.NODE_ATTRS_KEY, | |
| ], | |
| my_irreps_in={ | |
| AtomicDataDict.EDGE_EMBEDDING_KEY: Irreps( | |
| [ | |
| ( | |
| irreps_in[AtomicDataDict.EDGE_EMBEDDING_KEY].num_irreps, | |
| (0, 1), | |
| ) | |
| ] # (0, 1) is even (invariant) scalars. We are forcing the EDGE_EMBEDDING to be invariant scalars so we can use a dense network | |
| ) | |
| }, | |
| irreps_out={AtomicDataDict.NODE_FEATURES_KEY: irreps_out}, | |
| ) | |
| # === normalization module === | |
| self.avg_num_neighbors_norm = AvgNumNeighborsNorm( | |
| avg_num_neighbors=avg_num_neighbors, type_names=type_names | |
| ) | |
| self.use_sc = use_sc | |
| feature_irreps_in = self.irreps_in[AtomicDataDict.NODE_FEATURES_KEY] | |
| feature_irreps_out = self.irreps_out[AtomicDataDict.NODE_FEATURES_KEY] | |
| irreps_edge_attr = self.irreps_in[AtomicDataDict.EDGE_ATTRS_KEY] | |
| # - Build modules - | |
| self.linear_1 = Linear( | |
| irreps_in=feature_irreps_in, | |
| irreps_out=feature_irreps_in, | |
| internal_weights=True, | |
| shared_weights=True, | |
| ) | |
| irreps_mid = [] | |
| instructions = [] | |
| for i, (mul, ir_in) in enumerate(feature_irreps_in): | |
| for j, (_, ir_edge) in enumerate(irreps_edge_attr): | |
| for ir_out in ir_in * ir_edge: | |
| if ir_out in feature_irreps_out: | |
| k = len(irreps_mid) | |
| irreps_mid.append((mul, ir_out)) | |
| instructions.append((i, j, k, "uvu", True)) | |
| # We sort the output irreps of the tensor product so that we can simplify them | |
| # when they are provided to the second o3.Linear | |
| irreps_mid = Irreps(irreps_mid) | |
| irreps_mid, p, _ = irreps_mid.sort() | |
| # Permute the output indexes of the instructions to match the sorted irreps: | |
| instructions = [ | |
| (i_in1, i_in2, p[i_out], mode, train) | |
| for i_in1, i_in2, i_out, mode, train in instructions | |
| ] | |
| self.tp_scatter = TensorProductScatter( | |
| feature_irreps_in, | |
| irreps_edge_attr, | |
| irreps_mid, | |
| instructions, | |
| ) | |
| # init_irreps already confirmed that the edge embeddding is all invariant scalars | |
| self.edge_mlp = ScalarMLPFunction( | |
| input_dim=self.irreps_in[AtomicDataDict.EDGE_EMBEDDING_KEY].num_irreps, | |
| output_dim=self.tp_scatter.tp.weight_numel, | |
| hidden_layers_depth=radial_mlp_depth, | |
| hidden_layers_width=radial_mlp_width, | |
| nonlinearity="silu", # hardcode SiLU | |
| bias=False, | |
| forward_weight_init=True, | |
| ) | |
| self.linear_2 = Linear( | |
| # irreps_mid has uncoallesed irreps because of the uvu instructions, | |
| # but there's no reason to treat them seperately for the Linear | |
| # Note that normalization of o3.Linear changes if irreps are coallesed | |
| # (likely for the better) | |
| irreps_in=irreps_mid.simplify(), | |
| irreps_out=feature_irreps_out, | |
| internal_weights=True, | |
| shared_weights=True, | |
| ) | |
| self.sc = None | |
| if self.use_sc: | |
| self.sc = FullyConnectedTensorProduct( | |
| feature_irreps_in, | |
| self.irreps_in[AtomicDataDict.NODE_ATTRS_KEY], | |
| feature_irreps_out, | |
| ) | |
| self.ghost_exchange = NoOpGhostExchangeModule( | |
| field=AtomicDataDict.NODE_FEATURES_KEY, irreps_in=self.irreps_in | |
| ) | |
| self.is_first_layer = is_first_layer | |
| def _get_mliap_num_local(self, data: AtomicDataDict.Type) -> int: | |
| return data[AtomicDataDict.LMP_MLIAP_DATA_KEY].nlocal | |
| def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: | |
| if AtomicDataDict.LMP_MLIAP_DATA_KEY in data: | |
| num_local_nodes = self._get_mliap_num_local(data) | |
| else: | |
| num_local_nodes = AtomicDataDict.num_nodes(data) | |
| x = data[AtomicDataDict.NODE_FEATURES_KEY] | |
| # truncate if not first layer | |
| if not self.is_first_layer: | |
| x = x[:num_local_nodes] | |
| if self.sc is not None: | |
| node_attrs = data[AtomicDataDict.NODE_ATTRS_KEY] | |
| # truncate if not first layer | |
| if not self.is_first_layer: | |
| node_attrs = node_attrs[:num_local_nodes] | |
| sc = self.sc(x, node_attrs) | |
| x = self.linear_1(x) | |
| # normalize before TP-scatter | |
| data[AtomicDataDict.NODE_FEATURES_KEY] = x | |
| data = self.avg_num_neighbors_norm(data) | |
| x = data[AtomicDataDict.NODE_FEATURES_KEY] | |
| # === comms for ghost-exchange === | |
| # only done if not first layer | |
| # because initial embedding include ghosts since atom types come with ghosts | |
| if not self.is_first_layer: | |
| data[AtomicDataDict.NODE_FEATURES_KEY] = x | |
| data = self.ghost_exchange(data, ghost_included=False) | |
| x = data[AtomicDataDict.NODE_FEATURES_KEY] | |
| # === TP and scatter === | |
| x = self.tp_scatter( | |
| x=x, | |
| edge_attr=data[AtomicDataDict.EDGE_ATTRS_KEY], | |
| edge_weight=self.edge_mlp(data[AtomicDataDict.EDGE_EMBEDDING_KEY]), | |
| edge_dst=data[AtomicDataDict.EDGE_INDEX_KEY][0], | |
| edge_src=data[AtomicDataDict.EDGE_INDEX_KEY][1], | |
| )[:num_local_nodes] | |
| x = self.linear_2(x) | |
| if self.sc is not None: | |
| x = x + sc | |
| data[AtomicDataDict.NODE_FEATURES_KEY] = x | |
| return data | |