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
k3_node.models.materials.basis.SphericalBesselFunction
Bases: Layer
Spherical Bessel basis expansion j_0(k * r / cutoff).
Example
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
k3_node.models.materials.basis.ChebyshevRadialBasis
Bases: Layer
Chebyshev radial basis with polynomial cutoff envelope.
Example
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)
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
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
k3_node.models.materials.readout.EdgeSet2Set
Bases: Layer
Iterative content-based attention pooling (Set2Set) for edge features.