Skip to content

Materials & Crystal Graph Neural Networks

The k3_node.applications.materials module provides multi-backend Keras 3 ports of the MatGL family of interatomic potentials and crystal property predictors, plus built-in checkpoint downloaders that load official pretrained PyTorch weights (hosted on the Hugging Face Hub under the materialyze organization) directly into the Keras models.

All models consume a crystal graph input — a dict (or k3_node.data.Data-like object) with the following keys:

Key Shape Description
pos [N, 3] Atomic Cartesian coordinates.
edge_index [2, E] Bonded-neighbor pairs within the model's cutoff radius.
node_type [N] Atomic numbers (or a mapped element index).
batch [N] Graph id per node, for batched crystals.
line_edge_index [2, L] Bond-pair (angle) connectivity, required by 3-body models (M3GNet, CHGNet, GRACE).
state_attr [num_graphs, S] Optional global/state features (e.g. temperature), used by MEGNet.

Every model outputs a single scalar per crystal (shape ()/(1,) for a single graph, (num_graphs,) when batched) — typically formation energy, band gap, or another crystal-level property.


Models

MEGNet

MEGNet: a graph network with dedicated node/edge/global ("state") update blocks, alternating graph convolution with global-state pooling and broadcasting.

k3_node.models.materials.megnet.MEGNet

Bases: Model

MEGNet materials graph network supporting TensorFlow, PyTorch, and JAX.

Example
import numpy as np
from k3_node.models import MEGNet

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = MEGNet(dim_node_embedding=16, dim_edge_embedding=20, dim_state_embedding=2, nblocks=2,
               hidden_layer_sizes_input=(32, 16), hidden_layer_sizes_conv=(32, 16),
               hidden_layer_sizes_output=(16,))
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

Usage:

from k3_node.models.materials import MEGNet

model = MEGNet(
    dim_node_embedding=16,
    dim_edge_embedding=20,
    dim_state_embedding=2,
    nblocks=2,
    hidden_layer_sizes_input=(32, 16),
    hidden_layer_sizes_conv=(32, 16),
    hidden_layer_sizes_output=(16,),
)
out = model(crystal_graph)  # scalar property prediction

M3GNet

M3GNet: a materials 3-body graph network combining 2-body (bond) and 3-body (angle) interactions for interatomic potential and property prediction, used as MatGL's flagship universal potential.

k3_node.models.materials.m3gnet.M3GNet

Bases: Model

M3GNet materials potential model supporting 3-body angles and multibackend training.

Example
import numpy as np
from k3_node.models import M3GNet

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = M3GNet(dim_node_embedding=16, dim_edge_embedding=16, nblocks=2, units=16, max_n=3, max_l=3)
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

TensorNet

TensorNet: a Cartesian-tensor message-passing network that represents edge features as rank-0/1/2 tensors (scalar, vector, and symmetric-traceless tensor channels) for equivariant potential prediction.

k3_node.models.materials.tensornet.TensorNet

Bases: Model

Cartesian tensor-based equivariant GNN for molecular and crystal potentials.

Example
import numpy as np
from k3_node.models import TensorNet

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = TensorNet(units=16, nblocks=2, num_rbf=16)
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

CHGNet

CHGNet: a charge-informed 3-body graph network that additionally predicts atomic magnetic moments, useful for magnetism-aware materials properties.

k3_node.models.materials.chgnet.CHGNet

Bases: K3NodeHubMixin, Model

Crystal Hamiltonian Graph Neural Network (CHGNet) with charge and angular terms.

Example
import numpy as np
from k3_node.models import CHGNet

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = CHGNet(dim_atom_embedding=16, dim_bond_embedding=16, dim_angle_embedding=16, num_blocks=2,
               atom_conv_hidden_dims=(16,), bond_conv_hidden_dims=(16,))
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

SO3Net

An SO(3)-equivariant network built on real spherical harmonics convolutions, for rotation-equivariant potential prediction.

k3_node.models.materials.so3net.SO3Net

Bases: Model

SO(3)-equivariant representation model using spherical harmonics.

Example
import numpy as np
from k3_node.models import SO3Net

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = SO3Net(units=16, nblocks=2, lmax=2, num_rbf=16)
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

GRACE

A graph atomic cluster expansion (ACE)-style network combining a spherical-harmonic product basis with tensor-product message passing.

k3_node.models.materials.grace.GRACE

Bases: Model

Graph Atomic Cluster Expansion (GRACE) foundational interatomic potential.

Example
import numpy as np
from k3_node.models import GRACE

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = GRACE(cutoff=5.0, n_rad_base=6, lmax=2, embedding_size=8, max_order=2, nblocks=2,
              readout_hidden=(16,))
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

QET

A charge-equilibration-aware potential (couples a LinearQeq electronegativity-equilibration module with an ElectrostaticPotential energy term) for systems where long-range electrostatics matter.

k3_node.models.materials.qet.QET

Bases: TensorNet

Charge-Equilibration TensorNet (QET) model.

Example
import numpy as np
from k3_node.models import QET

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

model = QET(units=16, nblocks=2, num_rbf=16)
energy = model(structure)  # predicted property (e.g. energy) of the structure
print(tuple(energy.shape))  # (1,)

Model Wrappers

Potential

Wraps any of the above energy models as an interatomic potential, denormalizing predictions with data_mean/data_std and (optionally) computing per-atom forces via autodiff.

k3_node.models.materials.wrappers.Potential

Bases: Model

Interatomic potential wrapping an energy model and computing energies, forces, and stresses.

Example
import numpy as np
from k3_node.models import MEGNet, Potential

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

base = MEGNet(dim_node_embedding=8, dim_edge_embedding=16, nblocks=1)
potential = Potential(model=base, data_mean=-1.5, data_std=0.8)  # interatomic potential wrapper
print(tuple(potential(structure).shape))  # (1,)

TransformedTargetModel

A thin wrapper that applies pred * std + mean to a model's output — this is what pretrained MatGL checkpoints are usually distributed as, since targets are normalized during training.

k3_node.models.materials.wrappers.TransformedTargetModel

Bases: Model

Wraps a model and applies inverse transformation to predictions (e.g., mean/std denormalization).

Example
import numpy as np
from k3_node.models import MEGNet, TransformedTargetModel

# A 4-atom structure: positions, bonds (listed in both directions) and atomic numbers
structure = {
    "pos": np.array([[0.0, 0.0, 0.0], [1.0, 0.5, 0.0], [0.5, 1.2, 0.8], [1.5, 1.5, 1.0]], dtype="float32"),
    "edge_index": np.array([[0, 1, 1, 2, 2, 3, 3, 0], [1, 0, 2, 1, 3, 2, 0, 3]]),
    "line_edge_index": np.array([[0, 1, 2, 3], [1, 2, 3, 0]]),  # bond pairs forming angles
    "node_type": np.array([6, 8, 1, 6]),  # atomic numbers
    "batch": np.zeros(4, dtype="int32"),  # all atoms belong to structure 0
    "state_attr": np.zeros((1, 2), dtype="float32"),  # global state features
}

base = MEGNet(dim_node_embedding=8, dim_edge_embedding=16, nblocks=1)
model = TransformedTargetModel(model=base, mean=5.0, std=2.0)  # outputs base * std + mean
print(tuple(model(structure).shape))  # (1,)

Pretrained Checkpoints

load_model

The recommended one-call entry point: downloads (if needed), parses the MatGL model.json architecture spec, instantiates the matching model class with the right constructor arguments, and loads the pretrained state.pt weights — wrapping the result in a TransformedTargetModel when the checkpoint was trained on normalized targets.

k3_node.models.materials.io.load_model(name_or_path, **kwargs)

Convenience factory to download/load and instantiate any MatGL model.

Usage:

from k3_node.models.materials import load_model, get_available_pretrained_models

print(get_available_pretrained_models())
# e.g. ['CHGNet-PES-MatPES-PBE-2025.2.10', 'M3GNet-Eform-MP-2018.6.1', ...]

model = load_model("M3GNet-Eform-MP-2018.6.1")
formation_energy = model(crystal_graph)

get_available_pretrained_models

Lists the pretrained checkpoints published under the materialyze Hugging Face organization (falls back to a hardcoded list if offline).

k3_node.models.materials.io.get_available_pretrained_models()

Return list of available pretrained materials models.

download_matgl_checkpoint

Downloads a checkpoint's model.json (architecture) and state.pt (weights) files from the Hugging Face Hub without instantiating a model — use this if you want to inspect or manually load a checkpoint.

k3_node.models.materials.io.download_matgl_checkpoint(name_or_repo, folder='checkpoints', log=True)

Download matgl model files (model.json, state.pt) from Hugging Face Hub.

load_matgl_weights

Lower-level weight loader: copies a PyTorch state_dict (or a path to one) into an already-constructed Keras model, matching sublayers by name. Prefer load_model unless you need to load weights into a model you built yourself (e.g. for fine-tuning with a different output head — see the fine-tuning recipes).

k3_node.models.materials.io.load_matgl_weights(model, state_dict_or_path, log=False)

Load PyTorch checkpoint weights into multi-backend Keras 3 MatGL model.


Building Blocks

Lower-level components used internally by the models above — useful when assembling a custom materials architecture.

Radial & Angular Basis Functions

k3_node.models.materials.basis.GaussianExpansion

Bases: Layer

Gaussian Radial Basis Function expansion.

Example
import numpy as np
from k3_node.models import MatGLGaussianExpansion

r = np.array([0.9, 1.3, 2.1, 2.8, 3.5, 4.2, 1.1, 1.8], dtype="float32")  # 8 bond lengths

basis = MatGLGaussianExpansion(initial=0.0, final=5.0, num_centers=20, width=0.5)
print(tuple(basis(r).shape))  # (8, 20): expanded features per bond

k3_node.models.materials.basis.BondExpansion

Bases: Layer

Radial basis function dispatcher for pair distances.

Example
import numpy as np
from k3_node.models import MatGLBondExpansion

r = np.array([0.9, 1.3, 2.1, 2.8, 3.5, 4.2, 1.1, 1.8], dtype="float32")  # 8 bond lengths

basis = MatGLBondExpansion(rbf_type="SphericalBessel", max_n=3, max_l=3, cutoff=5.0)
print(tuple(basis(r).shape))  # (8, 9): expanded features per bond

k3_node.models.materials.basis.RadialBesselFunction

Bases: Layer

Zeroth-order spherical Bessel function radial basis with optional learnable roots.

Example
import numpy as np
from k3_node.models import MatGLRadialBesselFunction

r = np.array([0.9, 1.3, 2.1, 2.8, 3.5, 4.2, 1.1, 1.8], dtype="float32")  # 8 bond lengths

basis = MatGLRadialBesselFunction(max_n=3, cutoff=5.0)
print(tuple(basis(r).shape))  # (8, 3): expanded features per bond

k3_node.models.materials.basis.SphericalBesselFunction

Bases: Layer

Spherical Bessel basis expansion j_0(k * r / cutoff).

Example
import numpy as np
from k3_node.models import MatGLSphericalBesselFunction

r = np.array([0.9, 1.3, 2.1, 2.8, 3.5, 4.2, 1.1, 1.8], dtype="float32")  # 8 bond lengths

basis = MatGLSphericalBesselFunction(max_l=3, max_n=3, cutoff=5.0)
print(tuple(basis(r).shape))  # (8, 9): expanded features per bond

k3_node.models.materials.basis.SphericalBesselWithHarmonics

Bases: Layer

Spherical Bessel basis combined with angular Legendre polynomials / harmonics.

Example
import numpy as np
from k3_node.models import MatGLSphericalBesselWithHarmonics

r = np.array([1.1, 1.8, 2.5, 3.0], dtype="float32")  # distances of 4 triplets
theta = np.array([0.5, 1.2, 2.0, 2.8], dtype="float32")  # bond angles
basis = MatGLSphericalBesselWithHarmonics(max_n=3, max_l=3, cutoff=5.0)
print(tuple(basis(r, theta).shape))  # (4, 9): max_n * max_l three-body features

k3_node.models.materials.basis.FourierExpansion

Bases: Layer

Fourier expansion of scalar angular features into sine and cosine components.

Example
import numpy as np
from k3_node.models import MatGLFourierExpansion

x = np.array([0.1, 0.8, 1.6, 2.9], dtype="float32")  # e.g. angles
print(tuple(MatGLFourierExpansion(max_f=4)(x).shape))  # (4, 9): sine and cosine features

k3_node.models.materials.basis.ChebyshevRadialBasis

Bases: Layer

Chebyshev radial basis with polynomial cutoff envelope.

Example
import numpy as np
from k3_node.models import MatGLChebyshevRadialBasis

r = np.array([0.9, 1.3, 2.1, 2.8, 3.5, 4.2, 1.1, 1.8], dtype="float32")  # 8 bond lengths

basis = MatGLChebyshevRadialBasis(nfunc=6, cutoff=5.0)
print(tuple(basis(r).shape))  # (8, 6): expanded features per bond

Geometry Utilities

k3_node.models.materials.basis.compute_pair_vector_and_distance(pos, edge_index, pbc_offshift=None)

Compute pair displacement vectors and pairwise Euclidean distances.

Parameters:

Name Type Description Default
pos

Atomic positions of shape [num_nodes, 3].

required
edge_index

COO edge indices of shape [2, num_edges].

required
pbc_offshift

Periodic boundary condition offset displacement of shape [num_edges, 3] or None.

None

Returns:

Name Type Description
pair_vectors

Displacement vectors of shape [num_edges, 3].

bond_dists

Euclidean distances of shape [num_edges].

k3_node.models.materials.basis.compute_theta(pos, edge_index, line_edge_index, pbc_offshift=None)

Compute bond angles for triplets in directed line graph.

k3_node.models.materials.basis.compute_theta_and_phi(pos, edge_index, line_edge_index, pbc_offshift=None)

Compute theta (azimuthal) and phi (polar) angles for triplets.

k3_node.models.materials.basis.polynomial_cutoff(r, cutoff, exponent=3)

Envelope polynomial function that ensures a smooth cutoff.

k3_node.models.materials.basis.cosine_cutoff(r, cutoff)

Cosine cutoff function.

Core Layers

k3_node.models.materials.core.EmbeddingBlock

Bases: Layer

Embeddings for nodes (atoms), edges (bonds), and global states.

k3_node.models.materials.core.MLP

Bases: Layer

Multi-layer perceptron compatible with multi-backend Keras 3.

k3_node.models.materials.core.GatedMLP

Bases: Layer

Gated multi-layer perceptron: layer(x) * sigmoid(gate(x)).

k3_node.models.materials.core.SoftPlus2

Bases: Layer

SoftPlus2 activation: log(exp(x) + 1) - log(2). Zero at the origin.

k3_node.models.materials.core.SoftExponential

Bases: Layer

Soft exponential activation with learnable alpha.

Tensor Utilities (TensorNet)

k3_node.models.materials.core.vector_to_skewtensor(vector)

Create skew-symmetric 3x3 tensor from a 3D vector.

[0, -v_z, v_y]
[v_z, 0, -v_x]
[-v_y, v_x, 0]

k3_node.models.materials.core.vector_to_symtensor(vector)

Create symmetric traceless tensor from outer product of 3D vector with itself.

k3_node.models.materials.core.decompose_tensor(tensor)

Decompose 3x3 Cartesian tensor into scalar (I), skew-symmetric (A), and symmetric traceless (S).

k3_node.models.materials.core.new_radial_tensor(scalars, skew, traceless, f_I, f_A, f_S)

Multiply irreducible tensor components by radial invariant features.

k3_node.models.materials.core.tensor_norm(tensor)

Computes Frobenius norm squared across last two dimensions (3, 3).

Readout Heads

k3_node.models.materials.readout.ReduceReadOut

Bases: Layer

Pool node features into graph features via sum, mean, or max reduction.

Example
import numpy as np
from k3_node.models import MatGLReduceReadOut

node_feat = np.random.rand(6, 16).astype("float32")
batch = np.array([0, 0, 0, 1, 1, 1])  # two structures with 3 atoms each

print(tuple(MatGLReduceReadOut(op="mean")(node_feat, batch=batch, num_graphs=2).shape))  # (2, 16)

k3_node.models.materials.readout.WeightedReadOut

Bases: Layer

Feed node features through a GatedMLP to predict atomic properties.

Example
import numpy as np
from k3_node.models import MatGLWeightedReadOut

node_feat = np.random.rand(6, 16).astype("float32")
batch = np.array([0, 0, 0, 1, 1, 1])  # two structures with 3 atoms each

print(tuple(MatGLWeightedReadOut(in_feats=16, dims=[16], num_targets=1)(node_feat).shape))  # (6, 1): per-atom outputs

k3_node.models.materials.readout.WeightedAtomReadOut

Bases: Layer

Weighted atom readout for whole-graph properties with normalized learned weights.

Example
import numpy as np
from k3_node.models import MatGLWeightedAtomReadOut

node_feat = np.random.rand(6, 16).astype("float32")
batch = np.array([0, 0, 0, 1, 1, 1])  # two structures with 3 atoms each

readout = MatGLWeightedAtomReadOut(in_feats=16, dims=[16, 8])
print(tuple(readout(node_feat, batch=batch, num_graphs=2).shape))  # (2, 8)

k3_node.models.materials.readout.Set2SetReadOut

Bases: Layer

Iterative content-based attention pooling (Set2Set) for nodes.

Example
import numpy as np
from k3_node.models import MatGLSet2SetReadOut

node_feat = np.random.rand(6, 16).astype("float32")
batch = np.array([0, 0, 0, 1, 1, 1])  # two structures with 3 atoms each

print(tuple(MatGLSet2SetReadOut(in_channels=16)(node_feat, batch=batch, num_graphs=2).shape))  # (2, 32)

k3_node.models.materials.readout.EdgeSet2Set

Bases: Layer

Iterative content-based attention pooling (Set2Set) for edge features.

Example
import numpy as np
from k3_node.models import MatGLEdgeSet2Set

edge_feat = np.random.rand(8, 16).astype("float32")
edge_batch = np.array([0, 0, 0, 0, 1, 1, 1, 1])  # bonds of two structures
print(tuple(MatGLEdgeSet2Set(input_dim=16)(edge_feat, edge_batch=edge_batch, num_graphs=2).shape))  # (2, 32)