Skip to main content

fast_trimul

PyPI License: Apache 2.0 Python

Fused Triangle Multiplicative Update (AlphaFold2 / AlphaFold3 family) built on hand-written CUTLASS CuTe DSL kernels — a drop-in nn.Module for the structural-biology stacks (OpenFold, OpenFold-3, Boltz, Chai, Protenix).

  • Numerically matches the stock module (fp16 tolerance) — verified against OpenFold, OpenFold-3, Boltz-1, Protenix, and an AF3/Chai-style reference by loading their weights and comparing outputs.
  • Roughly halves peak memory versus the stock eager module (kernel fusion + CUDA-graph buffer reuse).
  • Fastest at small N, where per-launch overhead dominates and the captured CUDA graph removes it.
  • Modular, vendor-agnostic backends with a fallback chain (cuda → torch): the fast CUTLASS kernel when it fits, an always-correct pure-torch path otherwise, so an unsupported shape/dtype degrades gracefully instead of crashing. New hardware (TPU/Intel/…) is a plug-in, not a rewrite.
  • Drop-in on any shape, no whole-model compilation.

Run the shipped benchmark on your own GPU for numbers — see Benchmark below, and read Limitations for where torch.compile is the better choice.

Install

pip install fast_trimul          # or: uv pip install fast_trimul

Requires a CUDA GPU, torch, nvidia-cutlass-dsl, and cuda-python. Kernels JIT-compile on first use (one-time cost, then cached in-process).

Quick start

On Google Colab (Runtime → Change runtime type → GPU), install first:

!pip install -q uv
!uv pip install fast_trimul

Then use it:

import torch
from fast_trimul import FastTriangleMultiplication

module = FastTriangleMultiplication(d_z=128, d_c=128, mode="outgoing").cuda()
z = torch.randn(1, 256, 256, 128, device="cuda")          # (B, N, N, d_z)
mask = torch.ones(1, 256, 256, device="cuda")             # optional (B, N, N)
out = module(z, mask=mask)                                 # same dtype as z

For fastest inference at a fixed shape, capture a CUDA graph once — this removes the per-launch overhead of the internal kernels, which dominates the runtime at small N:

module.graphed(z, mask)      # capture once at this shape (inference only)
out = module(z, mask=mask)   # subsequent calls replay the graph

Low-level functional API (FlashAttention style):

from fast_trimul import functional
out = functional.triangle_multiplication(z, module._impl, mask=mask)

Load pretrained weights from a target library — one call, pick the source (names are remapped, and fused a/b projections are split, for you):

module.load_weights(ref.state_dict(), source="openfold")   # or: openfold3 / protenix / boltz / chai

The five named helpers (load_openfold_state_dict, …) still work as thin aliases. These target modules apply their residual (+ z) outside the triangle block, so build with residual=False when matching their output exactly:

module = FastTriangleMultiplication(d_z=128, d_c=128, mode="outgoing", residual=False).cuda()

Pick or force a backend (default is "auto" — fastest available, then torch):

from fast_trimul import list_backends
list_backends()                                            # e.g. ['torch', 'cuda']
FastTriangleMultiplication(d_z=128, backend="cuda")        # force the fast path (still falls back)
FastTriangleMultiplication(d_z=128, backend="torch")       # force the portable reference

Opt in explicitly, without any global monkeypatching (safe under strict runtime policies) — @accelerate on a function, or the scoped accelerated():

from fast_trimul import accelerate, accelerated

@accelerate                    # runs this function with the accelerated backend
def infer(z): ...

with accelerated("cuda"):      # scoped: reverts on exit, never leaks
    out = model(z)

One-line correctness check — build your library's TriMul, then:

import fast_trimul
fast_trimul.verify("openfold", my_openfold_trimul)         # -> True if outputs match (fp16 tol)

verify stays a one-liner even when a library changes its API, because you pass the reference module and only the (registered) name-remap is library-specific.

Benchmark

The package ships a benchmark that measures machine ceilings (memory bandwidth, fp16 tensor-core peak, launch floor), a per-iteration median timer, achieved TFLOP/s, peak memory, and a size sweep. It reports fast_trimul both un-graphed and graphed, next to torch.compile and an eager reference, so you can compare on your own hardware:

!pip install -q uv
!uv pip install fast_trimul
from fast_trimul.benchmark import run_benchmark
run_benchmark()              # or: run_benchmark(head_size=384, sweep=(128, 256, 512))

Or from a shell:

python -m fast_trimul.benchmark

It reports these variants:

  • fast no-graph — the kernel, fp16, un-graphed (shows the launch-overhead cost),
  • fast +graph — the same kernel with a captured CUDA graph (.graphed()),
  • compiletorch.compile(mode="reduce-overhead") and default mode,
  • torch eager — the eager reference.

Use CUDA events + synchronize() (as the shipped benchmark does) so timing reflects when the GPU finishes the work, not when the launch is queued. Warm up (or call .graphed()) before timing to exclude the one-time JIT/autotune cost.

Architecture

The library is a small stable front-end over a registry of interchangeable backends. CUDA is one backend; a pure-torch backend is the universal fallback. Adding new hardware or a new library is a plug-in (one decorated class/function), not a core edit.

  front-end   @accelerate  accelerated()  verify()          # stable, tiny
      │
  guard        contiguous · dtype · align · int64 strides   # NormalizedInput
      │
  dispatch     pick backend, then FALL BACK: cuda -> torch  # never crashes on a bad shape
      │
  backends     cuda_cute (CUTLASS)   torch_ref (portable)   # + future: xla/tpu, xpu/intel
Module Role
core/registry.py @backend / @weights_for decorators + O(1) lookup tables
core/context.py the input guardNormalizedInput (contiguous, dtype, alignment, int64 strides)
core/dispatch.py picks a backend and walks the fallback chain (cuda → torch)
core/decorators.py @accelerate + accelerated() — explicit, scoped, no global patching
backends/torch_ref.py universal pure-torch backend (any device torch supports)
backends/cuda_cute.py thin wrapper over the existing CUTLASS kernels (unchanged)
ops/triangle.py FastTriangleMultiplication (dispatch + graph capture + load_weights)
integrations/loaders.py the five per-library weight maps (@weights_for)
integrations/checks.py verify()

Everything the architecture adds — registries, guard, dispatch, decorators — is O(1) overhead; only the op itself scales with N. The CUDA kernels (_kernels.py) are untouched.

Adding a new backend

from fast_trimul.core.registry import backend

@backend("mybackend", dtypes={torch.float16}, min_align=8)
class MyBackend:
    def __init__(self, caps): self.caps = caps
    def execute(self, inp, params): ...   # inp.tensor is guarded (contiguous, right dtype)

That one file makes backend="mybackend" selectable and slots it into the fallback chain — no changes to the dispatcher or the module.

API

  • fast_trimul.FastTriangleMultiplication(d_z, d_c=None, mode="outgoing", residual=True, backend="auto") — the module. forward(z, mask=None), .graphed(z, mask=None), .load_weights(state_dict, source=...) (plus the named aliases .load_openfold_state_dict / .load_openfold3_state_dict / .load_protenix_state_dict / .load_boltz_state_dict).
  • fast_trimul.accelerate / fast_trimul.accelerated(backend="auto") — explicit opt-in decorator / scoped context manager.
  • fast_trimul.verify(source, reference, ...) — one-line correctness check against a reference module.
  • fast_trimul.list_backends() — backends registered on this machine.
  • fast_trimul.functional.triangle_multiplication(z, params, mask=None) — low-level functional call.
  • fast_trimul.core.registry.{backend, weights_for} — decorators to register a new backend or library weight-map.

Limitations (read before relying on it)

  • torch.compile(mode="reduce-overhead") is competitive and often faster above small N. On an A100 it is frequently faster per call in the mid-range and, on several stacks, uses similar peak memory. These kernels are not yet epilogue-fused (future work), so the reasons to prefer this are drop-in-ness and robustness, not raw latency: reduce-overhead needs static shapes and recompiles per sequence length (awkward for variable-length inputs) and can break on some models, whereas this is a plain nn.Module that works on any shape with no compilation step. Benchmark both on your workload.
  • First call is slow: JIT compile + GEMM autotune. On the first forward at a new shape, the GEMM configs are auto-tuned (one-time, cached). Disable with the env var FAST_TRIMUL_AUTOTUNE=0. Warm up (or call .graphed()) before timing.
  • fp16 only. bf16/fp32 inputs are cast to fp16 and back; keep the module in fp16 (do not call .float()/.bfloat16() on it).
  • Pretrained weights need name remapping. Each library names its projections/norms differently, so a strict checkpoint load will not line up. Automated for the common stacks via load_weights(sd, source=...): openfold (OpenFold/AF2), openfold3 (separate or fused variant), protenix, boltz, and chai (Boltz/Chai/AF3 fuse the a/b projections). Other stacks: supply a @weights_for remap.
  • Mask semantics are approximate. The mask is applied to the pair tensor in and out; validate against each library's exact masking before production use.
  • Backward is correct but not fast (torch recompute), so it helps inference more than training throughput.
  • The CUDA kernel is Ampere (sm80) tested; fp16 only. On other hardware, or for bf16/fp32, or non-multiple-of-8 N, the dispatcher falls back to the pure-torch backend (correct, slower). Hopper/Blackwell + fp8 are future work.
  • Building a module needs a CUDA GPU + CUTLASS (the fused kernel weights live in a CUTLASS module). import fast_trimul itself is lazy and does not require CUTLASS until you construct FastTriangleMultiplication.

License

Apache License 2.0 (this project) — see LICENSE. The GEMM core is derived from NVIDIA CUTLASS and is licensed under BSD 3-Clause — see NOTICE.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

fast_trimul-2.0.0.tar.gz (94.6 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

fast_trimul-2.0.0-py3-none-any.whl (49.2 kB view details)

Uploaded Python 3

File details

Details for the file fast_trimul-2.0.0.tar.gz.

File metadata

  • Download URL: fast_trimul-2.0.0.tar.gz
  • Upload date:
  • Size: 94.6 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.9.18 {"installer":{"name":"uv","version":"0.9.18","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":null,"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

File hashes

Hashes for fast_trimul-2.0.0.tar.gz
Algorithm Hash digest
SHA256 c580968753347cc5ad57eb23f406bf216340881570740148d41e3822a8cde66b
MD5 a441da9c37da3d7043170190f45ebc5a
BLAKE2b-256 c87ea16a458baf08d0574607a29288f7c2bc215cf8b44dcab2a300afe1413835

See more details on using hashes here.

File details

Details for the file fast_trimul-2.0.0-py3-none-any.whl.

File metadata

  • Download URL: fast_trimul-2.0.0-py3-none-any.whl
  • Upload date:
  • Size: 49.2 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.9.18 {"installer":{"name":"uv","version":"0.9.18","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":null,"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

File hashes

Hashes for fast_trimul-2.0.0-py3-none-any.whl
Algorithm Hash digest
SHA256 c97fd007c9568ed2ced26b2daf6b174eebae06d004857e469acf159737640607
MD5 4132d1b92e206a8515ef44ce906aba6f
BLAKE2b-256 22b7c26022d82ba1463bcb203666ed2bb08f77854454ed240fe589b32e35cf90

See more details on using hashes here.

Release history Release notifications | RSS feed

3.0.4

2 files

3.0.3

2 files

3.0.2

2 files

3.0.1

2 files

3.0.0

2 files

2.4.34

2 files

2.4.33

2 files

2.4.32

2 files

2.4.31

2 files

2.4.30

2 files

2.4.29

2 files

2.4.28

2 files

2.4.27

2 files

2.4.26

2 files

2.4.25

2 files

2.4.24

2 files

2.4.23

2 files

2.4.22

2 files

2.4.21

2 files

2.4.20

2 files

2.4.19

2 files

2.4.18

2 files

2.4.17

2 files

2.4.16

2 files

2.4.15

2 files

2.4.14

2 files

2.4.13

2 files

2.4.12

2 files

2.4.11

2 files

2.4.10

2 files

2.4.9

2 files

2.4.8

2 files

2.4.7

2 files

2.4.6

2 files

2.4.5

2 files

2.4.4

2 files

2.4.3

2 files

2.4.2

2 files

2.4.1

2 files

2.4.0

2 files

2.3.3

2 files

2.3.2

2 files

2.3.1

2 files

2.3.0

2 files

2.2.2

2 files

2.2.1

2 files

2.2.0

2 files

2.1.4

2 files

2.1.3

2 files

2.1.2

2 files

2.1.1

2 files

2.0.1

2 files

This release

2.0.0 This release

2 files

1.0.0

2 files

0.0.30

2 files

0.0.29

2 files

0.0.28

2 files

0.0.27

2 files

0.0.26

2 files

0.0.25

2 files

0.0.24

2 files

0.0.22

2 files

0.0.21

2 files

0.0.20

2 files

0.0.16

2 files

0.0.15

2 files

0.0.14

2 files

0.0.13

2 files

0.0.12

2 files

0.0.11

2 files

0.0.10

2 files

0.0.9

2 files

0.0.8

2 files

0.0.7

2 files

0.0.6

2 files

0.0.5

2 files

0.0.4

2 files

0.0.1

2 files

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page