| name | cuequivariance-torch |
| description | Execute equivariant tensor products in PyTorch using SegmentedPolynomial (naive/uniform_1d/fused_tp/indexed_linear), high-level operations (ChannelWiseTensorProduct, FullyConnectedTensorProduct, Linear, SymmetricContraction, SphericalHarmonics, Rotation), and layers (BatchNorm, FullyConnectedTensorProductConv). Use when writing PyTorch code with cuequivariance. |
cuequivariance_torch: Executing Equivariant Polynomials in PyTorch
Overview
cuequivariance_torch (imported as cuet) executes cuequivariance polynomials on GPU via PyTorch. It provides:
- Core primitive:
cuet.SegmentedPolynomial — torch.nn.Module with multiple CUDA backends
- High-level operations (
torch.nn.Module): ChannelWiseTensorProduct, FullyConnectedTensorProduct, Linear, SymmetricContraction, SphericalHarmonics, Rotation, Inversion
- Layers:
cuet.layers.BatchNorm, cuet.layers.FullyConnectedTensorProductConv (message passing)
- Utilities:
triangle_attention, triangle_multiplicative_update, attention_pair_bias (AlphaFold3-style)
- Export support:
onnx_custom_translation_table(), register_tensorrt_plugins()
Execution methods
| Method | Backend | Requirements |
|---|
"naive" | Pure PyTorch (einsum) | Always works, any platform |
"uniform_1d" | CUDA kernel | GPU, all segments uniform shape within each operand, single mode |
"fused_tp" | CUDA kernel | GPU, 3- or 4-operand contractions, float32/float64 |
"indexed_linear" | CUDA kernel | GPU, linear with indexed weights, sorted indices |
Core primitive: SegmentedPolynomial
import torch
import cuequivariance as cue
import cuequivariance_torch as cuet
e = cue.descriptors.spherical_harmonics(cue.SO3(1), [0, 1, 2])
poly = e.polynomial
sp = cuet.SegmentedPolynomial(poly, method="uniform_1d")
x = torch.randn(batch, 3, device="cuda")
[output] = sp([x])
Inputs, indexing, and scatter
e = cue.descriptors.channelwise_tensor_product(
16 * cue.Irreps("SO3", "0 + 1"),
cue.Irreps("SO3", "0 + 1"),
cue.Irreps("SO3", "0 + 1"),
)
poly = e.polynomial
sp = cuet.SegmentedPolynomial(poly, method="uniform_1d")
w = torch.randn(1, poly.inputs[0].size, device="cuda")
x1 = torch.randn(batch, poly.inputs[1].size, device="cuda")
x2 = torch.randn(batch, poly.inputs[2].size, device="cuda")
[out] = sp([w, x1, x2])
senders = torch.randint(0, num_nodes, (num_edges,), device="cuda")
[out] = sp([w, x1, x2], input_indices={1: senders})
receivers = torch.randint(0, num_nodes, (num_edges,), device="cuda")
[out] = sp(
[w, x1, x2],
input_indices={1: senders},
output_indices={0: receivers},
output_shapes={0: torch.empty(num_nodes, 1, device="cuda")},
)
Math dtype control
sp = cuet.SegmentedPolynomial(poly, method="fused_tp", math_dtype=torch.float32)
High-level operations
All operations are torch.nn.Module subclasses. They wrap SegmentedPolynomial and handle layout transposition automatically.
Memory layout
IrrepsLayout controls memory order within each (mul, ir) block:
cue.mul_ir: data ordered as (mul, ir.dim) — default, compatible with e3nn
cue.ir_mul: data ordered as (ir.dim, mul) — used internally by descriptors
Operations accept layout (applies to all), or per-operand layout_in1, layout_in2, layout_out.
ChannelWiseTensorProduct
Channel-wise tensor product: pairs channels of x1 with channels of x2.
tp = cuet.ChannelWiseTensorProduct(
cue.Irreps("SO3", "32x0 + 32x1"),
cue.Irreps("SO3", "0 + 1"),
layout=cue.mul_ir,
device="cuda",
dtype=torch.float32,
)
x1 = torch.randn(batch, tp.irreps_in1.dim, device="cuda")
x2 = torch.randn(batch, tp.irreps_in2.dim, device="cuda")
out = tp(x1, x2)
tp = cuet.ChannelWiseTensorProduct(
cue.Irreps("SO3", "32x0 + 32x1"),
cue.Irreps("SO3", "0 + 1"),
layout=cue.mul_ir,
shared_weights=False,
device="cuda",
)
w = torch.randn(batch, tp.weight_numel, device="cuda")
out = tp(x1, x2, weight=w)
out = tp(x1, x2, weight=w, indices_1=senders, indices_out=receivers, size_out=num_nodes)
Default method: "uniform_1d" if segments are uniform, else "naive".
FullyConnectedTensorProduct
All input irrep pairs contribute to all output irreps (dense contraction).
tp = cuet.FullyConnectedTensorProduct(
cue.Irreps("O3", "4x0e + 4x1o"),
cue.Irreps("O3", "0e + 1o"),
cue.Irreps("O3", "4x0e + 4x1o"),
layout=cue.mul_ir,
internal_weights=True,
device="cuda",
)
out = tp(x1, x2)
Default method: "fused_tp".
Linear
Equivariant linear layer (weight-only, no second input).
linear = cuet.Linear(
cue.Irreps("SO3", "4x0 + 2x1"),
cue.Irreps("SO3", "3x0 + 5x1"),
layout=cue.mul_ir,
internal_weights=True,
device="cuda",
)
out = linear(x)
linear = cuet.Linear(
irreps_in, irreps_out,
weight_classes=50,
internal_weights=True,
device="cuda",
)
out = linear(x, weight_indices=species_indices)
Default method: "naive". Use method="fused_tp" for CUDA acceleration.
SymmetricContraction
MACE-style symmetric contraction with element-indexed weights.
sc = cuet.SymmetricContraction(
cue.Irreps("O3", "32x0e + 32x1o"),
cue.Irreps("O3", "32x0e"),
contraction_degree=3,
num_elements=95,
layout=cue.ir_mul,
dtype=torch.float32,
device="cuda",
)
out = sc(x, indices)
Default method: "uniform_1d" if segments are uniform, else "naive".
SphericalHarmonics
sh = cuet.SphericalHarmonics(
ls=[0, 1, 2, 3],
normalize=True,
device="cuda",
)
vectors = torch.randn(batch, 3, device="cuda")
out = sh(vectors)
Default method: "uniform_1d".
Rotation and Inversion
rot = cuet.Rotation(
cue.Irreps("SO3", "4x0 + 2x1 + 1x2"),
layout=cue.ir_mul,
device="cuda",
)
gamma = torch.tensor([0.1], device="cuda")
beta = torch.tensor([0.2], device="cuda")
alpha = torch.tensor([0.3], device="cuda")
out = rot(gamma, beta, alpha, x)
encoded = cuet.encode_rotation_angle(angle, ell=3)
beta, alpha = cuet.vector_to_euler_angles(vector)
inv = cuet.Inversion(
cue.Irreps("O3", "4x0e + 2x1o"),
layout=cue.ir_mul,
device="cuda",
)
out = inv(x)
Layers
BatchNorm
Batch normalization for equivariant representations (adapted from e3nn).
bn = cuet.layers.BatchNorm(
cue.Irreps("O3", "4x0e + 4x1o"),
layout=cue.mul_ir,
eps=1e-5,
momentum=0.1,
affine=True,
)
out = bn(x)
FullyConnectedTensorProductConv
Message passing layer for equivariant GNNs (DiffDock-style).
conv = cuet.layers.FullyConnectedTensorProductConv(
in_irreps=cue.Irreps("O3", "4x0e + 4x1o"),
sh_irreps=cue.Irreps("O3", "0e + 1o"),
out_irreps=cue.Irreps("O3", "4x0e + 4x1o"),
mlp_channels=[16, 32, 32],
mlp_activation=torch.nn.ReLU(),
batch_norm=True,
layout=cue.ir_mul,
)
graph = ((src, dst), (num_src_nodes, num_dst_nodes))
out = conv(src_features, edge_sh, edge_emb, graph, reduce="mean")
out = conv(src_features, edge_sh, edge_emb, graph,
src_scalars=src_scalars, dst_scalars=dst_scalars)
Triangle operations (AlphaFold2-style)
Require cuequivariance_ops_torch.
kv_lengths = cuet.mask_to_kv_lengths(prefix_mask)
out = cuet.triangle_attention(q, k, v, bias, scale=scale, kv_lengths=kv_lengths)
out = cuet.triangle_attention(q, k, v, bias, mask=holey_mask, scale=scale)
out = cuet.triangle_multiplicative_update(
x,
mask=mask,
precision=cuet.TriMulPrecision.DEFAULT,
)
out, _ = cuet.attention_pair_bias(
single_repr, pair_repr, mask, num_heads,
w_ln_a, b_ln_a,
w_proj_q, b_proj_q, w_proj_k, w_proj_v,
w_proj_g, w_proj_o, w_proj_z,
w_ln_z=w_ln_z, b_ln_z=b_ln_z,
)
ONNX and TensorRT export
table = cuet.onnx_custom_translation_table()
onnx_program = torch.onnx.export(model, inputs, custom_translation_table=table)
cuet.register_tensorrt_plugins()
Complete GNN example
import torch
import cuequivariance as cue
import cuequivariance_torch as cuet
class SimpleGNN(torch.nn.Module):
def __init__(self, irreps_in, irreps_sh, irreps_out):
super().__init__()
self.tp = cuet.ChannelWiseTensorProduct(
irreps_in, irreps_sh, layout=cue.mul_ir,
shared_weights=False, device="cuda",
)
self.linear = cuet.Linear(
self.tp.irreps_out, irreps_out,
layout=cue.mul_ir, internal_weights=True, device="cuda",
)
self.sh = cuet.SphericalHarmonics(
ls=[ir.l for _, ir in irreps_sh], normalize=True, device="cuda",
)
def forward(self, node_feats, edge_vec, edge_index, num_nodes):
src, dst = edge_index
edge_sh = self.sh(edge_vec)
w = torch.randn(1, self.tp.weight_numel, device=node_feats.device)
messages = self.tp(
node_feats, edge_sh, weight=w,
indices_1=src, indices_2=None,
indices_out=dst, size_out=num_nodes,
)
return self.linear(messages)
Key file locations
| Component | Path |
|---|
SegmentedPolynomial | cuequivariance_torch/primitives/segmented_polynomial.py |
uniform_1d backend | cuequivariance_torch/primitives/segmented_polynomial_uniform_1d.py |
naive backend | cuequivariance_torch/primitives/segmented_polynomial_naive.py |
fused_tp backend | cuequivariance_torch/primitives/segmented_polynomial_fused_tp.py |
indexed_linear backend | cuequivariance_torch/primitives/segmented_polynomial_indexed_linear.py |
ChannelWiseTensorProduct | cuequivariance_torch/operations/tp_channel_wise.py |
FullyConnectedTensorProduct | cuequivariance_torch/operations/tp_fully_connected.py |
Linear | cuequivariance_torch/operations/linear.py |
SymmetricContraction | cuequivariance_torch/operations/symmetric_contraction.py |
SphericalHarmonics | cuequivariance_torch/operations/spherical_harmonics.py |
Rotation / Inversion | cuequivariance_torch/operations/rotation.py |
BatchNorm | cuequivariance_torch/layers/batchnorm.py |
FullyConnectedTensorProductConv | cuequivariance_torch/layers/tp_conv_fully_connected.py |
| Triangle operations | cuequivariance_torch/primitives/triangle.py |
| Layout transposition | cuequivariance_torch/primitives/transpose.py |