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