Skip to main content

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_normalized projects an already-normalized pair representation and applies an arbitrary Boolean [B,U,V] mask. The implementation is shape-generic; BF16 D=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_bias finalizes 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, which attention_pair_bias_masked_cache(..., is_cached_z_proj=True) consumes in place of attention_pair_bias. The canonical shape bounds and measurement scope are in _apb_route_policy.py: _pair_update_finalize_admits for exact capabilities (9, 0), (10, 0) and (12, 0), and _sm100_apb_block_kf_admits for 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 for D_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_multiplication combines 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 with B=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. Use pack_pairformer_grouped_projection to 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-cu12 0.12.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Built distributions (wheels)

Table of built distributions (wheels) for cuequivariance-ops-torch-cu12 0.12.0
File
cuequivariance_ops_torch_cu12-0.12.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl CPython 3.14 CPython 3.14 Linux glibc 2.27+ x86-64, Linux glibc 2.28+ x86-64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.14 CPython 3.14 Linux glibc 2.28+ ARM64, Linux glibc 2.26+ ARM64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl CPython 3.13 CPython 3.13 Linux glibc 2.28+ x86-64, Linux glibc 2.27+ x86-64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.13 CPython 3.13 Linux glibc 2.26+ ARM64, Linux glibc 2.28+ ARM64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl CPython 3.12 CPython 3.12 Linux glibc 2.27+ x86-64, Linux glibc 2.28+ x86-64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.12 CPython 3.12 Linux glibc 2.28+ ARM64, Linux glibc 2.26+ ARM64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl CPython 3.11 CPython 3.11 Linux glibc 2.27+ x86-64, Linux glibc 2.28+ x86-64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.11 CPython 3.11 Linux glibc 2.26+ ARM64, Linux glibc 2.28+ ARM64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl CPython 3.10 CPython 3.10 Linux glibc 2.27+ x86-64, Linux glibc 2.28+ x86-64 Details
cuequivariance_ops_torch_cu12-0.12.0-cp310-cp310-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.10 CPython 3.10 Linux glibc 2.28+ ARM64, Linux glibc 2.26+ ARM64 Details

Total release size: 10.2 MB

Release files / cuequivariance_ops_torch_cu12-0.12.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
Size 1.0 MB
Tags CPython 3.14 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
1cd74955462556b6d41c6cd26397f0f4b361a1e207ecdd34148decf0579f4779
BLAKE2b-256 checksum
How to use checksums
35b2104a0f22f47cfc4893875e2ec01d51e5bb660a2fbc482337cecc7031d262
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_cu12-0.12.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp314-cp314-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
Size 1.0 MB
Tags CPython 3.14 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64
SHA-256 checksum
How to use checksums
75b11f69908d98027e2f1bff8cdd3de45443272cf4cda5453f1f59830619c4b5
BLAKE2b-256 checksum
How to use checksums
4778543832020d2e0a7fa14411aa64f6c8a9dcaed85dd42c6d0b9eb7bd9fe476
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_cu12-0.12.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
Size 1.0 MB
Tags CPython 3.13 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
0490373224316c65f015cb27c374c93ca7fefa3c313751dede7ee75743376206
BLAKE2b-256 checksum
How to use checksums
8e5e7001fbe91a3a122ef574852a782de8a8dd9f5eaa0c29ae08bae51992edc7
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_cu12-0.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
Size 1.0 MB
Tags CPython 3.13 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64
SHA-256 checksum
How to use checksums
5e5f86c344db64ba60bab50b4d85f7862269648ff1fcc54d7623acaa95ffc67f
BLAKE2b-256 checksum
How to use checksums
30e07912934936b53aee2acfb8504591af59e3712d208cd2c81f2b2c8eb6fce7
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_cu12-0.12.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
Size 1.0 MB
Tags CPython 3.12 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
3bb7520ae0e54da53eaa2bd0df1547015af10df15ade9922591a5de9f62eb4e4
BLAKE2b-256 checksum
How to use checksums
85adcea0d0211de48fedc365e364a8d0d173f04b1c831adcffb55a990ec45ef0
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_cu12-0.12.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp312-cp312-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
Size 1.0 MB
Tags CPython 3.12 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64
SHA-256 checksum
How to use checksums
b74f702178911952d47f925a8fa90f99fce50ab1f1073247c7615fdffd28ba93
BLAKE2b-256 checksum
How to use checksums
3300d27b12c7a2f6afa09ca17a53f6c10d8f17880487ccd7c802af91727c9c55
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_cu12-0.12.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
Size 1.0 MB
Tags CPython 3.11 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
903784bc76dba9158dbbb04af7862afbda71403d031c0b3b092dfb943dcbef43
BLAKE2b-256 checksum
How to use checksums
a0285a3658c7cefd4c00b4d182696b4117e98ab2bc903b3187ca41eff6444f46
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_cu12-0.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
Size 1.0 MB
Tags CPython 3.11 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64
SHA-256 checksum
How to use checksums
c913e357df0375d795cbf938a19444995bc7a8c3863e24be59e719fabf12f5d7
BLAKE2b-256 checksum
How to use checksums
d8c170ddf77da95d3f835c619aaaf89403b73af45d18e875996a035799491fa6
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_cu12-0.12.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp310-cp310-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl
Size 1.0 MB
Tags CPython 3.10 Linux glibc 2.27+ x86-64 Linux glibc 2.28+ x86-64
SHA-256 checksum
How to use checksums
3de012d66f68e20daf067d30aa401b85635922cc2de227d92ca00055f1b1c2f6
BLAKE2b-256 checksum
How to use checksums
5fb5acb47af9c607e429ca1475e51aaddf26445c222d20638ffa5189fc0b4c30
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_cu12-0.12.0-cp310-cp310-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl

Download URL cuequivariance_ops_torch_cu12-0.12.0-cp310-cp310-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl
Size 1.0 MB
Tags CPython 3.10 Linux glibc 2.26+ ARM64 Linux glibc 2.28+ ARM64
SHA-256 checksum
How to use checksums
d0702719570ba2163f76d4de3e7717f7978c8c7b0016ba309d8519591ecc83cc
BLAKE2b-256 checksum
How to use checksums
79e9e7a2a36ce162a845413057693567a6c6b022558606a76522488fb8ab47ab
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/7.0.0 CPython/3.14.4

Release history Release notifications | RSS feed

This release

0.12.0 This release

10 release files

0.9.1

10 release files

0.9.0

10 release files

0.8.1

8 release files

0.8.0

8 release files

0.7.0

8 release files

0.6.1

6 release files

0.6.0

6 release files

0.5.1

4 release files

0.5.0

4 release files

0.4.0

4 release files

0.3.0

4 release files

0.2.0

3 release files

0.1.0

3 release files

0.0.0

3 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page