File size: 7,049 Bytes
3e02ab8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | # This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it.
from typing import Any, Dict, List, Optional
import torch
from e3nn.o3._irreps import Irreps
from e3nn.o3._spherical_harmonics import SphericalHarmonics
from onescience.datapipes.materials.nequip import AtomicDataDict
from onescience.datapipes.materials.nequip._key_registry import get_field_type
from .._graph_mixin import GraphModuleMixin
class AppendVectorFieldEmbed(GraphModuleMixin, torch.nn.Module):
"""Append embedded node or graph vector fields to node features.
Each field is embedded via solid harmonics up to ``l_max``.
The parity of the input vector must be specified per field: ``+1`` for axial vectors
(pseudovectors, e.g. spin, magnetic field) and ``-1`` for polar vectors (e.g. electric field).
Args:
vector_fields: dict mapping field name to its vector parity (+1 or -1).
l_max: maximum l for the solid harmonic embedding of each field.
append_to_node_attrs: if True, keep ``node_attrs`` equal to appended ``node_features``.
irreps_in: input irreps dictionary passed to ``GraphModuleMixin``.
"""
def __init__(
self,
vector_fields: Dict[str, int],
l_max: int,
append_to_node_attrs: bool = True,
irreps_in: Optional[Dict[str, Any]] = None,
):
super().__init__()
irreps_in = {} if irreps_in is None else dict(irreps_in)
self.append_to_node_attrs = append_to_node_attrs
assert AtomicDataDict.NODE_FEATURES_KEY in irreps_in, (
f"`{AtomicDataDict.NODE_FEATURES_KEY}` must be present in `irreps_in`"
)
if self.append_to_node_attrs:
assert AtomicDataDict.NODE_ATTRS_KEY in irreps_in, (
f"`{AtomicDataDict.NODE_ATTRS_KEY}` must be present in `irreps_in` when `append_to_node_attrs=True`"
)
assert len(vector_fields) > 0, "`vector_fields` cannot be empty"
assert all(p in (1, -1) for p in vector_fields.values()), (
"all parity values in `vector_fields` must be +1 (axial) or -1 (polar)"
)
# preserve insertion order for consistent forward indexing
self.vector_fields: List[str] = list(vector_fields.keys())
self.field_kinds: Dict[str, str] = self._validate_fields(self.vector_fields)
# per-field SH modules; e3nn infers irreps_in ("1e" or "1o") from the output irreps
sh_modules = []
extra_irreps = Irreps()
for field, parity in vector_fields.items():
required_irreps = Irreps("1e" if parity == 1 else "1o")
if field in irreps_in:
assert irreps_in[field] == required_irreps, (
f"`{field}` must have irreps {required_irreps} for parity {parity:+d}, "
f"but got {irreps_in[field]}"
)
else:
irreps_in[field] = required_irreps
# degree-l SH of a parity-p vector transforms as (l, p**l):
# axial (p=+1): all even — 0e, 1e, 2e, ...
# polar (p=-1): alternating — 0e, 1o, 2e, ...
# e3nn validates this and auto-infers irreps_in from these labels
field_sh_irreps = Irreps([(1, (l, parity**l)) for l in range(l_max + 1)])
# don't normalize SH for field vectors; this gives solid harmonics
sh_modules.append(
SphericalHarmonics(
field_sh_irreps, normalize=False, normalization="component"
)
)
extra_irreps += field_sh_irreps
self.sh_modules = torch.nn.ModuleList(sh_modules)
irreps_out = {
AtomicDataDict.NODE_FEATURES_KEY: (
irreps_in[AtomicDataDict.NODE_FEATURES_KEY] + extra_irreps
)
}
if self.append_to_node_attrs:
irreps_out[AtomicDataDict.NODE_ATTRS_KEY] = (
irreps_in[AtomicDataDict.NODE_ATTRS_KEY] + extra_irreps
)
required_irreps_in = [AtomicDataDict.NODE_FEATURES_KEY]
if self.append_to_node_attrs:
required_irreps_in.append(AtomicDataDict.NODE_ATTRS_KEY)
required_irreps_in.extend(self.vector_fields)
self._init_irreps(
irreps_in=irreps_in,
required_irreps_in=required_irreps_in,
irreps_out=irreps_out,
)
self.model_dtype = torch.get_default_dtype()
def __repr__(self) -> str:
lines = [f"{self.__class__.__name__}("]
for field, sh in zip(self.vector_fields, self.sh_modules):
lines.append(f" {field}: {sh.irreps_in} -> {sh.irreps_out},")
lines.append(
f" node_features: {self.irreps_in[AtomicDataDict.NODE_FEATURES_KEY]}"
f" -> {self.irreps_out[AtomicDataDict.NODE_FEATURES_KEY]}"
)
lines.append(")")
return "\n".join(lines)
@staticmethod
def _validate_fields(vector_fields: List[str]) -> Dict[str, str]:
assert len(vector_fields) > 0, "`vector_fields` cannot be empty"
field_kinds = {}
for field in vector_fields:
field_kind = get_field_type(field, error_on_unregistered=True)
assert field_kind in ("graph", "node"), (
f"`{field}` has field type `{field_kind}` but only graph/node fields can be appended"
)
field_kinds[field] = field_kind
return field_kinds
def _field_to_per_node(
self,
data: AtomicDataDict.Type,
field: str,
num_nodes: int,
) -> torch.Tensor:
value = data[field].view(-1, 3)
field_kind = self.field_kinds[field]
# short-circuit of node case
if field_kind == "node":
return value
# (num_graph, 3) -> (num_nodes, 3)
if AtomicDataDict.BATCH_KEY in data:
batch = data[AtomicDataDict.BATCH_KEY].view(-1)
return torch.index_select(value, 0, batch)
# unbatched case -> all nodes get same value
return value.expand(num_nodes, 3)
def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type:
node_features = data[AtomicDataDict.NODE_FEATURES_KEY]
embedded_fields = []
for i, sh in enumerate(self.sh_modules):
per_node_vector = self._field_to_per_node(
data=data,
field=self.vector_fields[i],
num_nodes=node_features.size(0),
)
embedded_fields.append(sh(per_node_vector).to(dtype=self.model_dtype))
# build the concatenation input list explicitly to satisfy TorchScript
cat_inputs = [node_features]
for embedded in embedded_fields:
cat_inputs.append(embedded)
node_features = torch.cat(cat_inputs, dim=1)
data[AtomicDataDict.NODE_FEATURES_KEY] = node_features
if self.append_to_node_attrs:
data[AtomicDataDict.NODE_ATTRS_KEY] = node_features
return data
|