Skip to content
AI360Xpert
Beta

Equivariant Graph Neural Networks

Equivariant Graph Neural Networks guarantee that rotating or translating 3D atomic coordinates rotates or translates the predicted vectors identically.

EGNN processes invariant scalar features alongside 3D coordinates using radial difference vectors to guarantee exact E(n) rotational and translational equivariance
EGNN processes invariant scalar features alongside 3D coordinates using radial difference vectors to guarantee exact E(n) rotational and translational equivariance

Why Does This Exist?

When applying machine learning to physical systems—such as molecular drug discovery, protein folding, material crystal design, and n-body gravitational simulations—data is fundamentally embedded in 3D Euclidean space. Molecules do not possess an arbitrary coordinate frame: if an ethanol molecule rotates 45 degrees in space, its internal quantum potential energy (a scalar) remains strictly unchanged, while the atomic force vectors acting on its carbon atoms (3D vectors) must rotate by that exact same 45 degrees.

Standard Graph Neural Networks fail this physical requirement. If you feed raw Cartesian coordinates (x,y,z)(x, y, z) into standard MLPs or MPNNs, the network treats the coordinates as arbitrary numbers; rotating the molecule changes the predictions completely, forcing models to rely on data augmentation that requires millions of rotated training samples without ever guaranteeing physical consistency. Earlier equivariant models (such as Cormorant or TFN) solved this by projecting representations onto spherical harmonics and computing Clebsch-Gordan tensor products, but these required complex implementations and high computational overhead. Equivariant Graph Neural Networks (EGNN; Satorras et al., 2021) solved this by demonstrating that exact E(n)E(n) equivariance can be achieved using simple scalar MLPs and radial vector differences.

Think of It Like This

A weather compass on a spinning sailboat

Imagine you are standing on the deck of a sailboat navigating through a fog bank. You have two instruments: a mercury thermometer and a magnetic compass.

  • Invariance: When the boat turns 90 degrees to starboard, the air temperature does not change. The thermometer reading is an invariant scalar: its value is unaffected by the orientation of the boat.
  • Equivariance: When the boat turns 90 degrees to starboard, the compass needle does not freeze relative to the boat; it rotates 90 degrees in your local field of view so that it continues pointing toward true magnetic North. The compass vector is an equivariant vector: transforming your reference frame transforms the measurement by the exact same group action.

An EGNN treats atomic charges and element types like thermometers (invariant) and atomic positions and force vectors like compasses (equivariant).

How It Actually Works

Dual-Stream Invariant and Equivariant Message Passing

An Equivariant Graph Neural Network operates over an nn-dimensional Euclidean group E(n)E(n), which consists of all translations, rotations, and reflections in Rn\mathbb{R}^n (typically n=3n = 3). A transformation g∈E(n)g \in E(n) acts on a coordinate x∈R3\mathbf{x} \in \mathbb{R}^3 via:

g⋅x=Rx+tg \cdot \mathbf{x} = \mathbf{R} \mathbf{x} + \mathbf{t}

where R∈R3×3\mathbf{R} \in \mathbb{R}^{3 \times 3} is an orthogonal matrix (R⊤R=I\mathbf{R}^\top \mathbf{R} = \mathbf{I}, det⁡(R)=±1\det(\mathbf{R}) = \pm 1) and t∈R3\mathbf{t} \in \mathbb{R}^3 is a translation vector.

Every node i∈Vi \in V carries two distinct streams of representations through layer ll:

  1. A scalar hidden state hi(l)∈Rd\mathbf{h}_i^{(l)} \in \mathbb{R}^d (invariant).
  2. A geometric coordinate vector xi(l)∈R3\mathbf{x}_i^{(l)} \in \mathbb{R}^3 (equivariant).

An EGNN layer updates both representations across four steps:

1. Invariant Squared Distance Computation

For every pair of connected nodes ii and jj, compute the Euclidean squared distance:

dij2=∥xi(l)−xj(l)∥22d_{ij}^2 = \left\| \mathbf{x}_i^{(l)} - \mathbf{x}_j^{(l)} \right\|_2^2

Because ∥Rxi+t−(Rxj+t)∥2=∥R(xi−xj)∥2=(xi−xj)⊤R⊤R(xi−xj)=dij2\|\mathbf{R} \mathbf{x}_i + \mathbf{t} - (\mathbf{R} \mathbf{x}_j + \mathbf{t})\|^2 = \|\mathbf{R}(\mathbf{x}_i - \mathbf{x}_j)\|^2 = (\mathbf{x}_i - \mathbf{x}_j)^\top \mathbf{R}^\top \mathbf{R} (\mathbf{x}_i - \mathbf{x}_j) = d_{ij}^2, the distance dij2d_{ij}^2 is strictly invariant under E(n)E(n).

2. Invariant Message Function ϕm\phi_m

Calculate edge messages using an MLP ϕm\phi_m operating on invariant inputs:

mij=ϕm(hi(l),hj(l),dij2,aij)\mathbf{m}_{ij} = \phi_m\left(\mathbf{h}_i^{(l)}, \mathbf{h}_j^{(l)}, d_{ij}^2, a_{ij}\right)

where aija_{ij} are optional edge attributes (like covalent bond types).

3. Equivariant Coordinate Update ϕx\phi_x

Update node positions by accumulating displacement vectors (xi(l)−xj(l))(\mathbf{x}_i^{(l)} - \mathbf{x}_j^{(l)}) scaled by a scalar MLP ϕx:Rm→R\phi_x: \mathbb{R}^m \to \mathbb{R}:

xi(l+1)=xi(l)+∑j∈N(i)1∣N(i)∣(xi(l)−xj(l))ϕx(mij)\mathbf{x}_i^{(l+1)} = \mathbf{x}_i^{(l)} + \sum_{j \in \mathcal{N}(i)} \frac{1}{|\mathcal{N}(i)|} (\mathbf{x}_i^{(l)} - \mathbf{x}_j^{(l)}) \phi_x(\mathbf{m}_{ij})

Under transformation x′=Rx+t\mathbf{x}' = \mathbf{R}\mathbf{x} + \mathbf{t}, the updated coordinate transforms as: xi′(l+1)=Rxi(l)+t+∑jR(xi(l)−xj(l))ϕx(mij)=Rxi(l+1)+t\mathbf{x}_i'^{(l+1)} = \mathbf{R}\mathbf{x}_i^{(l)} + \mathbf{t} + \sum_{j} \mathbf{R}(\mathbf{x}_i^{(l)} - \mathbf{x}_j^{(l)}) \phi_x(\mathbf{m}_{ij}) = \mathbf{R} \mathbf{x}_i^{(l+1)} + \mathbf{t} Equivariance holds without using spherical harmonics or group representation theory.

4. Invariant Node Feature Update ϕh\phi_h

Aggregate messages and update the scalar representations via MLP ϕh\phi_h:

mi=∑j∈N(i)mij,hi(l+1)=ϕh(hi(l),mi)\mathbf{m}_i = \sum_{j \in \mathcal{N}(i)} \mathbf{m}_{ij}, \qquad \mathbf{h}_i^{(l+1)} = \phi_h\left(\mathbf{h}_i^{(l)}, \mathbf{m}_i\right)

Worked Example

Consider a 2-node physical bond:

  • Node 1 at x1=[0.0,0.0,0.0]⊤\mathbf{x}_1 = [0.0, 0.0, 0.0]^\top, Node 2 at x2=[2.0,0.0,0.0]⊤\mathbf{x}_2 = [2.0, 0.0, 0.0]^\top.
  • Scalar features: h1=1.0,h2=1.0h_1 = 1.0, h_2 = 1.0.
  • Scalar coordinate function: ϕx(m)=0.1\phi_x(m) = 0.1 (constant force coefficient).
  1. Original Frame Update:

    • Relative displacement vector: x1−x2=[0.0−2.00.0−0.00.0−0.0]=[−2.00.00.0]\mathbf{x}_1 - \mathbf{x}_2 = \begin{bmatrix} 0.0 - 2.0 \\ 0.0 - 0.0 \\ 0.0 - 0.0 \end{bmatrix} = \begin{bmatrix} -2.0 \\ 0.0 \\ 0.0 \end{bmatrix}
    • Updated coordinate for node 1: x1(1)=x1(0)+(x1−x2)ϕx=[0.00.00.0]+[−2.00.00.0]×0.1=[−0.20.00.0]\mathbf{x}_1^{(1)} = \mathbf{x}_1^{(0)} + (\mathbf{x}_1 - \mathbf{x}_2) \phi_x = \begin{bmatrix} 0.0 \\ 0.0 \\ 0.0 \end{bmatrix} + \begin{bmatrix} -2.0 \\ 0.0 \\ 0.0 \end{bmatrix} \times 0.1 = \begin{bmatrix} -0.2 \\ 0.0 \\ 0.0 \end{bmatrix}
  2. Rotated Frame Verification: Now rotate the entire system 90∘90^\circ counterclockwise around the zz-axis (R=[0−10100001]\mathbf{R} = \begin{bmatrix} 0 & -1 & 0 \\ 1 & 0 & 0 \\ 0 & 0 & 1 \end{bmatrix}): x1rot=[0.00.00.0],x2rot=[0.02.00.0]\mathbf{x}_1^{\text{rot}} = \begin{bmatrix} 0.0 \\ 0.0 \\ 0.0 \end{bmatrix}, \quad \mathbf{x}_2^{\text{rot}} = \begin{bmatrix} 0.0 \\ 2.0 \\ 0.0 \end{bmatrix}

    • Rotated displacement vector: x1rot−x2rot=[0.0−2.00.0]\mathbf{x}_1^{\text{rot}} - \mathbf{x}_2^{\text{rot}} = \begin{bmatrix} 0.0 \\ -2.0 \\ 0.0 \end{bmatrix}
    • Updated coordinate under EGNN rule: x1rot,(1)=[0.00.00.0]+[0.0−2.00.0]×0.1=[0.0−0.20.0]\mathbf{x}_1^{\text{rot},(1)} = \begin{bmatrix} 0.0 \\ 0.0 \\ 0.0 \end{bmatrix} + \begin{bmatrix} 0.0 \\ -2.0 \\ 0.0 \end{bmatrix} \times 0.1 = \begin{bmatrix} 0.0 \\ -0.2 \\ 0.0 \end{bmatrix}
    • Checking Rx1(1)\mathbf{R} \mathbf{x}_1^{(1)}: R[−0.20.00.0]=[0.0−0.20.0]=x1rot,(1)\mathbf{R} \begin{bmatrix} -0.2 \\ 0.0 \\ 0.0 \end{bmatrix} = \begin{bmatrix} 0.0 \\ -0.2 \\ 0.0 \end{bmatrix} = \mathbf{x}_1^{\text{rot},(1)} The rotation commutes with the neural update.

Code

from typing import Tupleimport numpy as np

class SimpleEGNNLayer:    """Minimal Equivariant Graph Neural Network layer in NumPy."""
    def __init__(self, feat_dim: int) -> None:        rng = np.random.default_rng(seed=42)        # phi_m: predicts scalar message from (h_i, h_j, dist_sq)        self.w_m: np.ndarray = rng.standard_normal((2 * feat_dim + 1, feat_dim)) * 0.1        # phi_x: predicts scalar multiplier for radial displacement vector        self.w_x: np.ndarray = rng.standard_normal((feat_dim, 1)) * 0.1        # phi_h: updates scalar features        self.w_h: np.ndarray = rng.standard_normal((2 * feat_dim, feat_dim)) * 0.1
    def forward(        self,        h: np.ndarray,  # (N, feat_dim) scalar features        x: np.ndarray,  # (N, 3) 3D Cartesian coordinates        edges: list[Tuple[int, int]],    ) -> Tuple[np.ndarray, np.ndarray]:        num_nodes = h.shape[0]        coord_updates = np.zeros_like(x)        msg_sums = np.zeros_like(h)
        for i, j in edges:            # 1. Invariant squared distance            diff = x[i] - x[j]            dist_sq = np.sum(diff ** 2, keepdims=True)
            # 2. Invariant message            msg_input = np.concatenate([h[i], h[j], dist_sq], axis=0)            m_ij = np.maximum(0.0, msg_input @ self.w_m)  # ReLU            msg_sums[i] += m_ij
            # 3. Equivariant coordinate displacement            phi_x_val = m_ij @ self.w_x  # Scalar weight            coord_updates[i] += diff * phi_x_val
        # Update coordinates and scalar features        new_x = x + coord_updates        h_update_input = np.concatenate([h, msg_sums], axis=1)        new_h = np.maximum(0.0, h_update_input @ self.w_h)        return new_h, new_x

# Verification: Rotate coordinates and verify equivariancelayer = SimpleEGNNLayer(feat_dim=2)h_in = np.array([[1.0, 0.5], [1.0, 0.5]])x_orig = np.array([[0.0, 0.0, 0.0], [2.0, 0.0, 0.0]])edges = [(0, 1), (1, 0)]
# 90-degree z-rotation matrix Rrot_matrix = np.array([[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 1.0]])x_rot = x_orig @ rot_matrix.T
_, out_x_orig = layer.forward(h_in, x_orig, edges)_, out_x_rot = layer.forward(h_in, x_rot, edges)
# Commutative check: R @ out_x_orig vs out_x_rotexpected_rot_out = out_x_orig @ rot_matrix.Tdiff = np.max(np.abs(expected_rot_out - out_x_rot))print("Max Equivariance Discrepancy:", np.round(diff, 8))# -> Max Equivariance Discrepancy: 0.0 (exact numerical equivariance!)

Watch Out For

Chirality blindness under E(n) reflection equivariance

Because EGNN relies exclusively on squared Euclidean distances ∥xi−xj∥2\|\mathbf{x}_i - \mathbf{x}_j\|^2, it is equivariant to the full Euclidean group E(3)E(3), which includes spatial reflections (parity inversions x↦−x\mathbf{x} \mapsto -\mathbf{x}). In organic chemistry and pharmacology, enantiomers (mirror-image molecules like (R)(R)- and (S)(S)-thalidomide) share identical pairwise distances but have radically different biological binding activities.

If your problem domain involves chiral stereochemistry, an E(3)E(3)-equivariant model cannot distinguish between left-handed and right-handed stereocenters. To differentiate enantiomers, restrict the symmetry group to SE(3)SE(3) (Special Euclidean group with reflections excluded) by incorporating pseudoscalar signed triple products (xi−xj)⋅((xi−xk)×(xi−xl))(\mathbf{x}_i - \mathbf{x}_j) \cdot ((\mathbf{x}_i - \mathbf{x}_k) \times (\mathbf{x}_i - \mathbf{x}_l)) or dihedral angle encodings.

The Quick Version

  • Equivariant GNNs guarantee that rotating or translating 3D input coordinates transforms predicted vector outputs by that exact same geometric rotation or translation.
  • Scalar node features (like mass or charge) remain strictly invariant, while coordinate updates are scaled along relative displacement vectors.
  • EGNN achieves exact E(n)E(n) equivariance with standard MLPs and O(∣E∣)\mathcal{O}(|E|) complexity, avoiding computationally heavy spherical harmonics.