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:
- Core primitive:
cuex.segmented_polynomial()— JAX primitive with full AD/vmap/JIT support - Two data representations (both built on
segmented_polynomial):cuex.equivariant_polynomial()+RepArray— the original interface, a single contiguous array with representation metadatacuex.ir_dictmodule —dict[Irrep, Array]interface, usesIrDictPolynomialdescriptors, works naturally withjax.tree
- NNX layers:
cuex.nnxmodule — Flax NNXModulewrappers usingdict[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 |