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