Buckets:

hailstone-harmony's picture
|
download
raw
9.24 kB
# Stable Audio 3 — TensorRT engines
Prebuilt TensorRT engines for the Stable Audio 3 stack, organised by CUDA architecture.
```
tensorRT/sm_90/ Hopper — H100, H200
tensorRT/sm_120/ Blackwell — RTX 5090, RTX PRO 6000, RTX PRO 4500
```
TensorRT engines are **not portable across architectures** — an `sm_90` engine will not
load on an `sm_120` GPU. They are also tied to the TensorRT version they were built with
(these were built with **TensorRT 10.15.1.29**). If your GPU or TensorRT version is not
covered here, build from the ONNX in `onnx/` using the scripts in
[`optimized/tensorRT/build/`](https://github.com/Stability-AI/stable-audio-3/tree/main/optimized/tensorRT/build).
| directory | contents |
|---|---|
| `sa3-m/` | medium DiT (the main diffusion transformer) |
| `sa3-sm-music/`, `sa3-sm-sfx/` | small domain-specialised DiTs |
| `same-l/`, `same-s/` | autoencoder (SAME) encoder + decoder |
| `t5gemma/` | text conditioner, plus its tokenizer |
---
## The SAME-L attention kernel: two implementations
The SAME-L encoder and decoder use a custom sliding-window-attention plugin
(`samel::diff_attn_swa`) for differential attention, which TensorRT has no native op for.
That plugin ships **two implementations**, and *which one an engine uses is fixed when the
engine is built* — it is baked into the `.trt` file and cannot be changed at load time.
**AOT** (ahead-of-time) compiles a block-tiled tensor-core kernel to PTX and embeds it in
the engine. No Python runs during inference. **JIT** dispatches a Triton kernel through a
Python callback on every enqueue — 12 callbacks per decode, one per attention layer.
| | AOT | JIT |
|---|---|---|
| CUDA-graph capturable | **yes, both architectures** | sm_90 only — **fails on sm_120** |
| Python / Triton calls per decode | **0** | 12 |
| Triton needed at inference | no | **yes** |
| Engine size | +0.1% | baseline |
### Which files use which
| published engines | kernel | why |
|---|---|---|
| `sm_120/same-l/` | **AOT** (block-tiled MMA) | required — the JIT build is not capturable on sm_120 |
| `sm_90/same-l/` | **JIT** | works, and AOT measured no faster on Hopper, so these were left alone |
So on **sm_90 you need Triton installed to run** the `same-l` engines, and on **sm_120 you
do not**. Both implementations stay registered by the runtime, so either kind of engine
loads and runs regardless of which GPU you are on.
For anything you build yourself, **AOT is the default on every architecture** — it is the
only correct choice on sm_120, it is at parity on sm_90, and it drops the runtime Triton
dependency. Rebuilding the `sm_90` engines would therefore give you AOT rather than the
JIT files published here; that is expected and fine, not a mismatch.
### Why JIT breaks on sm_120
The runtime captures the whole pipeline into one CUDA graph. A JIT engine cannot be
captured on sm_120: `enqueueV3` returns `False`, TensorRT logs *"this TRT engine is not
stream capturable"*, and **the decode stage is silently omitted from the graph** — no
exception, no non-zero exit.
What you get is the pre-capture warm-up decode of zero latents: a constant wash of noise,
**byte-identical for every prompt and every seed**, with exit code 0. If you ever see
that, this is the first thing to check. The diagnostic is to render twice with different
seeds and compare the files — identical bytes mean the decode never ran.
The same engine captures without complaint on sm_90, which is why the problem did not
surface until Blackwell. Re-entering Python inside a captured region was always fragile
rather than supported, so do not assume a future architecture will tolerate it either.
The decode is inside the captured region on the default path: `sa3_trt.py` captures T5 →
DiT loop → decoder → PCM copy as one graph whenever `--cfg 1.0` and no inpaint/init-audio
is used, which is the default. `--no-mega-graph` runs eagerly and sidesteps capture
entirely, so it is a usable workaround on an affected engine — at the cost of the
graph-replay speedup.
> **The `sm_120` SAME-L engines were rebuilt AOT on 2026-07-31.** Anything published
> before that date was a JIT build and is affected, so if you have an `sm_120` copy cached
> from earlier, re-pull it. The `sm_90` files are unchanged — JIT captures correctly on
> Hopper, so they were never affected.
### Benchmarks
SAME-L decode, milliseconds, median of 7 on an otherwise idle GPU. `L` is the latent
sequence length — one latent is 4096 samples, so at 44.1 kHz L=1292 is a two-minute
render and L=4096 is about 6m20s.
**sm_120 · RTX PRO 4500 Blackwell**
| kernel | L=256 | L=1024 | L=1292 | L=2048 | L=4096 |
|---|---|---|---|---|---|
| JIT *(not capturable here)* | 55.1 | 227.2 | 289.1 | 457.7 | 920.2 |
| AOT `ptx` (scalar fallback) | 62.5 | 251.2 | 318.7 | 502.8 | 1008.8 |
| **AOT `mma` (default)** | 55.4 | **223.6** | **283.7** | **448.4** | **899.3** |
**sm_90 · H200**
| kernel | L=256 | L=1024 | L=1292 | L=2048 | L=4096 |
|---|---|---|---|---|---|
| JIT | 12.4 | 47.0 | 61.1 | 98.5 | 194.4 |
| AOT `ptx` (scalar fallback) | 12.7 | 47.3 | 60.7 | 96.6 | 195.5 |
| **AOT `mma` (default)** | 12.4 | 48.4 | 62.6 | 100.0 | 196.8 |
The default AOT kernel is at parity with JIT on both architectures — within a few percent
either way, and faster at long sequence lengths on sm_120 — while remaining capturable.
Accuracy against the FP32 decoder at L=1292: AOT `mma` **51.45 dB** PSNR, JIT 51.43, AOT
`ptx` 51.11. The default kernel is the most accurate of the three as well as the fastest.
Note the ~4.5× gap between the two architectures. That is silicon, not a software
regression — a workstation card against a flagship datacenter part. Measured on these two
GPUs, the H200 has **5.57×** the copy bandwidth (4084 vs 734 GB/s) and 5.08× the BF16
tensor throughput, so a bandwidth-bound decoder landing at 4.5× is extracting slightly
*more* of its silicon on sm_120 than the H200 does of its own.
### How the kernels differ
The default AOT kernel (`mma`) is block-tiled: 16 queries per block share a 64-wide K/V
tile in shared memory, and both attention products go through `tl.dot`, which lowers to
512 × `mma.sync.aligned.m16n8k8.row.col.f32.tf32.tf32.f32` — TF32 tensor cores. It is
written in Triton and compiled ahead of time, so Triton generates the fragment layouts
rather than us hand-writing them, but the result is a plain PTX blob inside the engine.
The fallback AOT kernel (`ptx`) is hand-written scalar FP32: one warp per query, no K/V
reuse, zero `mma` instructions. It exists only for targets where the block-tiled kernel
will not build, and costs about 10% on sm_120.
TF32 in the two attention products matches what the JIT path already did — it is not a new
precision loss. The differential subtraction stays in FP32 on FP32 accumulators, so the
cancellation-sensitive step is untouched.
### Building your own
```bash
SA3_SWA_PLUGIN=aot # default — PTX compiled into the engine
SA3_SWA_PLUGIN=jit # Triton through a Python callback per enqueue
SA3_SWA_AOT=mma # default AOT kernel — block-tiled, tensor cores
SA3_SWA_AOT=ptx # fallback AOT kernel — scalar FP32
```
Triton is needed to *build* an AOT engine (it is the code generator) but not to run one.
Three gotchas, all of which produce wrong output rather than a build failure:
- **Kernel parameter order is inputs → runtime scalars → outputs.** TensorRT's AOT
launcher passes the `extra` scalars *before* the output pointers. Declaring the output
ahead of them makes the kernel dereference a sequence length as an address; ours died
with an illegal memory access until the order was fixed.
- **`BLOCK_KV` must be ≥ `BLOCK_N + 2*WINDOW`.** Otherwise the K/V tile stops covering
every position the block's queries can attend to and attention contributions are
*silently dropped*. `BLOCK_N=32` with `BLOCK_KV=64` needs 66 and is therefore wrong,
even though it compiles and fits in shared memory.
- **Shared memory over 48 KB fails at enqueue, not at build.** `BLOCK_N=64` needs 64 KB;
the engine builds cleanly and then reports *"Failed to enqueue status -1"* and returns
zeros. `BLOCK_N=16` needs 40 KB, and is faster anyway — with a window of only ±17, a
wide K/V tile is mostly masked-out waste.
The engine filenames still say `triton_swa` because Triton remains the code generator;
only the compilation moved ahead of time.
---
## Other things worth knowing
**The medium DiT has three precisions.** `dit_fp16mixed.trt` is the default and the one to
use: FP16 attention core with FP32 RMSNorm and RoPE islands, ~4.3× faster than FP32 with
no audible quality cost. `dit_fp32.trt` is the reference. `dit_bf16.trt` is kept for
comparison — earlier builds of it degraded past about two minutes of audio because the
RoPE angle was computed in BF16, where one ULP exceeds 2π at long sequence lengths; the
published engine has the RoPE table baked in at FP32 precision.
**`dit_fp8.trt` is sm_90 only** and must be built weakly-typed; strong typing breaks
TensorRT's attention fusion. This is the *opposite* of the FP16-mixed rule, which requires
`STRONGLY_TYPED` or the builder re-casts the FP32 islands.

Xet Storage Details

Size:
9.24 kB
·
Xet hash:
6d3391635ff8b60bd00ef1b355f4ad7d8419c105c791be9dfba010262da6d869

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.