cuequivariance-ops-torch
Introduction
cuequivariance_ops_torch provides CUDA kernels for the
cuEquivariance project's PyTorch
components. As such, it contains pytorch bindings to optimized kernels that
cuEquivariance's operations map down to. In general, we advice that you access
those kernels through cuEquivariance, but you may also find them useful on
their own.
Installation
Please install using either pip install cuequivariance-ops-torch-cu12 or
pip install cuequivariance-ops-torch-cu13 (depending on the CUDA toolkit
you wish to use).
Documentation
For detailed usage information of the kernels, please refer to the doc-strings in their respective modules. For higher-level documentations, refer to cuEquivariance.
Pairformer inference primitives
The Torch package also exposes lower-level primitives for integrating Pairformer inference graphs:
attention_pair_bias_mask_from_normalizedprojects an already-normalized pair representation and applies an arbitrary Boolean[B,U,V]mask. The implementation is shape-generic; BF16D=256,H=16, and square sequence lengths 384 and 512 are the performance-validated configurations. FP16 and BF16 inputs use native low-precision tensor-core products with FP32 accumulation; FP32 inputs retain the higher-accuracy TF32x3 path.pair_update_finalize_masked_attention_biasfinalizes a pair residual and, on the shapes admitted for the device's exact compute capability, also returns the next layer's already-masked BF16 attention bias, whichattention_pair_bias_masked_cache(..., is_cached_z_proj=True)consumes in place ofattention_pair_bias. The canonical shape bounds and measurement scope are in_apb_route_policy.py:_pair_update_finalize_admitsfor exact capabilities (9, 0), (10, 0) and (12, 0), and_sm100_apb_block_kf_admitsfor the embedded block. The consumer computes any cache it is handed. On (10, 0) and (10, 3) both halves are taken over by the embedded Kernel Factory APB block (apb_kf_sm100_torch: six CUBINs per exact SM, sm_100a and sm_103a from one source, launched through the CUDA driver, no CuTeDSL at runtime) inside its admitted hull forD_z=384,H=12,DH=64, when the producer has no pair-projection bias and the consumer receives the packed QKV projection with bias, both q/k norm weights and the gate/output biases.pairformer_combined_triangle_multiplicationcombines incoming and outgoing triangle multiplication after caller-provided non-affine LayerNorm with epsilon 1e-5. Its fused direction join applies the affine middle LayerNorm with epsilon 0.03. This operation intentionally fails closed outside BF16 inference withB=1,D=256, sequence length 384 or 512, and an SM100 GPU. The caller guarantees that the second normalized carrier is the separately materialized spatial transpose of the first; the hot path validates metadata and distinct storage but does not compare tensor contents. It returns the base update before any caller-owned peri-LayerNorm, output mask, dropout, or residual addition; model-level dropout must be disabled for inference. Usepack_pairformer_grouped_projectionto convert grouped projection weights to the expected payload/gate layout.
Both operations are registered as torch.library custom operators with fake
implementations for torch.compile tracing. See their Python docstrings for
the complete input contracts.
Calling convention for declined cells
Call triangle_multiplicative_update unconditionally. Under torch.compile
a BF16 cell no kernel owner admits (any architecture, hidden width 256 with a
pair width in 256..512 -- the measured band; inference or training) runs the
package's compiler-visible reference formula, which is
written the way a model's own module is (one gated projection per operand under
autograd, one fused projection GEMM without it, contiguous contraction
operands) and compiles to the same kernels or better, so a declined call costs
no more than the caller's native path (measured on L40S, A100 and H200:
inference 1.10-1.50x, training 0.98-1.01x of RFD4's compiled module; the
generic Triton route the op used to take there measured 0.58-0.88x). An eager
no-grad call no owner admits runs the fused generic route from 2**15 pair rows
(B * S * S) and the reference formula below, where the generic route's
dispatch cost is not repaid. The shape predicate
(triangle_multiplicative_update_inference_shape_is_supported) exists for
callers that must decide before the tensors exist; it is not needed to avoid a
penalty.
Usage
You can import the library from python:
import cuequivariance_ops_torch
Kernels are primarily exposed as torch.nn.Module, but also provide a lower-level
interface as torch.library operators. Generally, the module is responsible
for proper input transformation and initialization, and the operator execute the
kernel.
This allows you to export models using this operations using torch.export,
and running inference on them using TensorRT.
Support and Feedback
Please contact the cuEquivariance developers for any issues you might encounter.
Metadata
Release files for cuequivariance-ops-torch-cu13 0.12.0
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Built distributions (wheels)
Total release size: 11.2 MB
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.14 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64 |
|
SHA-256 checksum How to use checksums |
ed87ac52a50b72833e4ac9aa885ae7f07d4e816ff6df88b6a6e7a92bce565714
|
|
BLAKE2b-256 checksum How to use checksums |
9752bdcabeef66c922c816e133f18233049621f6401719f8e3319d27a42fe835
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.14 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64 |
|
SHA-256 checksum How to use checksums |
e154dc75a31b9d2704cd7efef5d01fec873a9a103386746cbeb63e6ad3b4d7af
|
|
BLAKE2b-256 checksum How to use checksums |
cd31876168c4ca6af4746663c12ce30078e4b3b9447335917565956a05b1c69f
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.13 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64 |
|
SHA-256 checksum How to use checksums |
372c92a1d5e292296f3233e5aaaa514ae0aa9b028e9a6b8e7a9e92837ef0aeaf
|
|
BLAKE2b-256 checksum How to use checksums |
dad68239e65f25c2baac8d970e330b46cf710dbbcd88cadb39babe2456edb864
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.13 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64 |
|
SHA-256 checksum How to use checksums |
ae8aab3a8a27fe6a30c201861853f6290a81bdba3893dc2a5fe838f6755abc41
|
|
BLAKE2b-256 checksum How to use checksums |
6abe2443b46b6af1b0e0a4415e5b2db687112e84aa36b25b6984b1f5f7dbebd7
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.12 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64 |
|
SHA-256 checksum How to use checksums |
4d598751d8befbda60ad63fbff050e9b34077dcaa78fbf5e7e47921f51b0a88f
|
|
BLAKE2b-256 checksum How to use checksums |
2aacd1c37ae2605fb279261e7e4025b26f3e8c154645544480368d621e545c1c
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.12 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64 |
|
SHA-256 checksum How to use checksums |
0667057f25a9b3e80467e58f9f7fd645b267ffcbdc6c5d828f028025b45139ac
|
|
BLAKE2b-256 checksum How to use checksums |
56a6c83cd56b4501999a66a098c31a6397c6b7f31bc31f760e7fe805cf385282
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.11 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64 |
|
SHA-256 checksum How to use checksums |
4e2624f6722972404fd8631c2c19c161ddc710f20180a7eaf56721f1b6bc5331
|
|
BLAKE2b-256 checksum How to use checksums |
505eb4c4bc2c1e80b7d4bcb95ee5c3ccc601336602526a6d4280503556ff9b0d
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.11 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64 |
|
SHA-256 checksum How to use checksums |
c8da5921a3e6bb5b10fcdf06982768237ecc0b82146942532725f13fa4889599
|
|
BLAKE2b-256 checksum How to use checksums |
6d606a8d2a6a5988af8687bfdc9c461ff57df91ae6cdc419a1006f722ef774f5
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.10 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64 |
|
SHA-256 checksum How to use checksums |
84fc658378ccb9759071e7a814d49f528fa7924772fa02191913f120de0acea2
|
|
BLAKE2b-256 checksum How to use checksums |
3c380b58730d024a8970f15671ae93d6fa8009bba5a134d8c9c3f1b357fe21bc
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|
Release files / cuequivariance_ops_torch_cu13-0.12.0-cp310-cp310-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
| Download URL | cuequivariance_ops_torch_cu13-0.12.0-cp310-cp310-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl |
|---|---|
| Size | 1.1 MB |
| Tags | CPython 3.10 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64 |
|
SHA-256 checksum How to use checksums |
0dfc93c50258179a4a10dfc8019f649b7ee5346150eb991f3ded10734f3ed617
|
|
BLAKE2b-256 checksum How to use checksums |
260919a4cabc242e31bc7950f1869c9114d2fe4b0b3a969f80bd693800175dcf
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
twine/7.0.0 CPython/3.14.4
|