NequIP / model /nn /interaction_block.py
dadadaxi's picture
Upload folder using huggingface_hub
3e02ab8 verified
Raw
History Blame Contribute Delete
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
@torch.jit.unused
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