cuequivariance-jax

Installation
SKILL.md

cuequivariance_jax: Executing Equivariant Polynomials in JAX

Overview

cuequivariance_jax (imported as cuex) executes cuequivariance polynomials on GPU via JAX. It provides:

  1. Core primitive: cuex.segmented_polynomial() — JAX primitive with full AD/vmap/JIT support
  2. Two data representations (both built on segmented_polynomial):
    • cuex.equivariant_polynomial() + RepArray — the original interface, a single contiguous array with representation metadata
    • cuex.ir_dict module — dict[Irrep, Array] interface, uses IrDictPolynomial descriptors, works naturally with jax.tree
  3. NNX layers: cuex.nnx module — Flax NNX Module wrappers using dict[Irrep, Array]

Execution methods

Method Backend Requirements
"naive" Pure JAX Always works, any platform
"uniform_1d" CUDA kernel GPU, all segments uniform shape within each operand, single mode
"indexed_linear" CUDA kernel GPU, linear operations with cuex.Repeats indexing
Installs
1
GitHub Stars
403
First Seen
Jul 3, 2026
cuequivariance-jax — nvidia/cuequivariance