Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions optimized/cpu-amx/.gitignore
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
output/
.venv/

# The repo-root .gitignore has a blanket `build/` rule for generated output. Ours holds the
# ENGINE SOURCES (C++ + weight dumpers), so re-include them explicitly...
!build/
!build/**

# ...but never track bytecode. These rules come AFTER the re-include so they still win.
__pycache__/
**/__pycache__/
*.pyc
125 changes: 125 additions & 0 deletions optimized/cpu-amx/BUILD.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,125 @@
# Building the cpu-amx engines

The `cpu-amx` runtime loads six torch-free C++/AMX engines as `.so`s via ctypes. This doc explains
how to (re)build each `.so` and dump its weight blob. The runtime itself (CLI/gradio) needs only the
built `.so`s + weight blobs; you only need this if you're rebuilding from source.

All sources live under `build/`. Prebuilt `.so`s + weight blobs live in per-engine directories under
`$SA3_CPUAMX_HOME` (see "Where the runtime looks" below) — `backends.py` resolves that same variable.

## Components

| engine | source (`build/…`) | GEMM path | weight blob | notes |
|---|---|---|---|---|
| **T5Gemma encoder** | `t5gemma/` | oneDNN AMX-BF16 | `weights.bin` (563 MB bf16) | 12-layer Gemma2 encoder |
| **medium DiT** (int8) | `dit/` | AOT Triton (int8) + oneDNN | `core_L{N}.bin` | **the complex one** — AOT-compiled Triton kernels |
| **SAME-S decoder** (bf16) | `same_s_bf16/` | oneDNN AMX-BF16 | `weights.bin` | fastest / distilled 50M |
| **SAME-L decoder** (bf16) | `same_l_bf16/` | oneDNN AMX-BF16 | `weights.bin` (852 MB) | native 426M |
| **SAME-S decoder** (int8) | `same_s_int8/` | oneDNN AMX-INT8 (fused) | `weights.bin` (55 MB) | ships SQ+GPTQ ("improved") grid |
| **SAME-L decoder** (int8) | `same_l_int8/` | oneDNN AMX-INT8 (fused) | `weights.bin` (427 MB) | ships SQ+GPTQ ("improved") grid |

## Prerequisites

1. **CPU with AMX** (Sapphire Rapids / Emerald Rapids Xeon; reports `cpu_isa_avx10_1_512_amx`).
The engines fall back to VNNI/AVX2 but the headline speed needs AMX (see `LESSONS.md`).
2. **Static oneDNN with the OpenMP runtime**, path in `$ONEDNN_HOME` (`include/`, `lib/libdnnl.a`).
Rebuild from oneDNN source with `-DDNNL_CPU_RUNTIME=OMP -DDNNL_LIBRARY_TYPE=STATIC` if it's missing.
3. **g++** with C++17 + `-march=native` (must resolve AMX intrinsics).
4. **triton-cpu fork**, path in `$SA3_TRITON_CPU` — ONLY for rebuilding the DiT's AOT kernels.
5. A Python environment with numpy (for weight dumping).

## Weight sourcing

All weights derive from the HF release **`stabilityai/stable-audio-3-optimized`** → the `MLX/*.npz`
files (fp16 named arrays). Download with `huggingface_hub.hf_hub_download`:
- `MLX/t5gemma_f16.npz` → T5Gemma
- `MLX/same_s_decoder_f32.npz`, `MLX/same_l_decoder_f32.npz` → decoders
- the medium DiT int8 weights come from the campaign's `aot_stage2/weights_int8_L{N}.npz` (int8-quantized
offline from `MLX/dit_medium_f16.npz`).

## Build recipes

### T5Gemma, and both bf16 / int8 decoders (the easy five)

All five are a single `g++` + static-oneDNN link. The pattern (see each dir's `build.sh`, or the build
comment at the top of the `.cpp`):

```bash
ONE="$ONEDNN_HOME"
g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I"$ONE/include" \
<engine>.cpp -o <engine>.so "$ONE/lib/libdnnl.a" -ldl -lpthread -lm
```

Then dump the weight blob from the HF npz (each dir has its dumper):

```bash
# T5Gemma
python build/t5gemma/dump_weights.py <t5gemma_f16.npz> # -> weights.bin + weights_manifest.txt
# bf16 decoders
python build/same_s_bf16/dump_weights.py <same_s_decoder_f32.npz>
python build/same_l_bf16/dump_weights.py <same_l_decoder_f32.npz>
# int8 decoders (naive grid)
python build/same_s_int8/dump_weights_int8.py <same_s_decoder_f32.npz>
python build/same_l_int8/dump_weights_int8.py <same_l_decoder_f32.npz>
```

The dumpers write `weights.bin` (bf16 or per-out-channel int8) + `weights_manifest.txt` (the mmap layout
the C++ loader reads). Linear weights are pre-transposed to `(in,out)` so the row-major oneDNN matmul
consumes them directly. **The int8 "improved" (SQ+GPTQ) blobs that actually ship** are produced by the
separate calibration tooling (SmoothQuant α0.9 → GPTQ → re-quant), not included here; they
are byte-compatible drop-ins for the int8 `.so` and are what the shipped `weights.bin` symlinks point to.

### medium DiT (int8) — the AOT-Triton path

The DiT does its int8 GEMMs + fused int8 attention through **AOT-compiled Triton-CPU kernels** (not plain
oneDNN), because the fully-fused all-integer path is the CPU win (see `LESSONS.md`). Two stages:

**1. Compile the per-ISA kernel set** (`build/dit/compile_isa_all.py`). Run once per ISA in its OWN
subprocess with the env the launcher sets — the ISA is NOT in Triton's cache key, so each ISA needs an
isolated cache dir:

```bash
TRITON_CPU_BACKEND=1 \
TRITON_CPU_TARGET_FEATURES=+amx-tile,+amx-int8,+amx-bf16,+avx512f,... \
TRITON_CPU_AOT_FORCE_ASM_FEATURES=1 \
TRITON_CACHE_DIR=/tmp/triton_cache_amx \
python build/dit/compile_isa_all.py amx <out_dir>
```

Produces `<out_dir>/{so/, cpp_kernels.txt, kernels_abi.json, so_flash/}` matching the dispatch keys the
C++ driver expects. Repeat for `vnni` / `avx2` if you want the fallbacks (AMX is the shipped path).
(Note: `gemm_i8` at BK=256 blows LLVM up on AVX2 — it's compiled at BK=64 there, int8→int32 is
tiling-invariant-exact; `NOGEMM=1` skips it.)

**2. Dump the weight blob** for the sequence length(s) you need (`build/dit/gen_core.py`):

```bash
python build/dit/gen_core.py 320 # -> core_L320.bin + core_L320_manifest.txt (L-independent block weights)
```

`core_L{N}.bin` flattens the int8 block weights + a golden preamble. The block weights are
sequence-length-independent, so one `core_L*.bin` serves any length (the `L` in the name is just the
golden shape).

**3. Build the C++ driver** (`build/dit/dit_cpu_amx.cpp`) — links static oneDNN like the others and
`dlopen`s the AOT kernels from `<out_dir>`. See the build comment at the top of the `.cpp`.

⚠ A full clean DiT rebuild needs the triton-cpu fork + the int8-quantized weight npzs
(`aot_stage2/weights_int8_L{N}.npz`). If you only need to run (not rebuild), use the prebuilt
`dit_cpu_amx.so` + `core_L320.bin`.

## Where the runtime looks

`scripts/backends.py` loads each engine's `.so` + weight blob from an engine dir under
`$SA3_CPUAMX_HOME` (`t5gemma_cpu_amx/`, `dit_medium_cpu_amx/`, `same_{s,l}_cpu_amx/`,
`same_{s,l}_int8fused_cpu_amx/`, `same_{s,l}_encoder_cpu_amx/`). **These binaries are NOT in git — `scripts/weights.py` `ensure()`
downloads them from HF** (`stabilityai/stable-audio-3-optimized/cpu-amx/`, flat) into those dirs on first
use; a local build (this BUILD.md) satisfies them too and skips the download. Override the base dir with
`SA3_CPUAMX_HOME`, or `SA3_CPUAMX_NO_HF=1` to force local-only. The DiT `.so` bakes in absolute kernel/core
paths are read from `$SA3_CPUAMX_HOME` at load time, so a rebuilt `.so` is relocatable; binaries published before that change still carry the old absolute prefix.
env-configurable (like the other seven engines, which load relative to their dir) is the remaining
portability follow-up.

See `LESSONS.md` for the numerics gotchas each of these engines encodes (fp32 islands, the bf16-RoPE
long-sequence bug, the global-static-weights trap, the DiT teardown quirk, …). Read it before touching
the `.cpp`s.
214 changes: 214 additions & 0 deletions optimized/cpu-amx/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,214 @@
# sa3 cpu-amx — Stable Audio 3 on CPU (torch-free C++ AMX)

CPU-native inference for **Stable Audio 3 medium**, running the whole pipeline on
**torch-free C++ AMX engines** (Intel AMX / AVX-512). No PyTorch, MLX, TFLite, or
stable-audio-tools at runtime for text-to-audio — just the C++ `.so`s + numpy.

```
prompt ─▶ T5Gemma (C++ AMX) ─▶ numpy conditioner ─▶ DiT pingpong (C++ AMX int8)
─▶ SAME-S / SAME-L decoder (C++ AMX) ─▶ WAV
▲
audio-to-audio / inpaint: SAME-S / SAME-L encoder (C++ AMX) init-encode
```

The DiT 24-block int8 forward runs in `dit_cpu_amx.so`, T5Gemma in
`t5gemma_cpu_amx.so`, the decoder in `same_{s,l}_cpu_amx.so` (bf16) or
`same_{s,l}_int8fused_cpu_amx.so` (int8), and the encoder (audio→latent, for
`--init-audio`) in `same_{s,l}_encoder_cpu_amx.so` — all static oneDNN + AOT
kernels. The fp32 pre/post and the pingpong sampler are numpy. **Every mode is
100% torch-free** — text-to-audio, CFG, and (now) audio-to-audio / inpainting,
which encode the input through the C++ AMX SAME encoders (the old fp32 torch
autoencoder is kept only as a fallback).

## Documentation

- **[BUILD.md](BUILD.md)** — how to (re)build every `.so` + dump its weight blob (oneDNN, the AOT-Triton DiT, HF weight sourcing).
- **[TESTING.md](TESTING.md)** — the test suite: per-engine validations (`tests/`) + the CLI integration matrix; run `bash tests/run_all.sh`.
- **[LESSONS.md](LESSONS.md)** — hard-won findings (AMX-only speedups, fp32 attention islands, the bf16-RoPE bug, int8 quant, the global-static-weights trap, …). Read before touching the engines.

## One model, four modes

cpu-amx ships the **medium** DiT only (the int8 C++ core). For `sm-music` /
`sm-sfx`, use [`optimized/tflite`](../tflite) or [`optimized/mlx`](../mlx).

| `--decoder` | codec | notes |
|-------------|-------|-------|
| `same-l` (default) | native 426M medium codec | best fidelity |
| `same-s` | distilled 50M | faster, shares the medium latent space |

| mode | flags | example |
|------|-------|---------|
| text-to-audio | `--prompt P` | new clip from a description |
| audio-to-audio | `--prompt P --init-audio IN.wav --init-noise-level σ` | variation of an existing clip |
| inpainting | `--prompt P --init-audio IN.wav --inpaint-range "S,E"` | regenerate one span, keep the rest |
| CFG + negative | `--cfg 3.0 --negative-prompt P_NEG` | steer toward / away from prompts |

## Run

`./sa3` is a thin wrapper that picks a python interpreter (`$SA3_PYTHON`, else
`.venv/bin/python`, else `python3`) and runs `scripts/sa3_cpu_amx.py`.

```bash
# Text-to-audio (native codec)
./sa3 --prompt "A beautiful piano arpeggio grows into a cinematic climax" \
--dit medium --decoder same-l --seconds 30 --out piano.wav

# Faster distilled decoder
./sa3 --prompt "lofi house loop, 120 BPM" --dit medium --decoder same-s --seconds 15 --out lofi.wav

# int8 fused decoder (smaller/faster), more threads
./sa3 --prompt "techno beat" --dit medium --decoder same-s \
--decoder-precision int8 --threads 32 --out techno.wav

# Audio-to-audio variation (uses the fp32 torch AE to encode the input)
./sa3 --prompt "jazz fusion with electric piano" --dit medium --decoder same-l \
--init-audio funk.wav --init-noise-level 0.7 --out funk_jazz.wav

# Inpaint seconds 4-7 (kept region stays bit-exact in latent space)
./sa3 --prompt "explosive drum break" --dit medium --decoder same-l \
--init-audio funk.wav --inpaint-range "4,7" --out funk_drums.wav

# CFG + negative prompt
./sa3 --prompt "ambient drone" --cfg 3.0 --negative-prompt "drums, vocals" \
--dit medium --decoder same-l --out drone.wav

# Play after writing (ffplay/aplay/paplay/afplay), and all examples
./sa3 --prompt "rainforest" --dit medium --decoder same-l --play
./sa3 --help
```

Omit `--decoder` for an interactive picker. Omit `--prompt` for a stdin prompt.
Relative `--out` paths land in `output/`; absolute paths are used as-is. The
output path is printed as a `▸ saved` line at the end.

### Without the wrapper

```bash
python scripts/sa3_cpu_amx.py --prompt "..." --dit medium --decoder same-l
```

## Web UI (gradio)

```bash
./sa3-gradio # public gradio.live share link (same-l default)
./sa3-gradio --no-share # local-only (http://127.0.0.1:7860)
./sa3-gradio --decoder same-s
```

Every mode is wired: text-to-audio, CFG 0–10 + negative prompt + APG,
audio-to-audio (guide audio + init_noise_level), and inpainting (reference audio
+ start/end sliders). Each clip renders a 3-band tinted stereo mel spectrogram
(numpy port — no torch) with a click-to-seek playhead. Decoder and precision
switch from the dropdowns; WAVs land in `output/gradio/`. Extra UI packages
(`gradio`, `pillow`, `soundfile`) — `pip install -r requirements-gradio.txt`.

> Each generation runs in a **fresh subprocess** (of the tested CLI). The int8
> C++ DiT core is single-length per process, so this keeps changing seconds /
> steps safe at the cost of reloading the (mmap'd) weights each time.

## Precision & threads

MLX's `--dit-dtype` maps to two cpu-amx dials:

| flag | default | notes |
|------|---------|-------|
| `--decoder-precision` | `bf16` | `bf16` = best fidelity; `int8` = SmoothQuant+GPTQ fused w8a8 (smaller/faster) |
| `--threads` | 16 | threads for T5Gemma + the decoder |

The **DiT is int8-fixed** (the only shipped C++ core) and pinned to **1 thread**
— its `.so` heap-races at higher thread counts; at short clip lengths one thread
is already fast. T5Gemma is bf16 regardless.

## Flag reference

| Flag | Default | Notes |
|------|---------|-------|
| `--prompt` | (asks) | Text prompt; empty string = unconditional |
| `--negative-prompt` | — | CFG uncond branch; only used when `--cfg ≠ 1.0` |
| `--dit` | medium | **medium only**; `sm-music`/`sm-sfx` rejected (use tflite/mlx) |
| `--decoder` | same-l | `same-l` (native) or `same-s` (distilled) |
| `--decoder-precision` | bf16 | `bf16` or `int8` |
| `--threads` | 16 | T5Gemma + decoder threads (DiT is pinned to 1) |
| `--seconds` | 30 | Output length; `T_lat = ceil(seconds·44100/4096)` |
| `--steps` | 8 | Pingpong steps; 1 = single forward, 8 = sweet spot |
| `--seed` | random | Set for reproducibility; printed at the end |
| `--cfg` | 1.0 | Guidance scale; 1.0 = off, >1 toward prompt, <1 toward uncond |
| `--apg` | 1.0 | Adaptive Projected Guidance; only when `--cfg ≠ 1` |
| `--init-audio` | — | WAV input for audio-to-audio / inpaint (fp32 torch AE encode) |
| `--init-noise-level` | 1.0 | σmax; 0.4–0.8 typical for variation, 1.0 = full regen |
| `--inpaint-range` | — | `START,END` seconds; regenerate that span, keep the rest |
| `--free-models` | on | Free each model after last use; `--no-free-models` keeps them |
| `--out` / `-o` | auto | Relative → `output/`; absolute → as-is |
| `--play` | off | Play after writing (ffplay/aplay/paplay/afplay) |
| `--lora` | — | **not supported** in cpu-amx (accepted + ignored with a note) |

## Files

```
cpu-amx/
├── sa3 ← CLI wrapper (use this)
├── sa3-gradio ← web UI wrapper
├── README.md
├── requirements.txt ← numpy, sentencepiece, soundfile (+ torch for a2a/inpaint)
├── requirements-gradio.txt ← gradio, pillow, soundfile
├── assets/
│ ├── t5gemma_f16.npz ← SentencePiece tokenizer (4.2 MB)
│ └── cond_medium.npz ← conditioner weights (learned padding + seconds embedder)
├── output/ ← default landing zone for WAVs
└── scripts/
├── sa3_cpu_amx.py ← orchestrator CLI (invoked by ./sa3)
├── sa3_gradio.py ← web UI (invoked by ./sa3-gradio)
├── pipeline.py ← vendored numpy pipeline (tokenizer/conditioner/sampler/CFG/WAV)
├── backends.py ← C++ AMX engine loaders + fp32 torch AE encoder
├── spec.py ← mel-spectrogram renderer (numpy, torch-free)
├── examples.py ← shared examples block (--help)
└── test_all_configs.py ← full-stack self-test (assets + every CLI mode)
```

The heavy weights are **not** in this directory. Each engine lives in its own directory under
the *engine home*, and `scripts/weights.py` downloads the published binaries into it on first
use (from `stabilityai/stable-audio-3-optimized/cpu-amx/`):

```bash
export SA3_CPUAMX_HOME=~/.cache/stable-audio-3/cpu-amx # the default
```

| component | engine directory (under `$SA3_CPUAMX_HOME`) |
|-----------|---------------------------------------------|
| T5Gemma | `t5gemma_cpu_amx/` |
| DiT (medium, int8 / bf16) | `dit_medium_cpu_amx/` |
| SAME-S / SAME-L decoder (bf16) | `same_{s,l}_cpu_amx/` |
| SAME-S / SAME-L decoder (int8) | `same_{s,l}_int8fused_cpu_amx/` |
| SAME-S / SAME-L encoder (a2a/inpaint) | `same_{s,l}_encoder_cpu_amx/` |

Both the Python loader and the C++ engines read `$SA3_CPUAMX_HOME`, so the tree is
relocatable — nothing absolute is compiled in.

## Notes on the design

- **One process = one T_lat.** The int8 C++ DiT core allocates length-specific
scratch on first call and heap-corrupts if invoked at a *second* sequence
length in the same process; it also double-frees at teardown. The CLI runs one
generation per process and calls `os._exit(0)` before teardown, so this is
invisible. The gradio isolates every generation in a fresh CLI subprocess.
- **Encoder is the only torch step.** No C++/TFLite encoder exists for this
platform, so `--init-audio` (audio-to-audio / inpainting) loads the fp32 torch
SAME-L autoencoder to encode the input to latents (auto-picks a free CUDA GPU;
CPU fallback). It is imported lazily — pure text-to-audio never touches torch.
- **Inpainting is paste-back only.** The C++ DiT core does not expose the
`local_add_cond` (to-local-embed) conditioning channel, so the regenerated
span is conditioned on the prompt (+ noise) rather than the surrounding
context. Per-step paste-back keeps the latents **outside** the mask bit-exact.
- **CFG uncond convention.** With no negative prompt, the uncond branch runs the
DiT on the conditioner's **learned padding embedding** (conditioner applied to
an all-zero T5 hidden + mask) — the standard unconditional case, matching the
TensorRT / baked-TFLite releases. A negative prompt replaces it with that
prompt's conditioning. CFG is a sequential dual-pass (no batch-2 C++ DiT).
- **Verified fidelity.** The C++ T5Gemma matches TFLite at cosine 0.99997; the
full C++ pipeline matches the TFLite reference at ~43 dB / 0.997 correlation.

## License & attribution

Model weights derived from Stability AI's Stable Audio 3 checkpoints. T5Gemma
text encoder from Google. Use of the Stable Audio 3 weights is governed by the
**Stability AI Community License** — see <https://stability.ai/license>.
Binary file added optimized/cpu-amx/assets/cond_medium.npz
Binary file not shown.
Binary file added optimized/cpu-amx/assets/t5gemma_f16.npz
Binary file not shown.
Loading
Loading