Metadata-Version: 2.5
Name: fast_trimul
Version: 0.0.7
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: MIT
License-File: LICENSE
License-File: NOTICE
Keywords: alphafold,cute,cutlass,gpu-kernel,openfold,triangle-multiplication
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

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

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

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

Low-level functional API (FlashAttention style):

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

## Colab / Jupyter quickstart (with an event-based timer)

Install:

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

Run it and time it. The timer uses **CUDA events + `synchronize()`**, so it measures
when the GPU actually *finishes the work* — not when the launch is queued:

```python
import time, torch
from fast_trimul import FastTriangleMultiplication

assert torch.cuda.is_available(), "Need a CUDA GPU (Colab: Runtime -> Change runtime type -> GPU)."
print("GPU:", torch.cuda.get_device_name(0))

B, N, d_z, d_c = 1, 256, 128, 128
module = FastTriangleMultiplication(d_z=d_z, d_c=d_c, mode="outgoing").cuda()
z    = torch.randn(B, N, N, d_z, device="cuda")     # (B, N, N, d_z)
mask = torch.ones(B, N, N, device="cuda")           # optional (B, N, N)
print(f"input : {tuple(z.shape)}  {z.dtype}")

# first call includes the one-time CuTe JIT compile (wall clock is fine here)
t0 = time.perf_counter()
with torch.no_grad():
    out = module(z, mask=mask)
torch.cuda.synchronize()
print(f"first call (incl. JIT compile): {time.perf_counter()-t0:5.2f} s")
print(f"output: {tuple(out.shape)}  {out.dtype}   mean={out.mean():.4f}  std={out.std():.4f}")

def bench(fn, iters=50, warmup=10):
    for _ in range(warmup):                 # warmup: compiled + caches hot
        fn()
    torch.cuda.synchronize()
    start = torch.cuda.Event(enable_timing=True)
    end   = torch.cuda.Event(enable_timing=True)
    start.record()
    for _ in range(iters):
        fn()
    end.record()
    torch.cuda.synchronize()                # read a COMPLETED timestamp, not a queued one
    return start.elapsed_time(end) / iters  # ms per call (GPU timeline)

with torch.no_grad():
    ms = bench(lambda: module(z, mask=mask))
elems = z.numel()
print(f"\nsteady-state (CUDA events):")
print(f"  {ms*1e3:8.1f} us / call")
print(f"  {elems/1e6:6.1f}M elements  ->  {elems/(ms/1e3)/1e9:6.2f} Gelem/s")
```

## Benchmark: with vs without `torch.compile`

Times the **same op** three ways — `fast_trimul`, plain torch (**without** compile),
and **with** `torch.compile` — using the same weights and the event-based timer:

```python
import torch
from fast_trimul import FastTriangleMultiplication
from fast_trimul._kernels import TriangleMultiplicativeUpdate   # torch reference (same op)

torch.manual_seed(0)
torch.set_float32_matmul_precision("high")   # let torch use TF32 tensor cores
B, N, d_z, d_c, mode = 1, 256, 128, 128, "outgoing"

ref  = TriangleMultiplicativeUpdate(d_z, d_c, mode).cuda().eval()   # torch, fp32
fast = FastTriangleMultiplication(d_z, d_c, mode).cuda()
fast._impl.load_state_dict(ref.state_dict(), strict=False)          # same weights

z = torch.randn(B, N, N, d_z, device="cuda")

def bench(fn, iters=50, warmup=10):
    for _ in range(warmup): fn()
    torch.cuda.synchronize()
    s = torch.cuda.Event(enable_timing=True); e = torch.cuda.Event(enable_timing=True)
    s.record()
    for _ in range(iters): fn()
    e.record(); torch.cuda.synchronize()
    return s.elapsed_time(e) / iters

with torch.no_grad():
    err = (fast(z).float() - ref(z)).abs().max().item()

ref_compiled = torch.compile(ref)
with torch.no_grad():
    t_fast  = bench(lambda: fast(z))            # fast_trimul (fp16 kernels)
    t_eager = bench(lambda: ref(z))             # torch  WITHOUT compile (fp32)
    t_comp  = bench(lambda: ref_compiled(z))    # torch  WITH compile   (fp32)

print(f"max|fast - torch| = {err:.2e}   (fp16 vs fp32 -> fp16 rounding, not a bug)\n")
for name, ms in [("fast_trimul (fp16)",   t_fast),
                 ("torch eager (fp32)",   t_eager),
                 ("torch.compile (fp32)", t_comp)]:
    print(f"  {name:<22} {ms*1e3:8.1f} us/iter")
```

**Read the result honestly:** `fast_trimul` is fp16 while the torch baselines are
fp32, and `torch.compile` typically **wins above small N today** — the kernels are
not yet epilogue-fused (see *Limitations*). The point of this cell is to *measure*,
not to assume. For the fully fair fp16 comparison, run the torch reference with
`.half()` and `torch.compile(..., mode="reduce-overhead")`.

## Full benchmark (machine ceilings + roofline)

To **compare correctly after installing the library**, the package ships a
rigorous benchmark — measured machine ceilings (memory bandwidth, fp16 tensor-core
peak, launch floor), a per-iteration **median** timer (median / min / p95 / CV,
not a mean), roofline placement (% of peak), effective GB/s, achieved TFLOP/s, and
a size sweep — for `fast_trimul` vs torch eager vs `torch.compile`:

```python
from fast_trimul.benchmark import main
main()                       # or: main(head_size=384, sweep=(128, 256, 512))
```
or from a shell:
```bash
python -m fast_trimul.benchmark
```

It prints something like:

```text
GPU: NVIDIA A100-SXM4-40GB
  measured mem bandwidth peak :     1490 GB/s
  measured fp16 matmul peak   :      270 TFLOP/s
  launch-overhead floor       :      4.6 us

Head-to-head  N=256, d_z=128, d_c=128   (17.2 GFLOP/call, fp16 err vs torch = 3.8e-03)
  impl                     median     min    CV%  TFLOP/s    GB/s  %peak  vs eager
  fast_trimul (fp16)       ...
  torch eager (fp32)       ...
  torch.compile (fp32)     ...

Size sweep (median us/call, and fast_trimul TFLOP/s):
      N      fast     eager   compile  fast TFLOP/s
     64      ...
```

(Numbers are illustrative — run it on your GPU. It's fp16 kernel vs fp32
baselines; the script prints that caveat and points to the fully-fair fp16 run.)

## 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
```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")` — 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](NOTICE).
