Metadata-Version: 2.5
Name: fast_trimul
Version: 1.0.0
Summary: Fused Triangle Multiplicative Update (AlphaFold/OpenFold) on CUTLASS CuTe DSL kernels.
Project-URL: Homepage, https://github.com/tiagomonteiro0715/fast_trimul
Project-URL: Repository, https://github.com/tiagomonteiro0715/fast_trimul
Project-URL: Issues, https://github.com/tiagomonteiro0715/fast_trimul/issues
Author: Tiago Monteiro
License: Apache-2.0
License-File: LICENSE
License-File: NOTICE
Keywords: alphafold,cute,cutlass,gpu-kernel,openfold,triangle-multiplication
Classifier: License :: OSI Approved :: Apache Software License
Classifier: Programming Language :: Python :: 3
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
Requires-Python: >=3.10
Requires-Dist: cuda-python
Requires-Dist: nvidia-cutlass-dsl
Requires-Dist: torch>=2.2
Provides-Extra: test
Requires-Dist: pytest; extra == 'test'
Description-Content-Type: text/markdown

# fast_trimul

[![PyPI](https://img.shields.io/pypi/v/fast_trimul)](https://pypi.org/project/fast_trimul/)
[![License: Apache 2.0](https://img.shields.io/badge/License-Apache_2.0-blue.svg)](LICENSE)
[![Python](https://img.shields.io/pypi/pyversions/fast_trimul)](https://pypi.org/project/fast_trimul/)

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.
- **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

```bash
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:

```python
!pip install -q uv
!uv pip install fast_trimul
```

Then use it:

```python
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:

```python
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):

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

Load **pretrained weights** from a target library (parameter names are remapped for you):

```python
module.load_openfold_state_dict(ref.state_dict())    # OpenFold / AF2  (separate a/b projections)
module.load_openfold3_state_dict(ref.state_dict())   # OpenFold-3      (separate OR fused variant)
module.load_protenix_state_dict(ref.state_dict())    # Protenix        (OpenFold-style names, bias-free)
module.load_boltz_state_dict(ref.state_dict())       # Boltz-1 / Chai / AF3 (fused p_in/g_in, split for you)
```

These target modules apply their residual (`+ z`) **outside** the triangle block,
so build with `residual=False` when matching their output exactly:

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

## 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:

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

Or from a shell:
```bash
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()`),
* **`compile`** — `torch.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.

## Drop-in monkeypatch for the 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 note.

### OpenFold
```python
import fast_trimul.integrations as fti
fti.patch_openfold()          # patches Outgoing + Incoming
# ... now build your OpenFold model as usual ...
```
Equivalent manual form:
```python
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
```python
import fast_trimul.integrations as fti
fti.patch_boltz()
```
Manual form:
```python
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
```python
import fast_trimul.integrations as fti
fti.patch_protenix()
```
Manual form:
```python
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):
```python
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", residual=True)`
  — high-level module, `forward(z, mask=None)`, `.graphed(z, mask=None)`; weight
  loaders `.load_openfold_state_dict` / `.load_openfold3_state_dict` /
  `.load_protenix_state_dict` / `.load_boltz_state_dict`.
- `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)

- **`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: `load_openfold_state_dict` (OpenFold/AF2),
  `load_openfold3_state_dict` (OpenFold-3, separate or fused variant),
  `load_protenix_state_dict` (Protenix), and `load_boltz_state_dict` (Boltz-1 / Chai /
  AF3, which fuse the a/b projections). Other stacks: patch-then-train, or supply a
  parameter 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.
- **Ampere (sm80) tested.** Hopper/Blackwell + fp8 are future work.
- **`import fast_trimul` needs a CUDA GPU** (device properties are read at import).

## License

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