Skip to main content

fast_trimul

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

Honest status. The kernels are fp16 and numerically correct (they match PyTorch fp16 to fp16 tolerance). On a fair comparison (torch.compile(..., mode="reduce-overhead") in fp16) they are slower than torch.compile above small N today — the GEMMs are not yet epilogue-fused. The wins are: correctness, a drop-in API, and (with full fusion, future work) lower memory. Full GEMM epilogue fusion and a FlashAttention-style megakernel are future work — see Limitations.

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

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

Low-level functional API (FlashAttention style):

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

Drop-in monkeypatch for the 4 target libraries

Each helper replaces the library's TriMul class with an adapter matching its constructor. Patch before building the model. See Limitations for the pretrained-weight caveat.

OpenFold

import fast_trimul.integrations as fti
fti.patch_openfold()          # patches Outgoing + Incoming
# ... now build your OpenFold model as usual ...

Equivalent manual form:

import openfold.model.triangular_multiplicative_update as of_tri
from fast_trimul.integrations import adapter
of_tri.TriangleMultiplicationOutgoing = adapter("outgoing")
of_tri.TriangleMultiplicationIncoming = adapter("incoming")

Boltz-1 / BoltzDesign

import fast_trimul.integrations as fti
fti.patch_boltz()

Manual form:

import boltz.model.layers.triangular_mult as b_tri
from fast_trimul.integrations import adapter
b_tri.TriangleMultiplicationOutgoing = adapter("outgoing")
b_tri.TriangleMultiplicationIncoming = adapter("incoming")

Protenix

import fast_trimul.integrations as fti
fti.patch_protenix()

Manual form:

import protenix.model.modules.pairformer as p_tri
from fast_trimul.integrations import adapter
p_tri.TriangleMultiplication = adapter("outgoing")

Chai-1

Chai's module path is version-dependent, so patch the attribute explicitly (replace the import path with the one in your installed version):

from fast_trimul.integrations import adapter
import chai_lab.model.<...>.triangle_mult as c_tri   # <- verify path for your version
c_tri.TriangleMultiplicationOutgoing = adapter("outgoing")
c_tri.TriangleMultiplicationIncoming = adapter("incoming")

API

  • fast_trimul.nn.FastTriangleMultiplication(d_z, d_c=None, mode="outgoing") — high-level module, forward(z, mask=None).
  • fast_trimul.functional.triangle_multiplication(z, params, mask=None) — low-level functional call.
  • fast_trimul.integrations.{patch_openfold, patch_boltz, patch_protenix, adapter} — monkeypatch helpers.

Limitations (read before relying on it)

  • Slower than torch.compile (fp16) above small N. Correctness and drop-in compatibility come first; speed parity needs the epilogue fusion / megakernel (future work).
  • 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. Patch-then-train, or supply a parameter remap. Loading pretrained checkpoints is not yet automated.
  • 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.
  • Ampere (sm80) tested. Hopper/Blackwell + fp8 are future work.
  • import fast_trimul needs a CUDA GPU (device properties are read at import).

License

MIT (this project). 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-0.0.5.tar.gz (78.5 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-0.0.5-py3-none-any.whl (28.6 kB view details)

Uploaded Python 3

File details

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

File metadata

  • Download URL: fast_trimul-0.0.5.tar.gz
  • Upload date:
  • Size: 78.5 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-0.0.5.tar.gz
Algorithm Hash digest
SHA256 86e0d5d69b01763b01760ea2c3cef511be20c2c17303dc14ccd471455dca818d
MD5 9f17e33828ea7906cc8d634d55522d68
BLAKE2b-256 0e747062434b90c542216c6df11331f97aab5b57e230705eb398894ae0d05971

See more details on using hashes here.

File details

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

File metadata

  • Download URL: fast_trimul-0.0.5-py3-none-any.whl
  • Upload date:
  • Size: 28.6 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-0.0.5-py3-none-any.whl
Algorithm Hash digest
SHA256 d4a139d1b775d77dfe33dbb453c06ab22a40686a6f0511743c614ddd4affadc7
MD5 71bb2fe34ee8cc01efc59c7b5e7abba7
BLAKE2b-256 a3ca60a72c1a1d9fe53a13e36b04e9021b07b34c70026c406a57e8f7fc31fcc9

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

2.0.0

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

This release

0.0.5 This release

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