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-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)

Table of built distributions (wheels) for cuequivariance-ops-torch-cu13 0.12.0
File
cuequivariance_ops_torch_cu13-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_cu13-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_cu13-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.27+ x86-64, Linux glibc 2.28+ x86-64 Details
cuequivariance_ops_torch_cu13-0.12.0-cp313-cp313-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.13 CPython 3.13 Linux glibc 2.28+ ARM64, Linux glibc 2.26+ ARM64 Details
cuequivariance_ops_torch_cu13-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_cu13-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_cu13-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_cu13-0.12.0-cp311-cp311-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl CPython 3.11 CPython 3.11 Linux glibc 2.28+ ARM64, Linux glibc 2.26+ ARM64 Details
cuequivariance_ops_torch_cu13-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_cu13-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: 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

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

4 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