diff --git a/optimized/cpu-amx/.gitignore b/optimized/cpu-amx/.gitignore new file mode 100644 index 00000000..0f0963a2 --- /dev/null +++ b/optimized/cpu-amx/.gitignore @@ -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 diff --git a/optimized/cpu-amx/BUILD.md b/optimized/cpu-amx/BUILD.md new file mode 100644 index 00000000..508bd9b2 --- /dev/null +++ b/optimized/cpu-amx/BUILD.md @@ -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" \ + .cpp -o .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 # -> weights.bin + weights_manifest.txt +# bf16 decoders +python build/same_s_bf16/dump_weights.py +python build/same_l_bf16/dump_weights.py +# int8 decoders (naive grid) +python build/same_s_int8/dump_weights_int8.py +python build/same_l_int8/dump_weights_int8.py +``` + +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 +``` + +Produces `/{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 ``. 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. diff --git a/optimized/cpu-amx/README.md b/optimized/cpu-amx/README.md new file mode 100644 index 00000000..2d2fd122 --- /dev/null +++ b/optimized/cpu-amx/README.md @@ -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 . diff --git a/optimized/cpu-amx/assets/cond_medium.npz b/optimized/cpu-amx/assets/cond_medium.npz new file mode 100644 index 00000000..127449f6 Binary files /dev/null and b/optimized/cpu-amx/assets/cond_medium.npz differ diff --git a/optimized/cpu-amx/assets/t5gemma_f16.npz b/optimized/cpu-amx/assets/t5gemma_f16.npz new file mode 100644 index 00000000..aeb4e5b8 Binary files /dev/null and b/optimized/cpu-amx/assets/t5gemma_f16.npz differ diff --git a/optimized/cpu-amx/build/dit/compile_isa_all.py b/optimized/cpu-amx/build/dit/compile_isa_all.py new file mode 100644 index 00000000..5f2b1f95 --- /dev/null +++ b/optimized/cpu-amx/build/dit/compile_isa_all.py @@ -0,0 +1,203 @@ +#!/usr/bin/env python +"""Per-ISA AOT compile of the full fused-runtime Triton kernel set for the SA3-medium int8 DiT. + +Run once per ISA in its OWN subprocess with (set by the launcher): + TRITON_CPU_TARGET_FEATURES= (forces backend.cpu_features -> MLIR dot lowering) + TRITON_CPU_AOT_FORCE_ASM_FEATURES=1 (forces emitted ISA via translate_to_asm; REQUIRED) + TRITON_CACHE_DIR= (cpu_features is NOT in the cache key -> must isolate) +Argv: + +Produces /{so/, cpp_kernels.txt, kernels_abi.json, so_flash/_flash_bm{64,128}.so} +matching the dispatch keys the C++ driver expects (dump_bin.py mapping), so + dit_isa_forward --aotdir loads the per-ISA kernels. + +AVX2 note: gemm_i8 (BK=256) triggers a multi-minute / multi-GB LLVM blow-up on AVX2 (Stage-0 +Risk 2). It is NOT on the fused (oneDNN) runtime path, so for AVX2 we compile it at BK=64 +(int8->int32 is tiling-invariant-exact, Stage-0 D2) purely to keep the set complete + the +--gemm triton isolation runnable. Set NOGEMM=1 to skip gemm_i8 entirely. +""" +import os, sys, re, glob, json, shutil, time, subprocess +os.environ.setdefault("TRITON_CPU_BACKEND", "1") +import numpy as np, torch +sys.path.insert(0, os.environ.get("SA3_TRITON_CPU", + os.path.join(os.environ.get("SA3_CPUAMX_HOME", "engines"), "dit_triton"))) +import kernels as K, kernels_fused as KF +from model_p3 import DiTTritonP3 +from gen_reference import read_subset, CKPT + +TAG = sys.argv[1] +OUT = sys.argv[2] +SO_DIR = os.path.join(OUT, "so") +FLASH_DIR = os.path.join(OUT, "so_flash") +os.makedirs(SO_DIR, exist_ok=True) +os.makedirs(FLASH_DIR, exist_ok=True) +L = int(os.getenv("L", "128")) +NOGEMM = os.getenv("NOGEMM", "0") == "1" +torch.set_num_threads(8) + +# AVX2: shrink gemm_i8 contraction tile to dodge the BK=256 legalization blow-up (unused on the +# fused path; int8 GEMM is exact regardless of BK). AMX/VNNI keep the shipped BK=256. +if TAG == "avx2": + K.GEMM_I8 = {"BM": 32, "BN": 64, "BK": 64, "GROUP_M": 8} + print(f"[{TAG}] GEMM_I8 -> BK=64 (avoid AVX2 blow-up)", flush=True) + +print(f"[{TAG}] TRITON_CPU_TARGET_FEATURES={os.getenv('TRITON_CPU_TARGET_FEATURES')}") +print(f"[{TAG}] AOT_FORCE_ASM_FEATURES={os.getenv('TRITON_CPU_AOT_FORCE_ASM_FEATURES')}") +print(f"[{TAG}] TRITON_CACHE_DIR={os.getenv('TRITON_CACHE_DIR')} NOGEMM={NOGEMM}", flush=True) + +KERNELS = { + "_gemm_kernel": K._gemm_kernel, + "_quant_rows_kernel": K._quant_rows_kernel, + "_rmsnorm_kernel": K._rmsnorm_kernel, + "_rope_kernel": K._rope_kernel, + "_flash_diff_i8_kernel": K._flash_diff_i8_kernel, + "_rmsnorm_mod_q_kernel": KF._rmsnorm_mod_q_kernel, + "_gemm_i8_raw_kernel": KF._gemm_i8_raw_kernel, + "_deq_kernel": KF._deq_kernel, + "_deq_glu_q_kernel": KF._deq_glu_q_kernel, + "_deq_gate_res_kernel": KF._deq_gate_res_kernel, + "_deq_add_kernel": KF._deq_add_kernel, +} + + +def all_compiled(fn): + out = [] + def scan(o, d=0): + if d > 6 or o is None: return + if isinstance(getattr(o, "asm", None), dict) and o.asm: + out.append(o); return + if isinstance(o, dict): [scan(v, d + 1) for v in o.values()] + elif isinstance(o, (list, tuple, set)): [scan(v, d + 1) for v in o] + for a in ("cache", "device_caches"): + scan(getattr(fn, a, None)) + return out + + +def ctype_code(ty): + if ty.startswith("*"): return "P" + return {"i1": "i1", "i8": "i8", "i16": "i16", "i32": "i32", "i64": "i64", + "u32": "u32", "u64": "u64", "fp16": "f16", "bf16": "bf16", + "fp32": "f32", "fp64": "f64"}.get(ty, ty) + + +def define_line(ck): + for ln in ck.asm.get("llir", "").splitlines(): + if ln.startswith("define") and "@" in ln: + return ln[:800] + return "" + + +def locate_so(ck, name): + so = ck.asm.get("so", None) + if isinstance(so, (bytes, bytearray)): + p = os.path.join(SO_DIR, f"_tmp_{ck.hash}.so") + with open(p, "wb") as f: f.write(so) + return p + if isinstance(so, str) and os.path.exists(so): + return so + cache = os.getenv("TRITON_CACHE_DIR", os.path.expanduser("~/.triton/cache")) + hits = glob.glob(os.path.join(cache, ck.hash, f"{name}.so")) or \ + glob.glob(os.path.join(cache, "**", f"{name}.so"), recursive=True) + return hits[0] if hits else None + + +def mnem(sopath): + txt = subprocess.run(["objdump", "-d", sopath], capture_output=True, text=True).stdout + c = {} + for m in ("tdpbssd", "vpdpbusd", "vpmaddubsw", "vpmaddwd", "vpmulld", "ldtilecfg", "tileloadd"): + c[m] = len(re.findall(r"\b" + m + r"\b", txt)) + c["zmm"] = len(re.findall(r"%zmm\d+", txt)); c["ymm"] = len(re.findall(r"%ymm\d+", txt)) + return c + + +if __name__ == "__main__": + ref = np.load(os.path.join(os.environ.get("SA3_TRITON_CPU", + os.path.join(os.environ.get("SA3_CPUAMX_HOME", "engines"), "dit_triton")), + "ref", f"ref_L{L}.npz")) + inp = [torch.from_numpy(ref[k]) for k in ("x", "t", "cross", "gcond")] + sd = read_subset(CKPT) + m = DiTTritonP3(sd, T_lat=L, prec="int8", backend="tri") + t0 = time.time() + with torch.no_grad(): + m.forward(*inp, num_cpu_threads=8) + print(f"[{TAG}] warm int8/tri forward {time.time()-t0:.1f}s", flush=True) + + manifest = [] + seen = set() + for name, fn in KERNELS.items(): + for ck in all_compiled(fn): + if ck.hash in seen: continue + seen.add(ck.hash) + sig = ck.src.signature + consts = {k[0] if isinstance(k, tuple) else k: v for k, v in ck.src.constants.items()} + surv, cxpr = [], {} + for i, (an, ty) in enumerate(sig.items()): + if ty == "constexpr": cxpr[an] = consts.get(i) + else: surv.append((an, ctype_code(ty))) + dl = define_line(ck) + ndef = dl.count("ptr ") + len(re.findall(r"\bi\d+ %", dl)) + len(re.findall(r"\bfloat %", dl)) + len(re.findall(r"\bdouble %", dl)) + sp = locate_so(ck, name) + so_dst = None + if sp and os.path.exists(sp): + so_dst = os.path.join(SO_DIR, f"{name}__{ck.hash[:10]}.so") + shutil.copy(sp, so_dst) + manifest.append(dict(kernel=name, hash=ck.hash, + so=os.path.basename(so_dst) if so_dst else None, + surviving_args=surv, n_surviving=len(surv), n_define_args=ndef, + abi_ok=(ndef == len(surv) + 6), constexprs=cxpr)) + with open(os.path.join(OUT, "kernels_abi.json"), "w") as f: + json.dump(manifest, f, indent=2) + print(f"[{TAG}] {len(manifest)} specializations; abi_ok all: {all(r['abi_ok'] for r in manifest)}", flush=True) + + # ---- dispatch table (dump_bin.py mapping) ---- + def pick(kernel, **cx): + c = [r for r in manifest if r["kernel"] == kernel and all(r["constexprs"].get(k) == v for k, v in cx.items())] + return c[0] if c else None + rows = [] + def add(key, kernel, **cx): + r = pick(kernel, **cx) + if r is None or r["so"] is None: + print(f"[{TAG}] WARN no .so for {key} ({kernel} {cx}); skipping"); return + rows.append((key, r["so"], kernel)) + gf = [r for r in manifest if r["kernel"] == "_gemm_kernel" and r["constexprs"].get("HAS_BIAS") is False and "M" not in r["constexprs"]] + if gf: rows.append(("gemm_fp", gf[0]["so"], "_gemm_kernel")) + if not NOGEMM: add("gemm_i8", "_gemm_i8_raw_kernel") + add("quant", "_quant_rows_kernel") + add("rmsnorm", "_rmsnorm_kernel", HAS_G=True, BK=64) + add("rope", "_rope_kernel") + add("flash", "_flash_diff_i8_kernel") + add("rmsmodq_mod1", "_rmsnorm_mod_q_kernel", HAS_MOD=True) + add("rmsmodq_mod0", "_rmsnorm_mod_q_kernel", HAS_MOD=False) + for bn in (256, 4096, 8192): + add(f"deq_bn{bn}_bias0", "_deq_kernel", BN=bn, HAS_BIAS=False) + add("deqglu", "_deq_glu_q_kernel") + add("deqgate_bn2048_bias0", "_deq_gate_res_kernel", BN=2048, HAS_BIAS=False) + add("deqgate_bn2048_bias1", "_deq_gate_res_kernel", BN=2048, HAS_BIAS=True) + add("deqadd_bn2048", "_deq_add_kernel", BN=2048, HAS_LOCAL=True) + with open(os.path.join(OUT, "cpp_kernels.txt"), "w") as f: + for key, so, sym in rows: + f.write(f"{key} {so} {sym}\n") + print(f"[{TAG}] cpp_kernels.txt {len(rows)} dispatch slots", flush=True) + + # ---- retiled flash BM=64,128 (fast path; same ABI, only BM differs) ---- + H, D = 24, 64 + g = lambda s: (torch.randn(H, s, D) * 4) + for BM in (64, 128): + K.flash_diff_i8(g(1356), g(1356), g(1356), g(1356), g(1356), 0.125, num_cpu_threads=8, BM=BM, BN=64) + for ck in all_compiled(K._flash_diff_i8_kernel): + names = list(ck.src.signature.keys()); idx = {n: i for i, n in enumerate(names)} + c = dict(ck.src.constants) + def lk(nm): i = idx[nm]; return c.get((i,), c.get(i)) + bm, bn = lk('BM'), lk('BN') + if bn == 64 and bm in (64, 128): + sp = locate_so(ck, "_flash_diff_i8_kernel") + if sp: + dst = os.path.join(FLASH_DIR, f"_flash_bm{bm}.so"); shutil.copy(sp, dst) + print(f"[{TAG}] flash BM={bm} -> {dst}", flush=True) + + # ---- mnemonic proof for the two int8-dot kernels ---- + for key in ("gemm_i8", "flash"): + r = [x for x in rows if x[0] == key] + if r: + print(f"[{TAG}] mnem {key}: {mnem(os.path.join(SO_DIR, r[0][1]))}", flush=True) + print(f"[{TAG}] DONE", flush=True) diff --git a/optimized/cpu-amx/build/dit/dit_amx_forward3.cpp b/optimized/cpu-amx/build/dit/dit_amx_forward3.cpp new file mode 100644 index 00000000..af720e40 --- /dev/null +++ b/optimized/cpu-amx/build/dit/dit_amx_forward3.cpp @@ -0,0 +1,456 @@ +// Speedprove v3: optimized TORCH-FREE C++ driver for the SA3-medium int8 DiT. +// Same math/ABI as dit_amx_forward2.cpp (bit-exact), but the orchestration layer is rewritten: +// 1. Preallocated, reused scratch pools (NO per-call malloc, NO zero-init of overwritten buffers). +// 2. oneDNN dst-buffer + dnnl::memory-handle reuse (set_data_handle) alongside the primitive cache. +// 3. OMP-parallelized slice_to_heads/from_heads (now memcpy), ssg add, no post xt copy. +// Bit-exact vs the model_p3(int8/tri) golden (int8@int8->int32 is integer-exact). Timing harness @L. +// +// build: see RESULTS_speedprove.md +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +static const int E=1536, H=24, D=64, RD=32, MEM=64, CROSS=257, DEPTH=24; +static const int FF=6144; // FFN inner +static const float NORM_EPS=1e-5f, QK_EPS=1e-6f; +static const float SCALE=0.125f; // D**-0.5 = 64**-0.5 + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string AOT_S = std::string(sa3_home()) + "/dit_medium_cpu_amx/aot_stage2"; +static const char* AOT = AOT_S.c_str(); // so/ + cpp_kernels.txt (L-independent) +static std::string COREBASE = std::string(sa3_home()) + "/dit_medium_cpu_amx/core_L128"; // .bin + _manifest.txt + +static bool USE_ONEDNN=true; // int8 linear GEMM backend +static void* FLASH_OVERRIDE=nullptr; // optional alternate flash kernel .so entry +static int FLASH_BM=128; // flash query-block size (BM); constexpr in the .so. +// BM only tiles the (independent) query rows -> bit-exact vs BM=32, but 4x fewer per-tile kernel +// invocations, cutting the per-launch AMX-config/dispatch overhead. bm{64,128}.so live in so_flash/. +static std::string FLASH_SO_DIR_S = std::string(sa3_home()) + "/dit_medium_cpu_amx/so_flash"; +static const char* FLASH_SO_DIR = FLASH_SO_DIR_S.c_str(); + +static inline int cdiv(int a,int b){return (a+b-1)/b;} +static int npow2(int n){int p=1;while(p shp;}; +static std::map TEN; +static char* BASE=nullptr; +static void load_core(){ + std::string bin=COREBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(COREBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("core.bin mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static int8_t* I8 (const std::string&k){return (int8_t*)TEN.at(k).p;} + +// ------------------------- dlopen dispatch table ------------------------- +static std::map FN; +static void load_kernels(){ + std::ifstream kf(std::string(AOT)+"/cpp_kernels.txt"); std::string key,so,sym; + while(kf>>key>>so>>sym){ + std::string path=std::string(AOT)+"/so/"+so; + void* h=dlopen(path.c_str(),RTLD_NOW|RTLD_LOCAL); + if(!h){fprintf(stderr,"dlopen %s: %s\n",path.c_str(),dlerror());exit(1);} + void* f=dlsym(h,sym.c_str()); + if(!f){fprintf(stderr,"dlsym %s\n",sym.c_str());exit(1);} + FN[key]=f; + } + printf("dlopen'd %ld AOT kernels\n",(long)FN.size()); +} + +// ------------------------- kernel ABI typedefs (surviving args + 6 grid u32) ------------------------- +typedef void(*t_gemm_i8)(void*,void*,void*,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_gemm_fp)(void*,void*,void*,void*,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_quant)(void*,void*,void*,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_rms)(void*,void*,void*,int,int,int,int,float,u32,u32,u32,u32,u32,u32); +typedef void(*t_rope)(void*,void*,void*,void*,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_flash)(void*,void*,void*,void*,void*,void*,void*,void*,void*,void*,void*, + int,int,float,int,int,int,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_rmsmodq)(void*,void*,void*,void*,void*,void*,int,int,int,int,float,u32,u32,u32,u32,u32,u32); +typedef void(*t_deq)(void*,void*,void*,void*,void*,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_deqglu)(void*,void*,void*,void*,void*,void*,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_deqgate)(void*,void*,void*,void*,void*,void*,void*,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_deqadd)(void*,void*,void*,void*,void*,void*,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); + +// ------------------------- reusable scratch pools (allocated once, no zero-init) ------------------------- +template struct Pool{ + std::vector> slab; size_t elems=0; int cur=0; + void init(int n,size_t e){ elems=e; slab.assign(n,std::vector(e)); cur=0; } + T* get(){ if(cur>=(int)slab.size()){fprintf(stderr,"POOL OVERFLOW (%d slabs)\n",(int)slab.size());exit(3);} return slab[cur++].data(); } + void reset(){ cur=0; } +}; +static Pool HEADf; // [H,S,D]=S*E sized head-space tensors +static Pool WIDEf; // deq outputs: qkv[S,5E], q2[S,2E], kv[CROSS,3E], proj/post +static Pool QROW; // int8 row-quant activations up to [S,FF] +static Pool SROWf; // row-scale vectors [S]/[Mp] +static std::vector ACC; // single shared GEMM int32 accumulator (max S*2FF) +static std::vector XA,XB,RES,SSG; // x ping-pong, residual snapshot, ssg buffer +// flash internal scratch (reused every flash call; calls are sequential) +static std::vector FL_qmi,FL_kmi,FL_qdi,FL_kdi,FL_viq; +static std::vector FL_out,FL_qsm,FL_ksm,FL_qsd,FL_ksd,FL_vsc; + +static void alloc_scratch(int S){ + // Head-space & flash tensors span S rows (self-attn) OR CROSS rows (cross-attn K/V side); + // size them to max(S,CROSS) so short sequences (S= kv[CROSS,3E] for all S>=... (holds L128 too) + QROW .init(8, (size_t)S*FF); + SROWf.init(8, (size_t)Sm); + ACC.assign((size_t)S*2*FF,0); + XA.assign(se,0); XB.assign(se,0); RES.assign(se,0); SSG.assign((size_t)6*E,0); + FL_qmi.assign(sme,0);FL_kmi.assign(sme,0);FL_qdi.assign(sme,0);FL_kdi.assign(sme,0);FL_viq.assign(sme,0); + FL_out.assign(sme,0);FL_qsm.assign(hsm,0);FL_ksm.assign(hsm,0);FL_qsd.assign(hsm,0);FL_ksd.assign(hsm,0);FL_vsc.assign(hsm,0); + printf("scratch: HEADf16x%.1fMB WIDEf4x%.1fMB QROW8x%.1fMB ACC%.1fMB (S=%d Sm=%d)\n", + sme*4/1e6,(double)S*5*E*4/1e6,(double)S*FF/1e6,(double)S*2*FF*4/1e6,S,Sm); +} +static void reset_block_pools(){ HEADf.reset(); WIDEf.reset(); QROW.reset(); SROWf.reset(); } + +// ------------------------- oneDNN int8 GEMM (s8 x s8 -> s32): primitive + memory-handle cache --------- +static dnnl::engine* ENG=nullptr; +static dnnl::stream* STRM=nullptr; +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M MM; +static void onednn_init(){ + ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); + STRM=new dnnl::stream(*ENG); +} +// write s8[M,K] x s8[K,N] -> s32[M,N] into caller-provided dst (no alloc, no zero-init) +static void gemm_i8_onednn(const int8_t*a,const int8_t*b,int M,int N,int K,int32_t* dst){ + using dt=dnnl::memory::data_type; + MMKey key{M,N,K}; auto it=MM.find(key); + if(it==MM.end()){ + dnnl::memory::desc a_md({M,K},dt::s8, {K,1}); // A row-major [M,K] + dnnl::memory::desc b_md({K,N},dt::s8, {N,1}); // B row-major [K,N] (plain, matches _int_mm) + dnnl::memory::desc c_md({M,N},dt::s32,{N,1}); // C row-major [M,N] + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)a), + dnnl::memory(pd.weights_desc(),*ENG,(void*)b), + dnnl::memory(pd.dst_desc(),*ENG,(void*)dst)}; + it=MM.emplace(key,std::move(e)).first; + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)a); + e.bm.set_data_handle((void*)b); + e.cm.set_data_handle((void*)dst); + e.prim.execute(*STRM,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + STRM->wait(); +} + +// ------------------------- kernel wrappers (grid looped in C, OMP-parallel; write into dst) ---------- +static void gemm_i8_triton(const int8_t*a,const int8_t*b,int M,int N,int K,int32_t* dst){ + std::memset(dst,0,(size_t)M*N*sizeof(int32_t)); auto fn=(t_gemm_i8)FN.at("gemm_i8"); // triton path accumulates + u32 g=cdiv(M,32)*cdiv(N,64); + #pragma omp parallel for schedule(static) + for(u32 x=0;x memcpy) ------------------------- +static void slice_to_heads(const float*src,int S,int total,int coloff,float* o){ // [S,total]->[H,S,D] + #pragma omp parallel for collapse(2) schedule(static) + for(int s=0;s[S,E] + #pragma omp parallel for collapse(2) schedule(static) + for(int h=0;h run_forward(int S){ + int Mp=S-MEM; + float* x=XA.data(); float* xnext=XB.data(); // ping-pong for the running activation + std::memcpy(x,F32("x_init"),(size_t)S*E*sizeof(float)); + int8_t* ctx_i8=I8("ctx_i8"); float* ctx_s=F32("ctx_s"); float* gc=F32("gc"); + float* rc=F32("rope_cos"); float* rs=F32("rope_sin"); + for(int b=0;b no copy + int8_t* qx=QROW.get(); float* sx=SROWf.get(); quant_rows(xt,Mp,E,qx,sx); + int32_t* acc=ACC.data(); gemm_i8(qx,I8("pout.q"),Mp,256,E,acc); + float* proj=WIDEf.get(); deq(acc,sx,F32("pout.scale"),Mp,256,proj); + std::vector post((size_t)Mp*256); + gemm_fp(proj,F32("Wpost.wt"),Mp,256,256,post.data()); + #pragma omp parallel for schedule(static) + for(size_t i=0;i<(size_t)Mp*256;i++) post[i]+=proj[i]; + return post; +} + +static double check_mad(const std::vector&post,int Mp){ + float* gold=F32("out"); // shape [1,256,Mp] row-major + double mad=0; for(int c=0;c<256;c++)for(int mm=0;mm&post,int Mp,double&mad,double&cos,double&psnr){ + float* gold=F32("out"); // [256,Mp] row-major (c-major) + double dot=0,na=0,nb=0,se=0,sg=0; mad=0; + for(int c=0;c<256;c++)for(int mm=0;mm0&&nb>0)?dot/(std::sqrt(na)*std::sqrt(nb)):0.0; + psnr=(se>0)?10.0*std::log10(sg/se):1e9; +} + +int main(int argc,char**argv){ + int threads=0, reps=0, warmup=3, L=128; + for(int i=1;i0) COREBASE = std::string(sa3_home()) + "/dit_medium_cpu_amx/core_L" + std::to_string(L); + + if(syscall(SYS_arch_prctl,0x1023,18)!=0){fprintf(stderr,"arch_prctl AMX FAILED\n");return 1;} + omp_set_dynamic(0); omp_set_num_threads(threads); + #pragma omp parallel // ensure every OMP thread has AMX permission + { syscall(SYS_arch_prctl,0x1023,18); } + onednn_init(); + load_core(); load_kernels(); + // Select flash kernel: BM=32 uses the built-in AOT kernel; BM in {64,128} auto-loads the + // bit-exact wider-tile .so from so_flash/ (fewer per-tile launches). --flashso overrides. + if(FLASH_OVERRIDE==nullptr && FLASH_BM!=32){ + std::string fp=std::string(FLASH_SO_DIR)+"/_flash_bm"+std::to_string(FLASH_BM)+".so"; + void* h=dlopen(fp.c_str(),RTLD_NOW|RTLD_LOCAL); + if(h) FLASH_OVERRIDE=dlsym(h,"_flash_diff_i8_kernel"); + if(!FLASH_OVERRIDE){fprintf(stderr,"WARN flash BM=%d .so unavailable (%s); falling back to built-in BM=32\n",FLASH_BM,fp.c_str()); FLASH_BM=32;} + } + printf("flash: BM=%d src=%s\n",FLASH_BM,FLASH_OVERRIDE?"so_flash":"builtin(cpp_kernels)"); + int S=(int)TEN.at("x_init").shp[0]; // S from injected preamble + int Mp=S-MEM; + alloc_scratch(S); + printf("gemm=%s threads=%d L=%d S=%d Mp=%d dnnl_isa=%s\n", + USE_ONEDNN?"onednn":"triton",threads,L,S,Mp,dnnl_cpu_isa2str(dnnl_get_effective_cpu_isa())); + + if(reps<=0){ // correctness-only + auto post=run_forward(S); double mad,cos,psnr; check_quality(post,Mp,mad,cos,psnr); + bool ok=(cos>=0.9999); + printf("[check] gemm=%s max_abs_diff=%.6e cos=%.7f psnr=%.1fdB bit_exact=%s quality_ok(cos>=0.9999)=%s\n", + USE_ONEDNN?"onednn":"triton",mad,cos,psnr,mad==0.0?"YES":"no",ok?"YES":"NO"); + return ok?0:2; + } + // timing mode + std::vector ms; double mad0=-1; + for(int r=0;r(t1-t0).count()/1000.0; + if(r==0) mad0=check_mad(post,Mp); + if(r>=warmup) ms.push_back(m); + } + std::sort(ms.begin(),ms.end()); + double med=ms[ms.size()/2], mn=ms.front(); + printf("[time] gemm=%s threads=%d L=%d per-forward median=%.1f ms min=%.1f ms (reps=%d warmup=%d) bitexact_mad=%.3e\n", + USE_ONEDNN?"onednn":"triton",threads,L,med,mn,reps,warmup,mad0); + return 0; +} diff --git a/optimized/cpu-amx/build/dit/dit_cpu_amx.cpp b/optimized/cpu-amx/build/dit/dit_cpu_amx.cpp new file mode 100644 index 00000000..a4e4adba --- /dev/null +++ b/optimized/cpu-amx/build/dit/dit_cpu_amx.cpp @@ -0,0 +1,411 @@ +// dit_cpu_amx.cpp — torch-free C++ AMX int8 SA3-medium DiT as a PER-STEP CALLABLE .so. +// Refactor of dit_amx_forward3.cpp (the proven, bit-exact standalone driver): the 24-block +// int8/AMX forward is kept byte-identical, but the CLI/timing main() is replaced by a small +// C ABI so the core is callable per sampling step: +// dit_init(core_bin_base, threads) -> load int8 block weights ONCE + AMX/omp/oneDNN/kernels +// dit_forward(x_init, ctx_i8, ctx_s, gc, rope_cos, rope_sin, S, out_post) -> denoised core +// The block/pout/Wpost weights are sequence-length-INDEPENDENT, so ANY core_L*.bin serves as the +// weight source; the per-step preamble tensors (x_init/ctx/gc) and the length-dependent rope are +// supplied by the caller (numpy preamble). Kernels are shape-general (proven L128 .so runs any L). +// +// build: g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I$ONEINC \ +// dit_cpu_amx.cpp -o dit_cpu_amx.so $ONELIB/libdnnl.a -ldl -lpthread -lm +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +static const int E=1536, H=24, D=64, RD=32, MEM=64, CROSS=257, DEPTH=24; +static const int FF=6144; // FFN inner +static const float NORM_EPS=1e-5f, QK_EPS=1e-6f; +static const float SCALE=0.125f; // D**-0.5 = 64**-0.5 + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string AOT_S = std::string(sa3_home()) + "/dit_medium_cpu_amx/aot_stage2"; +static const char* AOT = AOT_S.c_str(); // so/ + cpp_kernels.txt (L-independent) +static std::string COREBASE = std::string(sa3_home()) + "/dit_medium_cpu_amx/core_L320"; // .bin + _manifest.txt (weights) + +static bool USE_ONEDNN=true; // int8 linear GEMM backend +static void* FLASH_OVERRIDE=nullptr; // optional alternate flash kernel .so entry +static int FLASH_BM=128; // flash query-block size (BM); constexpr in the .so. +static std::string FLASH_SO_DIR_S = std::string(sa3_home()) + "/dit_medium_cpu_amx/so_flash"; +static const char* FLASH_SO_DIR = FLASH_SO_DIR_S.c_str(); + +static inline int cdiv(int a,int b){return (a+b-1)/b;} +static int npow2(int n){int p=1;while(p shp;}; +static std::map TEN; +static char* BASE=nullptr; +static void load_core(){ + std::string bin=COREBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(COREBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("[dit_cpu_amx] core weights mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static int8_t* I8 (const std::string&k){return (int8_t*)TEN.at(k).p;} + +// ------------------------- dlopen dispatch table ------------------------- +static std::map FN; +static void load_kernels(){ + std::ifstream kf(std::string(AOT)+"/cpp_kernels.txt"); std::string key,so,sym; + while(kf>>key>>so>>sym){ + std::string path=std::string(AOT)+"/so/"+so; + void* h=dlopen(path.c_str(),RTLD_NOW|RTLD_LOCAL); + if(!h){fprintf(stderr,"dlopen %s: %s\n",path.c_str(),dlerror());exit(1);} + void* f=dlsym(h,sym.c_str()); + if(!f){fprintf(stderr,"dlsym %s\n",sym.c_str());exit(1);} + FN[key]=f; + } + printf("[dit_cpu_amx] dlopen'd %ld AOT kernels\n",(long)FN.size()); +} + +// ------------------------- kernel ABI typedefs (surviving args + 6 grid u32) ------------------------- +typedef void(*t_gemm_i8)(void*,void*,void*,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_gemm_fp)(void*,void*,void*,void*,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_quant)(void*,void*,void*,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_rms)(void*,void*,void*,int,int,int,int,float,u32,u32,u32,u32,u32,u32); +typedef void(*t_rope)(void*,void*,void*,void*,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_flash)(void*,void*,void*,void*,void*,void*,void*,void*,void*,void*,void*, + int,int,float,int,int,int,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_rmsmodq)(void*,void*,void*,void*,void*,void*,int,int,int,int,float,u32,u32,u32,u32,u32,u32); +typedef void(*t_deq)(void*,void*,void*,void*,void*,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_deqglu)(void*,void*,void*,void*,void*,void*,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_deqgate)(void*,void*,void*,void*,void*,void*,void*,int,int,int,int,int,u32,u32,u32,u32,u32,u32); +typedef void(*t_deqadd)(void*,void*,void*,void*,void*,void*,int,int,int,int,int,int,u32,u32,u32,u32,u32,u32); + +// ------------------------- reusable scratch pools (allocated once per S, no zero-init) ------------------------- +template struct Pool{ + std::vector> slab; size_t elems=0; int cur=0; + void init(int n,size_t e){ elems=e; slab.assign(n,std::vector(e)); cur=0; } + T* get(){ if(cur>=(int)slab.size()){fprintf(stderr,"POOL OVERFLOW (%d slabs)\n",(int)slab.size());exit(3);} return slab[cur++].data(); } + void reset(){ cur=0; } +}; +static Pool HEADf; +static Pool WIDEf; +static Pool QROW; +static Pool SROWf; +static std::vector ACC; +static std::vector XA,XB,RES,SSG; +static std::vector FL_qmi,FL_kmi,FL_qdi,FL_kdi,FL_viq; +static std::vector FL_out,FL_qsm,FL_ksm,FL_qsd,FL_ksd,FL_vsc; + +static void alloc_scratch(int S){ + int Sm=std::max(S,CROSS); size_t se=(size_t)S*E, sme=(size_t)Sm*E, hsm=(size_t)H*Sm; + HEADf.init(16, sme); + WIDEf.init(4, (size_t)S*5*E); + QROW .init(8, (size_t)S*FF); + SROWf.init(8, (size_t)Sm); + ACC.assign((size_t)S*2*FF,0); + XA.assign(se,0); XB.assign(se,0); RES.assign(se,0); SSG.assign((size_t)6*E,0); + FL_qmi.assign(sme,0);FL_kmi.assign(sme,0);FL_qdi.assign(sme,0);FL_kdi.assign(sme,0);FL_viq.assign(sme,0); + FL_out.assign(sme,0);FL_qsm.assign(hsm,0);FL_ksm.assign(hsm,0);FL_qsd.assign(hsm,0);FL_ksd.assign(hsm,0);FL_vsc.assign(hsm,0); + printf("[dit_cpu_amx] scratch alloc S=%d Sm=%d (HEADf %.1fMB, ACC %.1fMB)\n", + S,Sm,sme*4/1e6,(double)S*2*FF*4/1e6); +} +static void reset_block_pools(){ HEADf.reset(); WIDEf.reset(); QROW.reset(); SROWf.reset(); } + +// ------------------------- oneDNN int8 GEMM (s8 x s8 -> s32): primitive + memory-handle cache --------- +static dnnl::engine* ENG=nullptr; +static dnnl::stream* STRM=nullptr; +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M MM; +static void onednn_init(){ + ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); + STRM=new dnnl::stream(*ENG); +} +static void gemm_i8_onednn(const int8_t*a,const int8_t*b,int M,int N,int K,int32_t* dst){ + using dt=dnnl::memory::data_type; + MMKey key{M,N,K}; auto it=MM.find(key); + if(it==MM.end()){ + dnnl::memory::desc a_md({M,K},dt::s8, {K,1}); + dnnl::memory::desc b_md({K,N},dt::s8, {N,1}); + dnnl::memory::desc c_md({M,N},dt::s32,{N,1}); + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)a), + dnnl::memory(pd.weights_desc(),*ENG,(void*)b), + dnnl::memory(pd.dst_desc(),*ENG,(void*)dst)}; + it=MM.emplace(key,std::move(e)).first; + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)a); + e.bm.set_data_handle((void*)b); + e.cm.set_data_handle((void*)dst); + e.prim.execute(*STRM,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + STRM->wait(); +} + +// ------------------------- kernel wrappers (grid looped in C, OMP-parallel; write into dst) ---------- +static void gemm_i8_triton(const int8_t*a,const int8_t*b,int M,int N,int K,int32_t* dst){ + std::memset(dst,0,(size_t)M*N*sizeof(int32_t)); auto fn=(t_gemm_i8)FN.at("gemm_i8"); + u32 g=cdiv(M,32)*cdiv(N,64); + #pragma omp parallel for schedule(static) + for(u32 x=0;x memcpy) ------------------------- +static void slice_to_heads(const float*src,int S,int total,int coloff,float* o){ + #pragma omp parallel for collapse(2) schedule(static) + for(int s=0;s env/16. +int dit_init(const char* core_base, int threads){ + if(threads<=0){const char* e=getenv("OMP_NUM_THREADS"); threads=e?atoi(e):16;} + if(core_base && core_base[0]) COREBASE=core_base; + if(syscall(SYS_arch_prctl,0x1023,18)!=0){fprintf(stderr,"[dit_cpu_amx] arch_prctl AMX FAILED\n");return 1;} + omp_set_dynamic(0); omp_set_num_threads(threads); + #pragma omp parallel + { syscall(SYS_arch_prctl,0x1023,18); } + onednn_init(); + load_core(); load_kernels(); + if(FLASH_OVERRIDE==nullptr && FLASH_BM!=32){ + std::string fp=std::string(FLASH_SO_DIR)+"/_flash_bm"+std::to_string(FLASH_BM)+".so"; + void* h=dlopen(fp.c_str(),RTLD_NOW|RTLD_LOCAL); + if(h) FLASH_OVERRIDE=dlsym(h,"_flash_diff_i8_kernel"); + if(!FLASH_OVERRIDE){fprintf(stderr,"[dit_cpu_amx] WARN flash BM=%d .so unavailable; using built-in BM=32\n",FLASH_BM);FLASH_BM=32;} + } + printf("[dit_cpu_amx] init ok: threads=%d gemm=%s flash_bm=%d isa=%s\n", + threads,USE_ONEDNN?"onednn":"triton",FLASH_BM,dnnl_cpu_isa2str(dnnl_get_effective_cpu_isa())); + fflush(stdout); + return 0; +} + +// One 24-block forward. Pointers are caller-owned (void* for ctypes friendliness): +// x_init fp32 [S*E] ctx_i8 int8 [CROSS*E] ctx_s fp32 [CROSS] gc fp32 [6*E] +// rope_cos/rope_sin fp32 [S*RD] out_post fp32 [(S-MEM)*256] (row-major, caller-allocated) +void dit_forward(void* x_init,void* ctx_i8,void* ctx_s,void* gc, + void* rope_cos,void* rope_sin,int S,void* out_post){ + if(S!=CUR_S){ alloc_scratch(S); CUR_S=S; } + run_forward(S,(const float*)x_init,(const int8_t*)ctx_i8,(const float*)ctx_s, + (const float*)gc,(const float*)rope_cos,(const float*)rope_sin,(float*)out_post); +} + +// Expose constants so the python side can sanity-check its preamble shapes. +int dit_E(){return E;} int dit_MEM(){return MEM;} int dit_CROSS(){return CROSS;} int dit_RD(){return RD;} + +} // extern "C" diff --git a/optimized/cpu-amx/build/dit/gen_core.py b/optimized/cpu-amx/build/dit/gen_core.py new file mode 100644 index 00000000..2234b53e --- /dev/null +++ b/optimized/cpu-amx/build/dit/gen_core.py @@ -0,0 +1,34 @@ +#!/usr/bin/env python +"""Parameterized core.bin generator for the speedprove driver (read-only on the npz sources). +Flatten weights_int8_L{L}.npz + injected golden preamble (x_init, ctx_i8, ctx_s, gc) + golden +out -> raw core_L{L}.bin + core_L{L}_manifest.txt in aot_speedprove/. Same byte layout as +aot_stage2/dump_bin.py; cpp_kernels.txt/so/ are L-independent and reused from aot_stage2.""" +import os, sys +import numpy as np + +SRC = os.path.join(os.environ.get("SA3_CPUAMX_HOME", "engines"), + "dit_medium_cpu_amx", "aot_stage2") # read-only npz source +DST = os.path.join(os.environ.get("SA3_CPUAMX_HOME", "engines"), "dit_medium_cpu_amx") +L = int(sys.argv[1]) if len(sys.argv) > 1 else 1292 +DT = {np.dtype("float32"): "f32", np.dtype("int8"): "i8", np.dtype("int32"): "i32"} + +W = dict(np.load(f"{SRC}/weights_int8_L{L}.npz")) +G = dict(np.load(f"{SRC}/golden_L{L}.npz")) +arrays = {} +arrays.update({k: v for k, v in W.items()}) +for k in ("x_init", "ctx_i8", "ctx_s", "gc", "out"): + arrays[k] = G[k] + +buf = bytearray() +man = [] +for name, a in arrays.items(): + a = np.ascontiguousarray(a) + off = len(buf) + buf += a.tobytes() + man.append((name, DT[a.dtype], off, a.size, list(a.shape))) +with open(f"{DST}/core_L{L}.bin", "wb") as f: + f.write(buf) +with open(f"{DST}/core_L{L}_manifest.txt", "w") as f: + for name, dt, off, n, shp in man: + f.write(f"{name} {dt} {off} {n} {' '.join(map(str, shp))}\n") +print(f"core_L{L}.bin {len(buf)} bytes ({len(buf)/1e9:.2f} GB), {len(man)} arrays; x_init={arrays['x_init'].shape} out={arrays['out'].shape}") diff --git a/optimized/cpu-amx/build/same_l_bf16/dump_weights.py b/optimized/cpu-amx/build/same_l_bf16/dump_weights.py new file mode 100644 index 00000000..0d0c4805 --- /dev/null +++ b/optimized/cpu-amx/build/same_l_bf16/dump_weights.py @@ -0,0 +1,102 @@ +#!/usr/bin/env python3 +"""Dump SAME-L fp32 npz -> flat {weights.bin + weights_manifest.txt} for the C++ engine. + +SAME-L = native 426M medium decoder: 12 blocks, dim=1536, 24 heads, hd=64, banded +SWA attention (window +/-17), sin-gate FF from block 5, plain Linear(1536->512) output map. + +Layout choices (mirror same_s_cpu_amx/dump_weights.py): + * Linears: stored **bf16** in oneDNN matmul weight layout [K=in, N=out] (= W.T of the + npz [out,in]). AMX-BF16 GEMM src[M,K]bf16 x wei[K,N]bf16 -> dst[M,N]f32. + * Output map = plain Conv1d(k=1) == Linear: mapping.weight npz is [out=512,in=1536,k=1] + (weight_norm ALREADY fused in export). Reshape -> [512,1536], store bf16 W.T = [1536,512]. + SIMPLER than SAME-S (no k=3 im2col conv). + * DyT (alpha scalar, gamma/beta), biases, new_tokens, running_std: fp32 (the + cancellation-fragile differential-attention elementwise stays fp32). + +bf16 conversion uses torch's round-to-nearest-even (build-time tool; the *runtime* is +torch-free). Manifest line: name dtype byte_offset nelem d0 d1 ... +""" + +import os + +# Paths come from the environment so nothing local is baked in. +# SA3_CPUAMX_HOME where the engine dirs live (default ./engines) +# SA3_REPO checkout providing the reference weights to dump +HOME = os.environ.get("SA3_CPUAMX_HOME", os.path.abspath("engines")) +REPO = os.environ.get("SA3_REPO", os.path.abspath(".")) +import os +import numpy as np +import torch + +SA3 = REPO +WEIGHTS = os.path.join(SA3, "models", "mlx", "same_l_decoder_f32.npz") +OUT = os.path.join(HOME, "same_l_cpu_amx") +NB = 12 + +raw = dict(np.load(WEIGHTS)) +def a(k): # fp32 numpy + return raw[k].astype(np.float32) + +blob = bytearray() +lines = [] + +def put(name, arr, dt): + """dt in {'bf16','f32'}. arr is numpy fp32 (any shape); stored C-contiguous.""" + global blob + arr = np.ascontiguousarray(arr.astype(np.float32)) + off = len(blob) + if dt == "bf16": + t = torch.from_numpy(arr).to(torch.bfloat16) + u = t.view(torch.uint16).numpy().ravel() + blob += u.tobytes() + elif dt == "f32": + blob += arr.ravel().tobytes() + else: + raise ValueError(dt) + shp = " ".join(str(s) for s in arr.shape) + lines.append(f"{name} {dt} {off} {arr.size} {shp}") + +def put_lin(name, w_oi): + """w_oi = npz [out,in]; store bf16 W.T = [in,out] for oneDNN wei[K,N].""" + put(name, np.ascontiguousarray(w_oi.T), "bf16") + +# ---- top-level ---- +put("running_std", a("running_std"), "f32") +put_lin("project_in.w", a("project_in.weight")) # [256,1536] +put("project_in.b", a("project_in.bias"), "f32") # [1536] +put("new_tokens", a("new_tokens").reshape(-1), "f32") # [1536] + +for b in range(NB): + p = f"blocks.{b}" + put(f"b{b}.pre.alpha", a(f"{p}.pre_norm.alpha"), "f32") + put(f"b{b}.pre.gamma", a(f"{p}.pre_norm.gamma"), "f32") + put(f"b{b}.pre.beta", a(f"{p}.pre_norm.beta"), "f32") + put_lin(f"b{b}.qkv.w", a(f"{p}.attn.to_qkv.weight")) # [1536,7680] + put(f"b{b}.qn.alpha", a(f"{p}.attn.q_norm.alpha"), "f32") + put(f"b{b}.qn.gamma", a(f"{p}.attn.q_norm.gamma"), "f32") # [64] + put(f"b{b}.qn.beta", a(f"{p}.attn.q_norm.beta"), "f32") + put(f"b{b}.kn.alpha", a(f"{p}.attn.k_norm.alpha"), "f32") + put(f"b{b}.kn.gamma", a(f"{p}.attn.k_norm.gamma"), "f32") + put(f"b{b}.kn.beta", a(f"{p}.attn.k_norm.beta"), "f32") + put_lin(f"b{b}.out.w", a(f"{p}.attn.to_out.weight")) # [1536,1536] + put(f"b{b}.ff.alpha", a(f"{p}.ff_norm.alpha"), "f32") + put(f"b{b}.ff.gamma", a(f"{p}.ff_norm.gamma"), "f32") + put(f"b{b}.ff.beta", a(f"{p}.ff_norm.beta"), "f32") + put_lin(f"b{b}.glu.w", a(f"{p}.ff.glu_proj.weight")) # [1536,9216] + put(f"b{b}.glu.b", a(f"{p}.ff.glu_proj.bias"), "f32") # [9216] + put_lin(f"b{b}.proj.w", a(f"{p}.ff.proj_out.weight")) # [4608,1536] + put(f"b{b}.proj.b", a(f"{p}.ff.proj_out.bias"), "f32") # [1536] + +# output map: plain Linear(1536->512). mapping.weight [512,1536,1] -> [512,1536] -> W.T [1536,512] +put_lin("map.w", a("mapping.weight").reshape(512, 1536)) # -> bf16 [1536,512] +put("map.b", a("mapping.bias"), "f32") # [512] + +with open(os.path.join(OUT, "weights.bin"), "wb") as f: + f.write(blob) +with open(os.path.join(OUT, "weights_manifest.txt"), "w") as f: + f.write("\n".join(lines) + "\n") + +print(f"wrote weights.bin ({len(blob)/1e6:.1f} MB), {len(lines)} arrays") +print("first/last few manifest lines:") +for l in lines[:5] + ["..."] + lines[-4:]: + print(" ", l) diff --git a/optimized/cpu-amx/build/same_l_bf16/same_l_chunk.inc b/optimized/cpu-amx/build/same_l_bf16/same_l_chunk.inc new file mode 100644 index 00000000..d229925d --- /dev/null +++ b/optimized/cpu-amx/build/same_l_bf16/same_l_chunk.inc @@ -0,0 +1,58 @@ +// same_l_chunk.inc — chunked / chunk-parallel decode (the PRIMARY SAME-L decode path). +// Included at the end of same_l_cpu_amx.cpp (shares all statics). +// +// SAME-L uses a banded SWA mask (window +-17 internal tokens) => the receptive field over 12 +// blocks is a few latent tokens. Tiling [1,256,T] into C-latent interiors with `overlap` latent +// tokens of context each side and concatenating the interior audio reproduces the whole decode +// to high PSNR. Overlap floor is 8 (campaign: ovl 6->64 dB, 8->82, plateau). C=64 is the ship size. +// (Unlike the TFLite dense-[S,S]-mask decoder this dodges no quadratic — the C++ band attention is +// already linear — but chunking bounds working-set/RoPE and enables chunk-parallelism.) + +struct ChunkL{int a,b,lo,hi;}; +static std::vector chunk_plan_l(int T,int C,int overlap){ + std::vector plan; int a=0; + while(a plan=chunk_plan_l(T,C,overlap); + int nch=(int)plan.size(); int L=SIN*T; + if(!parallel){ + static std::vector sublat, subpat; + for(auto&ck:plan){ + int w=ck.hi-ck.lo; + sublat.assign((size_t)LAT*w,0); + for(int c=0;c sublat((size_t)LAT*w), subpat((size_t)OUTCH*SIN*w); + for(int c=0;c no O(T^2)). +// * GLOBAL RoPE positions over the internal sequence (relative RoPE => chunk-local == global). +// * sin-gate FF for blocks 5..11 (value*sin(pi*gate)); silu for blocks 0..4. +// * plain Linear(1536->512) output map (no WNConv1d / im2col). +// * NO midpoint-shift (that is a SAME-S 34-token-chunk trick). +// oneDNN AMX-BF16 GEMMs for every linear; fp32 C++ for the cancellation-fragile elementwise +// (DyT norms, RoPE, GLU) and the differential band attention. +// +// samel_init(weights_base, threads) -> mmap bf16 weights + AMX/omp/oneDNN +// samel_forward(latent[1,256,T], T, out_patches[1,512,16T]) -> whole decode (linear band attn) +// samel_forward_chunked(latent, T, C, overlap, parallel, out) -> chunked decode (PRIMARY: C=64,ovl=8) +// samel_unpatch(patches[1,512,L], L, pcm[1,2,256L]) -> torch-free unpatch +// +// build: g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I$ONEINC \ +// same_l_cpu_amx.cpp -o same_l_cpu_amx.so $ONELIB/libdnnl.a -ldl -lpthread -lm +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +// ── architecture constants (same_l_decoder_torch.py) ── +static const int LAT=256, DIM=1536, H=24, HD=64, RD=32, HALF=16; +static const int NB=12, FF=4608, GLU2=9216, QKV=7680, OUTCH=512; +static const int SUB=17, SIN=16; // 17 internal tok/latent, keep 16 (drop slot 0) +static const int SIN_START=5; // blocks >=5 use sin(pi*gate) gate +static const int BAND=17; // SWA half-window (BLOCK_SIZE=SUB_CHUNK_SIZE=17) +static const float SCALE=0.125f; // HD**-0.5 = 64**-0.5 + +// ── optional phase profiler (SAMEL_PROF=1) ── +static double PROF[8]={0}; static long PROFN=0; +static const char* PROFLBL[8]={"gemm","attn","dyt","rope","glu","cast","misc",""}; +static bool PROF_ON=false; +static inline double wt(){return omp_get_wtime();} +#define TB(id) do{ if(PROF_ON){double _n=wt(); PROF[id]+=_n-_pt; _pt=_n;} }while(0) + +typedef uint16_t bf16; +static inline bf16 f2b(float f){ // round-to-nearest-even f32->bf16 + uint32_t x; std::memcpy(&x,&f,4); + uint32_t r=x+0x7fff+((x>>16)&1); return (bf16)(r>>16); +} +static inline float b2f(bf16 h){ uint32_t x=(uint32_t)h<<16; float f; std::memcpy(&f,&x,4); return f; } + +// ── fast vectorizable transcendentals (pure float, no libm/branches -> AVX-512 auto-vec) ── +static inline float vexp(float x){ + x = x<-87.0f?-87.0f:(x>88.0f?88.0f:x); + float z = x*1.442695041f; + float n = std::floor(z+0.5f); + float f = z-n; + float p = 1.0f+f*(0.6931472f+f*(0.2402265f+f*(0.0555041f+f*(0.0096181f+f*0.0013333f)))); + uint32_t bits=(uint32_t)(((int)n+127)<<23); float s; std::memcpy(&s,&bits,4); + return p*s; +} +static inline float vtanh(float x){ return 1.0f-2.0f/(vexp(2.0f*x)+1.0f); } // == Triton DyT kernel +static inline float vsilu(float x){ return x/(1.0f+vexp(-x)); } +// sin(pi*g): range-reduce g to r in [-0.5,0.5] (g=k+r), sin(pi*g)=(-1)^k * sin(pi*r); +// pi*r in [-pi/2,pi/2] via degree-9 Taylor (err ~2e-6 << bf16). branchless -> SIMD-clean. +static inline float vsinpi(float g){ + float k = std::floor(g+0.5f); + float r = g - k; // [-0.5,0.5] + float s = 1.0f - 2.0f*(float)(((long)k)&1L); // (-1)^k + float y = 3.14159265358979f*r; // [-pi/2,pi/2] + float y2 = y*y; + float p = y*(1.0f + y2*(-0.16666667f + y2*(0.00833333f + y2*(-0.00019841f + y2*2.75573e-6f)))); + return s*p; +} +static inline int cdiv(int a,int b){return (a+b-1)/b;} + +// ------------------------- mmap weights.bin + manifest ------------------------- +struct Ten{void* p; std::string dt; long n; std::vector shp;}; +static std::map TEN; +static char* BASE=nullptr; + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string WBASE = std::string(sa3_home()) + "/same_l_cpu_amx/weights"; +static void load_weights(){ + std::string bin=WBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(WBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("[samel] weights mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static bf16* BF (const std::string&k){return (bf16*)TEN.at(k).p;} + +// ------------------------- oneDNN bf16 matmul (bf16 x bf16 -> f32): primitive+handle cache +static dnnl::engine* ENG=nullptr; +static void onednn_init(){ ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); } +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M mm; dnnl::stream* strm=nullptr; }; +static std::vector MMC; +static void mmcache_init(int nworkers){ + MMC.clear(); MMC.resize(std::max(1,nworkers)+1); + for(auto& c:MMC) c.strm=new dnnl::stream(*ENG); +} +// src A[M,K] bf16, wei B[K,N] bf16 (row-major), dst C[M,N] f32. cache slot `slot`. +static void gemm_bf16(int slot,const bf16*A,const bf16*B,int M,int N,int K,float* C){ + using dt=dnnl::memory::data_type; + MMCache& c=MMC[slot]; + MMKey key{M,N,K}; auto it=c.mm.find(key); + if(it==c.mm.end()){ + #pragma omp critical(mmcreate) + { + it=c.mm.find(key); + if(it==c.mm.end()){ + dnnl::memory::desc a_md({M,K},dt::bf16,{K,1}); + dnnl::memory::desc b_md({K,N},dt::bf16,{N,1}); + dnnl::memory::desc c_md({M,N},dt::f32, {N,1}); + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)A), + dnnl::memory(pd.weights_desc(),*ENG,(void*)B), + dnnl::memory(pd.dst_desc(),*ENG,(void*)C)}; + it=c.mm.emplace(key,std::move(e)).first; + } + } + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)A); e.bm.set_data_handle((void*)B); e.cm.set_data_handle((void*)C); + e.prim.execute(*c.strm,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + c.strm->wait(); +} + +// ------------------------- GLOBAL RoPE table (positions 0..cap-1, 16 freqs) ------------------------- +// SAME-L RoPE runs over the whole internal sequence. Relative RoPE => a chunk's local positions +// give identical band-attention scores as global positions, so decode() uses absolute index m. +static float RINV[HALF]; +static std::vector RCOS, RSIN; // [cap*HALF] +static int RCAP=0; +static void rope_invfreq(){ for(int i=0;iRCAP){ int nc=S+S/4; RCOS.resize((size_t)nc*HALF); RSIN.resize((size_t)nc*HALF); + rope_fill(RCAP,nc); RCAP=nc; } } +} + +// ------------------------- reusable arena (per worker slot) ------------------------- +struct Arena{ + int cap=0; + std::vector x,xt,h,qkv,ao,glu; // f32 activations + std::vector srcb; // bf16 GEMM src staging + void ensure(int M){ + if(M<=cap) return; cap=M; + x.assign((size_t)M*DIM,0); xt.assign((size_t)M*DIM,0); h.assign((size_t)M*DIM,0); + qkv.assign((size_t)M*QKV,0); ao.assign((size_t)M*DIM,0); glu.assign((size_t)M*GLU2,0); + srcb.assign((size_t)M*GLU2,0); // wide enough for any src (K<=FF=4608) and glu val stage + } +}; +static std::vector AR; + +// ------------------------- fp32 elementwise kernels ------------------------- +// DyT: out[m,j] = gamma[j]*tanh(alpha*x[m,j]) + beta[j] +static void dyt(const float*x,float*o,int M,int K,float alpha,const float*g,const float*b,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;m G; + G.resize((size_t)5*S*HD); + float* qg=G.data(); float* kg=qg+(size_t)S*HD; float* vg=kg+(size_t)S*HD; + float* qdg=vg+(size_t)S*HD; float* kdg=qdg+(size_t)S*HD; + for(int t=0;tmm)mm=dm; if(dd>md)md=dd; + } + int n=j1-j0+1; float zm=0,zd=0; + #pragma omp simd reduction(+:zm,zd) + for(int j=0;j patches[512,16T]) --------- +// slot: arena/matmul cache index. par: true -> kernels use omp-for (slot 0); false -> serial. +static void decode(int slot,bool par,const float* latent,int T,float* out_patches){ + Arena& A=AR[slot]; + int S=SUB*T; // internal tokens after new-token expansion (17T) + A.ensure(S); + ensure_rope(S); + float* x=A.x.data(); float* xt=A.xt.data(); float* h=A.h.data(); + float* qkv=A.qkv.data(); float* ao=A.ao.data(); float* glu=A.glu.data(); + bf16* srcb=A.srcb.data(); + float rstd=F32("running_std")[0]; + + // project_in: src [T,256] bf16 = (latent^T * running_std) rows, GEMM -> [T,1536], +bias + // latent is [1,256,T] channel-major: latent[c*T + t] + #pragma omp parallel for schedule(static) if(par) + for(int t=0;t=SIN_START); + char pb[8]; snprintf(pb,sizeof pb,"b%d.",b); + auto W=[&](const char*n){return std::string(pb)+n;}; + // ---- attention: h = pre_norm(x); qkv = to_qkv(h); dyt+rope; band-attn; ao=to_out; x += ao + dyt(x,h,S,DIM,F32(W("pre.alpha"))[0],F32(W("pre.gamma")),F32(W("pre.beta")),par); TB(2); + to_bf16(h,srcb,(size_t)S*DIM,par); TB(5); + gemm_bf16(slot,srcb,BF(W("qkv.w")),S,QKV,DIM,qkv); TB(0); + dyt_slice(qkv,S,QKV,0*DIM,F32(W("qn.alpha"))[0],F32(W("qn.gamma")),F32(W("qn.beta")),par); + dyt_slice(qkv,S,QKV,3*DIM,F32(W("qn.alpha"))[0],F32(W("qn.gamma")),F32(W("qn.beta")),par); + dyt_slice(qkv,S,QKV,1*DIM,F32(W("kn.alpha"))[0],F32(W("kn.gamma")),F32(W("kn.beta")),par); + dyt_slice(qkv,S,QKV,4*DIM,F32(W("kn.alpha"))[0],F32(W("kn.gamma")),F32(W("kn.beta")),par); TB(2); + rope_slice(qkv,S,QKV,0*DIM,par); rope_slice(qkv,S,QKV,1*DIM,par); + rope_slice(qkv,S,QKV,3*DIM,par); rope_slice(qkv,S,QKV,4*DIM,par); TB(3); + diff_attn_banded(qkv,S,ao,par); TB(1); + to_bf16(ao,srcb,(size_t)S*DIM,par); TB(5); + gemm_bf16(slot,srcb,BF(W("out.w")),S,DIM,DIM,h); TB(0); + resadd(x,h,S,par); TB(6); + // ---- FFN: h = ff_norm(x); glu = glu_proj(h); val*act(gate); proj_out; x += . + dyt(x,h,S,DIM,F32(W("ff.alpha"))[0],F32(W("ff.gamma")),F32(W("ff.beta")),par); TB(2); + to_bf16(h,srcb,(size_t)S*DIM,par); TB(5); + gemm_bf16(slot,srcb,BF(W("glu.w")),S,GLU2,DIM,glu); TB(0); + addbias(glu,F32(W("glu.b")),S,GLU2,par); TB(6); + // fuse GLU gate + bf16 requant: h_ff = value*act(gate) -> bf16 for proj GEMM + // blocks 0..4: act=silu; blocks 5..11: act=sin(pi*gate) + #pragma omp parallel for schedule(static) if(par) + for(int m=0;m [16T, DIM] + int L=SIN*T; + float* y=A.h.data(); // reuse h as [16T,DIM] staging (16T*1536 <= S*1536) + #pragma omp parallel for schedule(static) if(par) + for(int t=0;t512). y[L,1536] bf16 -> GEMM -> [L,512] f32, +bias, transpose + to_bf16(y,srcb,(size_t)L*DIM,par); + float* cy=A.qkv.data(); // [L,512] f32 (L*512 <= S*QKV) + gemm_bf16(slot,srcb,BF("map.w"),L,OUTCH,DIM,cy); + addbias(cy,F32("map.b"),L,OUTCH,par); + #pragma omp parallel for schedule(static) if(par) + for(int ch=0;ch out_patches [1,512,16T] f32 (caller-allocated) +void samel_forward(const float* latent,int T,float* out_patches){ + decode(0,true,latent,T,out_patches); +} + +// torch-free unpatch: patches[1,512,L] -> pcm[1,2,256L]. reshape[2,256,L]->transpose->[2,L,256]->[2,256L] +void samel_unpatch(const float* patches,int L,float* pcm){ + #pragma omp parallel for collapse(2) schedule(static) + for(int st=0;st<2;st++)for(int l=0;l0?100*PROF[i]/tot:0); + printf(" total=%.1fms\n",tot*1e3); + for(int i=0;i<8;i++) PROF[i]=0; PROFN=0; fflush(stdout); +} + +} // extern "C" + +// ---- chunk-parallel decode (PRIMARY path) appended below via include ---- +#include "same_l_chunk.inc" diff --git a/optimized/cpu-amx/build/same_l_int8/build.sh b/optimized/cpu-amx/build/same_l_int8/build.sh new file mode 100755 index 00000000..8c5c658f --- /dev/null +++ b/optimized/cpu-amx/build/same_l_int8/build.sh @@ -0,0 +1,8 @@ +#!/bin/bash +set -e +ONE=${ONEDNN_HOME:?set ONEDNN_HOME to a static oneDNN+OpenMP build} +cd "$(dirname "$0")" +g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I"$ONE/include" \ + same_l_int8fused_cpu_amx.cpp -o same_l_int8fused_cpu_amx.so \ + "$ONE/lib/libdnnl.a" -ldl -lpthread -lm +echo "built $(ls -la same_l_int8fused_cpu_amx.so | awk '{print $5}') bytes" diff --git a/optimized/cpu-amx/build/same_l_int8/dump_weights_int8.py b/optimized/cpu-amx/build/same_l_int8/dump_weights_int8.py new file mode 100644 index 00000000..10b368ee --- /dev/null +++ b/optimized/cpu-amx/build/same_l_int8/dump_weights_int8.py @@ -0,0 +1,110 @@ +#!/usr/bin/env python3 +"""Dump SAME-L fp32 npz -> flat {weights.bin + weights_manifest.txt} for the INT8 C++ engine. + +Mirrors ../same_l_cpu_amx/dump_weights.py, but Linears + the plain Linear(1536->512) output +map are stored **w8: per-output-channel symmetric int8** (name.q int8 [K,N] + name.scale f32 [N]), +exactly the DiT/SAME-S-int8 scheme. Runtime: per-row dynamic int8 activation x per-channel int8 +weight -> s8s8->s32 AMX GEMM -> deq. Banded differential attention stays fp32 (Stage 1). + +DyT/biases/new_tokens/running_std: fp32. Optional SmoothQuant (SQ=1) folds per-in-channel s[K] +into `to_qkv` (+ 1/s into pre_norm gamma/beta), read from sq_scales.npz (calibrate_sq.py). +""" + +import os + +# Paths come from the environment so nothing local is baked in. +# SA3_CPUAMX_HOME where the engine dirs live (default ./engines) +# SA3_REPO checkout providing the reference weights to dump +HOME = os.environ.get("SA3_CPUAMX_HOME", os.path.abspath("engines")) +REPO = os.environ.get("SA3_REPO", os.path.abspath(".")) +import os +import numpy as np + +SA3 = REPO +WEIGHTS = os.path.join(SA3, "models", "mlx", "same_l_decoder_f32.npz") +OUT = os.path.join(HOME, "same_l_int8_cpu_amx") +NB = 12 +USE_SQ = os.environ.get("SQ", "0") == "1" +SQ_NPZ = os.path.join(OUT, "sq_scales.npz") + +raw = dict(np.load(WEIGHTS)) +def a(k): + return raw[k].astype(np.float32) + +sq = dict(np.load(SQ_NPZ)) if (USE_SQ and os.path.exists(SQ_NPZ)) else {} +if USE_SQ: + print(f"SmoothQuant ON: {'loaded '+SQ_NPZ if sq else 'NO sq_scales.npz -> plain'}") + +blob = bytearray() +lines = [] + +def put(name, arr, dt): + global blob + off = len(blob) + if dt == "i8": + arr = np.ascontiguousarray(arr.astype(np.int8)); blob += arr.tobytes() + elif dt == "f32": + arr = np.ascontiguousarray(arr.astype(np.float32)); blob += arr.ravel().tobytes() + else: + raise ValueError(dt) + lines.append(f"{name} {dt} {off} {arr.size} " + " ".join(str(s) for s in arr.shape)) + +def quant_w(wt): + amax = np.abs(wt).max(axis=0) + scale = np.maximum(amax / 127.0, 1e-12).astype(np.float32) + q = np.clip(np.round(wt / scale[None, :]), -127, 127).astype(np.int8) + return np.ascontiguousarray(q), np.ascontiguousarray(scale) + +def put_lin_i8(name, w_oi, smooth_s=None): + wt = np.ascontiguousarray(w_oi.T).astype(np.float32) # [K=in, N=out] + if smooth_s is not None: + wt = wt * smooth_s[:, None] + q, scale = quant_w(wt) + put(name + ".q", q, "i8") + put(name + ".scale", scale, "f32") + +# ---- top-level ---- +put("running_std", a("running_std"), "f32") +put_lin_i8("project_in", a("project_in.weight")) # [256,1536] +put("project_in.b", a("project_in.bias"), "f32") # [1536] +put("new_tokens", a("new_tokens").reshape(-1), "f32") # [1536] + +for b in range(NB): + p = f"blocks.{b}" + pg = a(f"{p}.pre_norm.gamma").copy() + pb = a(f"{p}.pre_norm.beta").copy() + s_qkv = sq.get(f"b{b}.qkv.s") if USE_SQ else None + if s_qkv is not None: + pg = pg / s_qkv + pb = pb / s_qkv + put(f"b{b}.pre.alpha", a(f"{p}.pre_norm.alpha"), "f32") + put(f"b{b}.pre.gamma", pg, "f32") + put(f"b{b}.pre.beta", pb, "f32") + put_lin_i8(f"b{b}.qkv", a(f"{p}.attn.to_qkv.weight"), smooth_s=s_qkv) # [1536,7680] + put(f"b{b}.qn.alpha", a(f"{p}.attn.q_norm.alpha"), "f32") + put(f"b{b}.qn.gamma", a(f"{p}.attn.q_norm.gamma"), "f32") # [64] + put(f"b{b}.qn.beta", a(f"{p}.attn.q_norm.beta"), "f32") + put(f"b{b}.kn.alpha", a(f"{p}.attn.k_norm.alpha"), "f32") + put(f"b{b}.kn.gamma", a(f"{p}.attn.k_norm.gamma"), "f32") + put(f"b{b}.kn.beta", a(f"{p}.attn.k_norm.beta"), "f32") + put_lin_i8(f"b{b}.out", a(f"{p}.attn.to_out.weight")) # [1536,1536] + put(f"b{b}.ff.alpha", a(f"{p}.ff_norm.alpha"), "f32") + put(f"b{b}.ff.gamma", a(f"{p}.ff_norm.gamma"), "f32") + put(f"b{b}.ff.beta", a(f"{p}.ff_norm.beta"), "f32") + put_lin_i8(f"b{b}.glu", a(f"{p}.ff.glu_proj.weight")) # [1536,9216] + put(f"b{b}.glu.b", a(f"{p}.ff.glu_proj.bias"), "f32") # [9216] + put_lin_i8(f"b{b}.proj", a(f"{p}.ff.proj_out.weight")) # [4608,1536] (NEVER smoothed) + put(f"b{b}.proj.b", a(f"{p}.ff.proj_out.bias"), "f32") # [1536] + +# output map: plain Linear(1536->512). mapping.weight [512,1536,1] -> [512,1536] -> int8 W.T [1536,512] +put_lin_i8("map", a("mapping.weight").reshape(512, 1536)) # int8 [1536,512] + scale [512] +put("map.b", a("mapping.bias"), "f32") # [512] + +with open(os.path.join(OUT, "weights.bin"), "wb") as f: + f.write(blob) +with open(os.path.join(OUT, "weights_manifest.txt"), "w") as f: + f.write("\n".join(lines) + "\n") + +print(f"wrote weights.bin ({len(blob)/1e6:.1f} MB), {len(lines)} arrays (SQ={'on' if sq else 'off'})") +for l in lines[:4] + ["..."] + lines[-4:]: + print(" ", l) diff --git a/optimized/cpu-amx/build/same_l_int8/same_l_int8fused_chunk.inc b/optimized/cpu-amx/build/same_l_int8/same_l_int8fused_chunk.inc new file mode 100644 index 00000000..e88daca8 --- /dev/null +++ b/optimized/cpu-amx/build/same_l_int8/same_l_int8fused_chunk.inc @@ -0,0 +1,49 @@ +// same_l_int8fused_chunk.inc — chunked / chunk-parallel decode (FUSED int8 SAME-L engine). +// Included at the end of same_l_int8fused_cpu_amx.cpp (shares all statics). Identical tiling to the +// naive int8 engine's chunk .inc (banded receptive field; C=64, overlap=8 ship size); calls the +// FUSED decode(). Quality-neutral vs whole (per-chunk int8 requant only). + +struct ChunkL{int a,b,lo,hi;}; +static std::vector chunk_plan_l(int T,int C,int overlap){ + std::vector plan; int a=0; + while(a plan=chunk_plan_l(T,C,overlap); + int nch=(int)plan.size(); int L=SIN*T; + if(!parallel){ + static std::vector sublat, subpat; + for(auto&ck:plan){ + int w=ck.hi-ck.lo; + sublat.assign((size_t)LAT*w,0); + for(int c=0;c sublat((size_t)LAT*w), subpat((size_t)OUTCH*SIN*w); + for(int c=0;c512) +// map — same int8 weights (reused byte-identically) + same requant grid + same fp32 attention island. +// The ONLY change is the dataflow: it COPIES the DiT fused-epilogue design (dit_cpu_amx.cpp). +// +// naive : dyt->fp32 ; quant->i8 ; GEMM ; deq->fp32 ; addbias ; act ; quant->i8 ; ... (Q/DQ per GEMM) +// FUSED : dyt_q (norm+quant->i8 ONE pass) ; GEMM ; deqglu_q (deq+bias+GLU+quant->i8 ONE pass) ; +// GEMM ; deq_res (deq+bias+residual ONE pass). Activation STAYS int8 between GLU & proj. +// +// Only standalone Q/DQ left = the 2 around the fp32 attention island. Residual stream x stays fp32 +// (exactly like the DiT xnext / the bf16 engine) -> requant grid IDENTICAL to naive -> quality matches. +// +// build: g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I$ONEINC \ +// same_l_int8fused_cpu_amx.cpp -o same_l_int8fused_cpu_amx.so \ +// -L$ONELIB -ldnnl -Wl,-rpath,$ONELIB -ldl -lpthread -lm +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +// ── architecture constants (same_l_decoder_torch.py) ── +static const int LAT=256, DIM=1536, H=24, HD=64, RD=32, HALF=16; +static const int NB=12, FF=4608, GLU2=9216, QKV=7680, OUTCH=512; +static const int SUB=17, SIN=16; +static const int SIN_START=5; // blocks >=5 use sin(pi*gate) gate +static const int BAND=17; +static const float SCALE=0.125f; + +// ── optional phase profiler (SAMEL_PROF=1) ── +static double PROF[8]={0}; static long PROFN=0; +static const char* PROFLBL[8]={"gemm","attn","dyt_q","rope","deqglu","quant","deq_res",""}; +static bool PROF_ON=false; +static inline double wt(){return omp_get_wtime();} +#define TB(id) do{ if(PROF_ON){double _n=wt(); PROF[id]+=_n-_pt; _pt=_n;} }while(0) + +// ── fast vectorizable transcendentals ── +static inline float vexp(float x){ + x = x<-87.0f?-87.0f:(x>88.0f?88.0f:x); + float z = x*1.442695041f; + float n = std::floor(z+0.5f); + float f = z-n; + float p = 1.0f+f*(0.6931472f+f*(0.2402265f+f*(0.0555041f+f*(0.0096181f+f*0.0013333f)))); + uint32_t bits=(uint32_t)(((int)n+127)<<23); float s; std::memcpy(&s,&bits,4); + return p*s; +} +static inline float vtanh(float x){ return 1.0f-2.0f/(vexp(2.0f*x)+1.0f); } +static inline float vsilu(float x){ return x/(1.0f+vexp(-x)); } +static inline float vsinpi(float g){ + float k = std::floor(g+0.5f); + float r = g - k; + float s = 1.0f - 2.0f*(float)(((long)k)&1L); + float y = 3.14159265358979f*r; + float y2 = y*y; + float p = y*(1.0f + y2*(-0.16666667f + y2*(0.00833333f + y2*(-0.00019841f + y2*2.75573e-6f)))); + return s*p; +} +static inline int cdiv(int a,int b){return (a+b-1)/b;} +static inline int8_t q127(float v){ v=std::rintf(v); return (int8_t)(v>127.0f?127.0f:(v<-127.0f?-127.0f:v)); } + +// ------------------------- mmap weights.bin + manifest ------------------------- +struct Ten{void* p; std::string dt; long n; std::vector shp;}; +static std::map TEN; +static char* BASE=nullptr; + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string WBASE = std::string(sa3_home()) + "/same_l_int8fused_cpu_amx/weights"; +static void load_weights(){ + std::string bin=WBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(WBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("[samel_i8fused] weights mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static int8_t* I8 (const std::string&k){return (int8_t*)TEN.at(k).p;} + +// ------------------------- oneDNN int8 matmul (s8 x s8 -> s32): primitive+handle cache +static dnnl::engine* ENG=nullptr; +static void onednn_init(){ ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); } +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M mm; dnnl::stream* strm=nullptr; }; +static std::vector MMC; +static void mmcache_init(int nworkers){ + MMC.clear(); MMC.resize(std::max(1,nworkers)+1); + for(auto& c:MMC) c.strm=new dnnl::stream(*ENG); +} +static void gemm_i8(int slot,const int8_t*A,const int8_t*B,int M,int N,int K,int32_t* C){ + using dt=dnnl::memory::data_type; + MMCache& c=MMC[slot]; + MMKey key{M,N,K}; auto it=c.mm.find(key); + if(it==c.mm.end()){ + #pragma omp critical(mmcreate) + { + it=c.mm.find(key); + if(it==c.mm.end()){ + dnnl::memory::desc a_md({M,K},dt::s8,{K,1}); + dnnl::memory::desc b_md({K,N},dt::s8,{N,1}); + dnnl::memory::desc c_md({M,N},dt::s32,{N,1}); + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)A), + dnnl::memory(pd.weights_desc(),*ENG,(void*)B), + dnnl::memory(pd.dst_desc(),*ENG,(void*)C)}; + it=c.mm.emplace(key,std::move(e)).first; + } + } + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)A); e.bm.set_data_handle((void*)B); e.cm.set_data_handle((void*)C); + e.prim.execute(*c.strm,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + c.strm->wait(); +} + +// ------------------------- GLOBAL RoPE table (growable) ------------------------- +static float RINV[HALF]; +static std::vector RCOS, RSIN; +static int RCAP=0; +static void rope_invfreq(){ for(int i=0;iRCAP){ int nc=S+S/4; RCOS.resize((size_t)nc*HALF); RSIN.resize((size_t)nc*HALF); + rope_fill(RCAP,nc); RCAP=nc; } } +} + +// ------------------------- reusable arena (per worker slot) ------------------------- +// vs naive: dropped the [S,GLU2] fp32 `glu` buffer (deqglu_q reads int32 acc -> writes int8). +struct Arena{ + int cap=0; + std::vector x,h,qkv,ao; // f32 activations (residual stream stays fp32) + std::vector srcq; // int8 GEMM-src staging (>= max K = FF) + std::vector acc; // int32 GEMM accumulator (>= max N = GLU2) + std::vector ascl; // per-row activation scale [S] + void ensure(int M){ + if(M<=cap) return; cap=M; + x.assign((size_t)M*DIM,0); h.assign((size_t)M*DIM,0); + qkv.assign((size_t)M*QKV,0); ao.assign((size_t)M*DIM,0); + srcq.assign((size_t)M*GLU2,0); // headroom (widest src is FF AR; + +// ------------------------- fp32 attention-prep elementwise (identical to bf16/naive engine) ------------------------- +static void dyt_slice(float*base,int M,int W,int coloff,float alpha,const float*g,const float*b,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mamax)amax=av; } + float sc=amax>0.0f? amax/127.0f : 1e-12f; s[m]=sc; float inv=1.0f/sc; + #pragma omp simd + for(int k=0;kamax)amax=au; } + float sc=amax>0.0f? amax/127.0f : 1e-12f; s[m]=sc; float inv=1.0f/sc; + #pragma omp simd + for(int j=0;j=5). Replaces naive deq+addbias+glu(->fp32)+quant. +static void deqglu_q(const int32_t*acc,const float*as,const float*bs,const float*bias, + int M,int Fdim,bool use_sin,int8_t* q,float* s,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mamax)amax=au; + } + }else{ + #pragma omp simd reduction(max:amax) + for(int j=0;jamax)amax=au; + } + } + float sc=amax>0.0f? amax/127.0f : 1e-12f; s[m]=sc; float inv=1.0f/sc; + #pragma omp simd + for(int j=0;j G; + G.resize((size_t)5*S*HD); + float* qg=G.data(); float* kg=qg+(size_t)S*HD; float* vg=kg+(size_t)S*HD; + float* qdg=vg+(size_t)S*HD; float* kdg=qdg+(size_t)S*HD; + for(int t=0;tmm)mm=dm; if(dd>md)md=dd; + } + int n=j1-j0+1; float zm=0,zd=0; + #pragma omp simd reduction(+:zm,zd) + for(int j=0;j patches[512,16T]) --------- +static void decode(int slot,bool par,const float* latent,int T,float* out_patches){ + Arena& A=AR[slot]; + int S=SUB*T; // internal tokens (17T) + A.ensure(S); + ensure_rope(S); + float* x=A.x.data(); float* h=A.h.data(); + float* qkv=A.qkv.data(); float* ao=A.ao.data(); + int8_t* srcq=A.srcq.data(); float* ascl=A.ascl.data(); int32_t* acc=A.acc.data(); + float rstd=F32("running_std")[0]; + + // project_in (w8a8): f32 (latent^T*running_std) rows -> per-row quant -> GEMM -> deq+bias -> h[T,1536] + #pragma omp parallel for schedule(static) if(par) + for(int t=0;tamax)amax=av; } + float sc=amax>0.0f? amax/127.0f : 1e-12f; ascl[t]=sc; float inv=1.0f/sc; + int8_t* qo=srcq+(size_t)t*LAT; + for(int c=0;c=SIN_START); + char pb[8]; snprintf(pb,sizeof pb,"b%d.",b); + auto W=[&](const char*n){return std::string(pb)+n;}; + // ---- attention ---- + dyt_q(x,S,DIM,F32(W("pre.alpha"))[0],F32(W("pre.gamma")),F32(W("pre.beta")),srcq,ascl,par); TB(2); + gemm_i8(slot,srcq,I8(W("qkv.q")),S,QKV,DIM,acc); TB(0); + deq(acc,ascl,F32(W("qkv.scale")),S,QKV,qkv,par); TB(6); // the ONE dequant into the fp32 island + dyt_slice(qkv,S,QKV,0*DIM,F32(W("qn.alpha"))[0],F32(W("qn.gamma")),F32(W("qn.beta")),par); + dyt_slice(qkv,S,QKV,3*DIM,F32(W("qn.alpha"))[0],F32(W("qn.gamma")),F32(W("qn.beta")),par); + dyt_slice(qkv,S,QKV,1*DIM,F32(W("kn.alpha"))[0],F32(W("kn.gamma")),F32(W("kn.beta")),par); + dyt_slice(qkv,S,QKV,4*DIM,F32(W("kn.alpha"))[0],F32(W("kn.gamma")),F32(W("kn.beta")),par); TB(2); + rope_slice(qkv,S,QKV,0*DIM,par); rope_slice(qkv,S,QKV,1*DIM,par); + rope_slice(qkv,S,QKV,3*DIM,par); rope_slice(qkv,S,QKV,4*DIM,par); TB(3); + diff_attn_banded(qkv,S,ao,par); TB(1); // fp32 attention island + quant_rows(ao,S,DIM,srcq,ascl,par); TB(5); // the ONE quant out of the island + gemm_i8(slot,srcq,I8(W("out.q")),S,DIM,DIM,acc); TB(0); + deq_res(acc,ascl,F32(W("out.scale")),nullptr,x,S,DIM,par); TB(6); // FUSED deq + residual (x += .) + // ---- FFN ---- + dyt_q(x,S,DIM,F32(W("ff.alpha"))[0],F32(W("ff.gamma")),F32(W("ff.beta")),srcq,ascl,par); TB(2); + gemm_i8(slot,srcq,I8(W("glu.q")),S,GLU2,DIM,acc); TB(0); + deqglu_q(acc,ascl,F32(W("glu.scale")),F32(W("glu.b")),S,FF,use_sin,srcq,ascl,par); TB(4); // deq+bias+GLU+quant->i8 (STAYS i8) + gemm_i8(slot,srcq,I8(W("proj.q")),S,DIM,FF,acc); TB(0); // int8 GEMM directly on int8 + deq_res(acc,ascl,F32(W("proj.scale")),F32(W("proj.b")),x,S,DIM,par); TB(6); // FUSED deq + bias + residual + }; + + for(int b=0;b y[16T,DIM] + int L=SIN*T; + float* y=A.h.data(); + #pragma omp parallel for schedule(static) if(par) + for(int t=0;t512) w8a8 -> [L,512] ; fused deq+bias+transpose -> out_patches[512,L] + quant_rows(y,L,DIM,srcq,ascl,par); + gemm_i8(slot,srcq,I8("map.q"),L,OUTCH,DIM,acc); + { const float* bs=F32("map.scale"); const float* bias=F32("map.b"); + #pragma omp parallel for schedule(static) if(par) + for(int l=0;l0?100*PROF[i]/tot:0); + printf(" total=%.1fms\n",tot*1e3); + for(int i=0;i<8;i++) PROF[i]=0; PROFN=0; fflush(stdout); +} + +} // extern "C" + +// ---- chunk-parallel decode (PRIMARY path) appended below via include ---- +#include "same_l_int8fused_chunk.inc" diff --git a/optimized/cpu-amx/build/same_s_bf16/dump_weights.py b/optimized/cpu-amx/build/same_s_bf16/dump_weights.py new file mode 100644 index 00000000..10f7d15d --- /dev/null +++ b/optimized/cpu-amx/build/same_s_bf16/dump_weights.py @@ -0,0 +1,100 @@ +#!/usr/bin/env python3 +"""Dump SAME-S fp32 npz -> flat {weights.bin + weights_manifest.txt} for the C++ engine. + +Layout choices (match model_bf16.py SHIP config exactly): + * Linears: stored **bf16** in oneDNN matmul weight layout [K=in, N=out] (= W.T of the + npz [out,in]). AMX-BF16 GEMM src[M,K]bf16 x wei[K,N]bf16 -> dst[M,N]f32. + * Conv (WNConv1d 768->512 k=3, weight_norm PRE-FUSED in the npz): stored bf16 as + [K=768*3, N=512] = conv_w.permute(1,2,0).reshape(2304,512) — exactly model_bf16's + im2col weight `cw`, whose col index is in*3+k (in-major, k-minor). + * DyT (alpha scalar, gamma/beta), biases, new_tokens, running_std: fp32 (the + cancellation-fragile elementwise stays fp32). + +bf16 conversion uses torch's round-to-nearest-even (build-time tool; the *runtime* is +torch-free). Manifest line: name dtype byte_offset nelem d0 d1 ... +""" + +import os + +# Paths come from the environment so nothing local is baked in. +# SA3_CPUAMX_HOME where the engine dirs live (default ./engines) +# SA3_REPO checkout providing the reference weights to dump +HOME = os.environ.get("SA3_CPUAMX_HOME", os.path.abspath("engines")) +REPO = os.environ.get("SA3_REPO", os.path.abspath(".")) +import os, sys +import numpy as np +import torch + +SA3 = REPO +WEIGHTS = os.path.join(SA3, "models", "mlx", "same_s_decoder_f32.npz") +OUT = os.path.join(HOME, "same_s_cpu_amx") +NB = 6 + +raw = dict(np.load(WEIGHTS)) +def a(k): # fp32 numpy + return raw[k].astype(np.float32) + +blob = bytearray() +lines = [] + +def put(name, arr, dt): + """dt in {'bf16','f32'}. arr is numpy fp32 (any shape); stored C-contiguous.""" + global blob + arr = np.ascontiguousarray(arr.astype(np.float32)) + off = len(blob) + if dt == "bf16": + t = torch.from_numpy(arr).to(torch.bfloat16) + u = t.view(torch.uint16).numpy().ravel() + blob += u.tobytes() + elif dt == "f32": + blob += arr.ravel().tobytes() + else: + raise ValueError(dt) + shp = " ".join(str(s) for s in arr.shape) + lines.append(f"{name} {dt} {off} {arr.size} {shp}") + +def put_lin(name, w_oi): + """w_oi = npz [out,in]; store bf16 W.T = [in,out] for oneDNN wei[K,N].""" + put(name, np.ascontiguousarray(w_oi.T), "bf16") + +# ---- top-level ---- +put("running_std", a("running_std"), "f32") +put_lin("project_in.w", a("project_in.weight")) # [256,768] +put("project_in.b", a("project_in.bias"), "f32") # [768] +put("new_tokens", a("new_tokens").reshape(-1), "f32") # [768] + +for b in range(NB): + p = f"blocks.{b}" + put(f"b{b}.pre.alpha", a(f"{p}.pre_norm.alpha"), "f32") + put(f"b{b}.pre.gamma", a(f"{p}.pre_norm.gamma"), "f32") + put(f"b{b}.pre.beta", a(f"{p}.pre_norm.beta"), "f32") + put_lin(f"b{b}.qkv.w", a(f"{p}.attn.to_qkv.weight")) # [768,3840] + put(f"b{b}.qn.alpha", a(f"{p}.attn.q_norm.alpha"), "f32") + put(f"b{b}.qn.gamma", a(f"{p}.attn.q_norm.gamma"), "f32") # [64] + put(f"b{b}.qn.beta", a(f"{p}.attn.q_norm.beta"), "f32") + put(f"b{b}.kn.alpha", a(f"{p}.attn.k_norm.alpha"), "f32") + put(f"b{b}.kn.gamma", a(f"{p}.attn.k_norm.gamma"), "f32") + put(f"b{b}.kn.beta", a(f"{p}.attn.k_norm.beta"), "f32") + put_lin(f"b{b}.out.w", a(f"{p}.attn.to_out.weight")) # [768,768] + put(f"b{b}.ff.alpha", a(f"{p}.ff_norm.alpha"), "f32") + put(f"b{b}.ff.gamma", a(f"{p}.ff_norm.gamma"), "f32") + put(f"b{b}.ff.beta", a(f"{p}.ff_norm.beta"), "f32") + put_lin(f"b{b}.glu.w", a(f"{p}.ff.glu_proj.weight")) # [768,4608] + put(f"b{b}.glu.b", a(f"{p}.ff.glu_proj.bias"), "f32") # [4608] + put_lin(f"b{b}.proj.w", a(f"{p}.ff.proj_out.weight")) # [2304,768] + put(f"b{b}.proj.b", a(f"{p}.ff.proj_out.bias"), "f32") # [768] + +# conv: [512,768,3]=[out,in,k] -> [in,k,out]=[768,3,512] -> [2304,512] (row=in*3+k) +cw = np.ascontiguousarray(a("mapping.weight").transpose(1, 2, 0).reshape(768 * 3, 512)) +put("conv.w", cw, "bf16") # already [K=2304, N=512] +put("conv.b", a("mapping.bias"), "f32") # [512] + +with open(os.path.join(OUT, "weights.bin"), "wb") as f: + f.write(blob) +with open(os.path.join(OUT, "weights_manifest.txt"), "w") as f: + f.write("\n".join(lines) + "\n") + +print(f"wrote weights.bin ({len(blob)/1e6:.1f} MB), {len(lines)} arrays") +print("first/last few manifest lines:") +for l in lines[:4] + ["..."] + lines[-4:]: + print(" ", l) diff --git a/optimized/cpu-amx/build/same_s_bf16/same_s_chunk.inc b/optimized/cpu-amx/build/same_s_bf16/same_s_chunk.inc new file mode 100644 index 00000000..e4c0415d --- /dev/null +++ b/optimized/cpu-amx/build/same_s_bf16/same_s_chunk.inc @@ -0,0 +1,57 @@ +// same_s_chunk.inc — cache-blocked / chunk-parallel decode (milestone C). +// Included at the end of same_s_cpu_amx.cpp (shares all statics). +// +// SAME-S decode is position-invariant with receptive field = 2 latent tokens, so tiling +// [1,256,T] into C-token interiors with overlap=2 context each side and concatenating the +// interior audio reproduces the whole decode (bit-exact in fp32; quality-neutral in bf16). + +struct Chunk{int a,b,lo,hi;}; +static std::vector chunk_plan(int T,int C,int overlap){ + std::vector plan; int a=0; + while(a plan=chunk_plan(T,C,overlap); + int nch=(int)plan.size(); int L=SIN*T; + if(!parallel){ + // sequential: reuse slot 0 (16-thread), one sub scratch pair + static std::vector sublat, subpat; + for(auto&ck:plan){ + int w=ck.hi-ck.lo; + sublat.assign((size_t)LAT*w,0); + for(int c=0;c sublat((size_t)LAT*w), subpat((size_t)OUTCH*SIN*w); + for(int c=0;c keep fp32). +// +// sames_init(weights_base, threads) -> mmap bf16 weights ONCE + AMX/omp/oneDNN +// sames_forward(latent[1,256,T], T, out_patches[1,512,16T]) -> whole decode +// sames_forward_chunked(latent, T, C, overlap, parallel, out) -> cache-blocked decode +// sames_unpatch(patches[1,512,L], L, pcm[1,2,256L]) -> torch-free unpatch +// +// build: g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I$ONEINC \ +// same_s_cpu_amx.cpp -o same_s_cpu_amx.so $ONELIB/libdnnl.a -ldl -lpthread -lm +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +// ── architecture constants (same_s_decoder_torch.py) ── +static const int LAT=256, DIM=768, H=12, HD=64, RD=32, HALF=16; +static const int NB=6, FF=2304, GLU2=4608, QKV=3840, OUTCH=512; +static const int SUB=17, ECH=34, SHIFT=17, SIN=16; // 17 tok/latent, 34-tok chunk +static const float SCALE=0.125f; // HD**-0.5 = 64**-0.5 + +// ── optional phase profiler (SAMES_PROF=1) ── +static double PROF[8]={0}; static long PROFN=0; +static const char* PROFLBL[8]={"gemm","attn","dyt","rope","glu","cast","misc",""}; +static bool PROF_ON=false; +static inline double wt(){return omp_get_wtime();} +#define TB(id) do{ if(PROF_ON){double _n=wt(); PROF[id]+=_n-_pt; _pt=_n;} }while(0) + +typedef uint16_t bf16; +static inline bf16 f2b(float f){ // round-to-nearest-even f32->bf16 + uint32_t x; std::memcpy(&x,&f,4); + uint32_t r=x+0x7fff+((x>>16)&1); return (bf16)(r>>16); +} +static inline float b2f(bf16 h){ uint32_t x=(uint32_t)h<<16; float f; std::memcpy(&f,&x,4); return f; } + +// ── fast vectorizable transcendentals (pure float, no libm/branches -> AVX-512 auto-vec) ── +// exp: 2^f degree-5 minimax on [-0.5,0.5] + ldexp via exponent bits. ~2e-7 rel error. +static inline float vexp(float x){ + x = x<-87.0f?-87.0f:(x>88.0f?88.0f:x); + float z = x*1.442695041f; + float n = std::floor(z+0.5f); + float f = z-n; + float p = 1.0f+f*(0.6931472f+f*(0.2402265f+f*(0.0555041f+f*(0.0096181f+f*0.0013333f)))); + uint32_t bits=(uint32_t)(((int)n+127)<<23); float s; std::memcpy(&s,&bits,4); + return p*s; +} +static inline float vtanh(float x){ return 1.0f-2.0f/(vexp(2.0f*x)+1.0f); } // == Triton DyT kernel +static inline float vsilu(float x){ return x/(1.0f+vexp(-x)); } +static inline int cdiv(int a,int b){return (a+b-1)/b;} + +// ------------------------- mmap weights.bin + manifest ------------------------- +struct Ten{void* p; std::string dt; long n; std::vector shp;}; +static std::map TEN; +static char* BASE=nullptr; + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string WBASE = std::string(sa3_home()) + "/same_s_cpu_amx/weights"; +static void load_weights(){ + std::string bin=WBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(WBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("[sames] weights mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static bf16* BF (const std::string&k){return (bf16*)TEN.at(k).p;} + +// ------------------------- oneDNN bf16 matmul (bf16 x bf16 -> f32): primitive+handle cache +static dnnl::engine* ENG=nullptr; +static void onednn_init(){ ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); } +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M0 = per-worker for chunk-parallel. +struct MMCache{ std::map mm; dnnl::stream* strm=nullptr; }; +static std::vector MMC; +static void mmcache_init(int nworkers){ + MMC.clear(); MMC.resize(std::max(1,nworkers)+1); + for(auto& c:MMC) c.strm=new dnnl::stream(*ENG); +} +// src A[M,K] bf16, wei B[K,N] bf16 (row-major), dst C[M,N] f32. cache slot `slot`. +static void gemm_bf16(int slot,const bf16*A,const bf16*B,int M,int N,int K,float* C){ + using dt=dnnl::memory::data_type; + MMCache& c=MMC[slot]; + MMKey key{M,N,K}; auto it=c.mm.find(key); + if(it==c.mm.end()){ + // primitive creation is rare (cached per shape); guard it so concurrent chunk + // workers can build primitives from the shared engine safely. Execution stays lock-free. + #pragma omp critical(mmcreate) + { + it=c.mm.find(key); + if(it==c.mm.end()){ + dnnl::memory::desc a_md({M,K},dt::bf16,{K,1}); + dnnl::memory::desc b_md({K,N},dt::bf16,{N,1}); + dnnl::memory::desc c_md({M,N},dt::f32, {N,1}); + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)A), + dnnl::memory(pd.weights_desc(),*ENG,(void*)B), + dnnl::memory(pd.dst_desc(),*ENG,(void*)C)}; + it=c.mm.emplace(key,std::move(e)).first; + } + } + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)A); e.bm.set_data_handle((void*)B); e.cm.set_data_handle((void*)C); + e.prim.execute(*c.strm,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + c.strm->wait(); +} + +// ------------------------- RoPE table for one 34-token chunk (positions 0..33) ------------------------- +static float RCOS[ECH*HALF], RSIN[ECH*HALF]; // [pos, i] i=0..15 +static void build_rope(){ + for(int i=0;i0 = chunk-parallel). +struct Arena{ + int cap=0; // token capacity (max M) + std::vector x,xt,h,qkv,ao,glu; // f32 activations + std::vector srcb; // bf16 GEMM src staging (max K=2304) + void ensure(int M){ + if(M<=cap) return; cap=M; + x.assign((size_t)M*DIM,0); xt.assign((size_t)M*DIM,0); h.assign((size_t)M*DIM,0); + qkv.assign((size_t)M*QKV,0); ao.assign((size_t)M*DIM,0); glu.assign((size_t)M*GLU2,0); + srcb.assign((size_t)M*GLU2,0); // big enough for any src (K<=2304) and glu val stage + } +}; +static std::vector AR; + +// ------------------------- fp32 elementwise kernels (OMP over slot's own thread budget) ------------------------- +// par=true: kernels use `#pragma omp parallel for` (slot 0, all threads). par=false: run serial +// (already inside an outer chunk-parallel region -> no nested fork; oneDNN also serializes there). + +// DyT: out[m,j] = gamma[j]*tanh(alpha*x[m,j]) + beta[j] (full-DIM norm, gk==K) +static void dyt(const float*x,float*o,int M,int K,float alpha,const float*g,const float*b,int gk,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;m bf16 (M*K) +static void to_bf16(const float*x,bf16*o,size_t n,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(size_t i=0;imm)mm=dm; if(dd>md)md=dd; + } + float zm=0,zd=0; + #pragma omp simd reduction(+:zm,zd) + for(int j=0;j patches[512,16T]) --------- +// slot: arena/matmul cache index (0 = 16-thread; >0 = a chunk-parallel worker running serial). +// par: true -> kernels use omp-for (slot 0); false -> serial (inside outer chunk-parallel). +static void decode(int slot,bool par,const float* latent,int T,float* out_patches){ + Arena& A=AR[slot]; + int iT=SUB*T; // internal tokens after new-token expansion (17T) + int M2=iT+ECH; // padded token count for the shifted 2nd half + A.ensure(M2); + float* x=A.x.data(); float* xt=A.xt.data(); float* h=A.h.data(); + float* qkv=A.qkv.data(); float* ao=A.ao.data(); float* glu=A.glu.data(); + bf16* srcb=A.srcb.data(); + float rstd=F32("running_std")[0]; + + // project_in: build src [T,256] bf16 = (latent^T * running_std) rows, then GEMM -> [T,768] + // latent is [1,256,T] channel-major: latent[c*T + t] + #pragma omp parallel for schedule(static) if(par) + for(int t=0;t bf16 staging for proj GEMM + #pragma omp parallel for schedule(static) if(par) + for(int m=0;m [16T, DIM] + int L=SIN*T; + float* y=A.h.data(); // reuse h as [16T,DIM] staging (16T*768 <= M2*768) + #pragma omp parallel for schedule(static) if(par) + for(int t=0;t512 k3 pad1 via im2col: cols[L, 2304] bf16 (col index c*3+k), GEMM -> [L,512], +bias, transpose + bf16* cols=srcb; // [L, 2304] bf16 (L*2304 <= M2*GLU2) + #pragma omp parallel for schedule(static) if(par) + for(int l=0;l=0&&j patches[512,L] (out_patches = [1,512,L], ch-major) + #pragma omp parallel for schedule(static) if(par) + for(int ch=0;ch out_patches [1,512,16T] f32 (caller-allocated) +void sames_forward(const float* latent,int T,float* out_patches){ + decode(0,true,latent,T,out_patches); +} + +// torch-free unpatch: patches[1,512,L] -> pcm[1,2,256L]. reshape[2,256,L]->transpose->[2,L,256]->[2,256L] +void sames_unpatch(const float* patches,int L,float* pcm){ + #pragma omp parallel for collapse(2) schedule(static) + for(int st=0;st<2;st++)for(int l=0;l0?100*PROF[i]/tot:0); + printf(" total=%.1fms\n",tot*1e3); + for(int i=0;i<8;i++) PROF[i]=0; PROFN=0; fflush(stdout); +} + +} // extern "C" + +// ---- chunk-parallel decode is appended in a second translation section below via include ---- +#include "same_s_chunk.inc" diff --git a/optimized/cpu-amx/build/same_s_int8/build.sh b/optimized/cpu-amx/build/same_s_int8/build.sh new file mode 100755 index 00000000..0056b00e --- /dev/null +++ b/optimized/cpu-amx/build/same_s_int8/build.sh @@ -0,0 +1,10 @@ +#!/bin/bash +# Build the FUSED int8 SAME-S engine, STATIC-linked against the threaded OMP oneDNN (clean ldd, +# same threaded GEMM as the naive/bf16 engines). ONE=... points at the OMP static oneDNN build. +set -e +ONE=${ONEDNN_HOME:?set ONEDNN_HOME to a static oneDNN+OpenMP build} +cd "$(dirname "$0")" +g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I"$ONE/include" \ + same_s_int8fused_cpu_amx.cpp -o same_s_int8fused_cpu_amx.so \ + "$ONE/lib/libdnnl.a" -ldl -lpthread -lm +echo "built $(ls -la same_s_int8fused_cpu_amx.so | awk '{print $5}') bytes" diff --git a/optimized/cpu-amx/build/same_s_int8/dump_weights_int8.py b/optimized/cpu-amx/build/same_s_int8/dump_weights_int8.py new file mode 100644 index 00000000..4cca8c18 --- /dev/null +++ b/optimized/cpu-amx/build/same_s_int8/dump_weights_int8.py @@ -0,0 +1,129 @@ +#!/usr/bin/env python3 +"""Dump SAME-S fp32 npz -> flat {weights.bin + weights_manifest.txt} for the INT8 C++ engine. + +Mirrors ../same_s_cpu_amx/dump_weights.py, but the Linears (and the WNConv1d) are stored +**w8: per-output-channel symmetric int8** (exactly the DiT engine's scheme, model.py W): + wt = W.T = [K=in, N=out] + scale[N] = max(|wt|.amax(axis=0) / 127, 1e-12) # per out-channel + q[K,N] = clip(round(wt / scale[None,:]), -127, 127).int8 +Runtime: activation is quantized per-row (dynamic symmetric int8) and s8s8->s32 AMX GEMM, +then dequant o[m,n] = acc[m,n]*a_scale[m]*w_scale[n]. Attention stays fp32 (Stage 1). + +DyT (alpha/gamma/beta), biases, new_tokens, running_std: fp32 (cancellation-fragile). + +Optional SmoothQuant (SQ=1 env): fold per-input-channel smoothing s[K] into `to_qkv` weight +(W_hat[k,n]=s[k]*W[k,n]) and DIVIDE it out of the activation by folding 1/s[k] into the +pre_norm gamma/beta that produce the to_qkv input (h=pre_norm(x) feeds ONLY to_qkv, so this is +exact + runtime-free). s[k] read from sq_scales.npz (built by calibrate_sq.py). proj_out is +NEVER smoothed (214x outliers -> backfires, per smoothquant_test/SMOOTHQUANT_W8A8.md). + +Manifest line: name dtype byte_offset nelem d0 d1 ... (dtype in {i8,f32}) +""" + +import os + +# Paths come from the environment so nothing local is baked in. +# SA3_CPUAMX_HOME where the engine dirs live (default ./engines) +# SA3_REPO checkout providing the reference weights to dump +HOME = os.environ.get("SA3_CPUAMX_HOME", os.path.abspath("engines")) +REPO = os.environ.get("SA3_REPO", os.path.abspath(".")) +import os +import numpy as np + +SA3 = REPO +WEIGHTS = os.path.join(SA3, "models", "mlx", "same_s_decoder_f32.npz") +OUT = os.path.join(HOME, "same_s_int8_cpu_amx") +NB = 6 +USE_SQ = os.environ.get("SQ", "0") == "1" +SQ_NPZ = os.path.join(OUT, "sq_scales.npz") + +raw = dict(np.load(WEIGHTS)) +def a(k): + return raw[k].astype(np.float32) + +sq = dict(np.load(SQ_NPZ)) if (USE_SQ and os.path.exists(SQ_NPZ)) else {} +if USE_SQ: + print(f"SmoothQuant ON: {'loaded '+SQ_NPZ if sq else 'NO sq_scales.npz -> plain'}") + +blob = bytearray() +lines = [] + +def put(name, arr, dt): + global blob + off = len(blob) + if dt == "i8": + arr = np.ascontiguousarray(arr.astype(np.int8)) + blob += arr.tobytes() + elif dt == "f32": + arr = np.ascontiguousarray(arr.astype(np.float32)) + blob += arr.ravel().tobytes() + else: + raise ValueError(dt) + shp = " ".join(str(s) for s in arr.shape) + lines.append(f"{name} {dt} {off} {arr.size} {shp}") + +def quant_w(wt): + """wt [K,N] fp32 -> (q int8 [K,N], scale f32 [N]) per out-channel symmetric.""" + amax = np.abs(wt).max(axis=0) + scale = np.maximum(amax / 127.0, 1e-12).astype(np.float32) + q = np.clip(np.round(wt / scale[None, :]), -127, 127).astype(np.int8) + return np.ascontiguousarray(q), np.ascontiguousarray(scale) + +def put_lin_i8(name, w_oi, smooth_s=None): + """w_oi = npz [out,in]; store int8 W.T=[in,out] per-out-channel + f32 scale. + smooth_s: optional per-input-channel s[K] to fold (W_hat[k,n]=s[k]*W[k,n]).""" + wt = np.ascontiguousarray(w_oi.T).astype(np.float32) # [K=in, N=out] + if smooth_s is not None: + wt = wt * smooth_s[:, None] + q, scale = quant_w(wt) + put(name + ".q", q, "i8") + put(name + ".scale", scale, "f32") + +# ---- top-level ---- +put("running_std", a("running_std"), "f32") +put_lin_i8("project_in", a("project_in.weight")) # [256,768] +put("project_in.b", a("project_in.bias"), "f32") # [768] +put("new_tokens", a("new_tokens").reshape(-1), "f32") # [768] + +for b in range(NB): + p = f"blocks.{b}" + pg = a(f"{p}.pre_norm.gamma").copy() + pb = a(f"{p}.pre_norm.beta").copy() + s_qkv = sq.get(f"b{b}.qkv.s") if USE_SQ else None + if s_qkv is not None: + pg = pg / s_qkv # h' = h/s -> fold 1/s into gamma,beta (h feeds ONLY to_qkv) + pb = pb / s_qkv + put(f"b{b}.pre.alpha", a(f"{p}.pre_norm.alpha"), "f32") + put(f"b{b}.pre.gamma", pg, "f32") + put(f"b{b}.pre.beta", pb, "f32") + put_lin_i8(f"b{b}.qkv", a(f"{p}.attn.to_qkv.weight"), smooth_s=s_qkv) # [768,3840] + put(f"b{b}.qn.alpha", a(f"{p}.attn.q_norm.alpha"), "f32") + put(f"b{b}.qn.gamma", a(f"{p}.attn.q_norm.gamma"), "f32") + put(f"b{b}.qn.beta", a(f"{p}.attn.q_norm.beta"), "f32") + put(f"b{b}.kn.alpha", a(f"{p}.attn.k_norm.alpha"), "f32") + put(f"b{b}.kn.gamma", a(f"{p}.attn.k_norm.gamma"), "f32") + put(f"b{b}.kn.beta", a(f"{p}.attn.k_norm.beta"), "f32") + put_lin_i8(f"b{b}.out", a(f"{p}.attn.to_out.weight")) # [768,768] + put(f"b{b}.ff.alpha", a(f"{p}.ff_norm.alpha"), "f32") + put(f"b{b}.ff.gamma", a(f"{p}.ff_norm.gamma"), "f32") + put(f"b{b}.ff.beta", a(f"{p}.ff_norm.beta"), "f32") + put_lin_i8(f"b{b}.glu", a(f"{p}.ff.glu_proj.weight")) # [768,4608] + put(f"b{b}.glu.b", a(f"{p}.ff.glu_proj.bias"), "f32") # [4608] + put_lin_i8(f"b{b}.proj", a(f"{p}.ff.proj_out.weight")) # [2304,768] (NEVER smoothed) + put(f"b{b}.proj.b", a(f"{p}.ff.proj_out.bias"), "f32") # [768] + +# conv: [512,768,3]=[out,in,k] -> [in,k,out]=[768,3,512] -> [2304,512] (row=in*3+k) +cw = np.ascontiguousarray(a("mapping.weight").transpose(1, 2, 0).reshape(768 * 3, 512)) +q, scale = quant_w(cw) +put("conv.q", q, "i8") +put("conv.scale", scale, "f32") +put("conv.b", a("mapping.bias"), "f32") # [512] + +with open(os.path.join(OUT, "weights.bin"), "wb") as f: + f.write(blob) +with open(os.path.join(OUT, "weights_manifest.txt"), "w") as f: + f.write("\n".join(lines) + "\n") + +print(f"wrote weights.bin ({len(blob)/1e6:.1f} MB), {len(lines)} arrays (SQ={'on' if sq else 'off'})") +for l in lines[:4] + ["..."] + lines[-5:]: + print(" ", l) diff --git a/optimized/cpu-amx/build/same_s_int8/same_s_int8fused_chunk.inc b/optimized/cpu-amx/build/same_s_int8/same_s_int8fused_chunk.inc new file mode 100644 index 00000000..6ff20390 --- /dev/null +++ b/optimized/cpu-amx/build/same_s_int8/same_s_int8fused_chunk.inc @@ -0,0 +1,49 @@ +// same_s_int8fused_chunk.inc — cache-blocked / chunk-parallel decode (FUSED int8 engine). +// Included at the end of same_s_int8fused_cpu_amx.cpp (shares all statics). Position-invariant +// tiling (receptive field = 2 latent tokens); identical to the naive int8 engine's chunk .inc +// except it calls the FUSED decode(). Quality-neutral vs whole (per-chunk int8 requant only). + +struct Chunk{int a,b,lo,hi;}; +static std::vector chunk_plan(int T,int C,int overlap){ + std::vector plan; int a=0; + while(a plan=chunk_plan(T,C,overlap); + int nch=(int)plan.size(); int L=SIN*T; + if(!parallel){ + static std::vector sublat, subpat; + for(auto&ck:plan){ + int w=ck.hi-ck.lo; + sublat.assign((size_t)LAT*w,0); + for(int c=0;c sublat((size_t)LAT*w), subpat((size_t)OUTCH*SIN*w); + for(int c=0;cfp32 ; quant->i8 ; GEMM ; deq->fp32 ; addbias ; act ; quant->i8 ; ... (Q/DQ per GEMM) +// FUSED : dyt_q (norm+quant->i8 ONE pass) ; GEMM ; deqglu_q (deq+bias+GLU+quant->i8 ONE pass) ; +// GEMM ; deq_res (deq+bias+residual ONE pass). Activation STAYS int8 between GLU & proj. +// +// The ONLY standalone Q/DQ left is the 2 around the mandatory fp32 attention island (deq(qkv)->fp32 +// in, quant(attn-out)->i8 out) — every other requant is folded into a norm / GLU / residual pass +// that already existed. Residual stream x stays fp32 (exactly like the DiT xnext / the bf16 engine), +// so the requant grid is IDENTICAL to naive int8 -> quality must match (verified). +// +// sames_init(weights_base, threads) -> mmap int8 weights ONCE + AMX/omp/oneDNN +// sames_forward(latent[1,256,T], T, out_patches[1,512,16T]) -> whole decode +// sames_forward_chunked(latent, T, C, overlap, parallel, out) -> cache-blocked decode +// sames_unpatch(patches[1,512,L], L, pcm[1,2,256L]) -> torch-free unpatch +// +// build: g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I$ONEINC \ +// same_s_int8fused_cpu_amx.cpp -o same_s_int8fused_cpu_amx.so \ +// -L$ONELIB -ldnnl -Wl,-rpath,$ONELIB -ldl -lpthread -lm +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +// ── architecture constants (same_s_decoder_torch.py) ── +static const int LAT=256, DIM=768, H=12, HD=64, RD=32, HALF=16; +static const int NB=6, FF=2304, GLU2=4608, QKV=3840, OUTCH=512; +static const int SUB=17, ECH=34, SHIFT=17, SIN=16; // 17 tok/latent, 34-tok chunk +static const float SCALE=0.125f; // HD**-0.5 = 64**-0.5 + +// ── optional phase profiler (SAMES_PROF=1) ── +static double PROF[8]={0}; static long PROFN=0; +static const char* PROFLBL[8]={"gemm","attn","dyt_q","rope","deqglu","quant","deq_res",""}; +static bool PROF_ON=false; +static inline double wt(){return omp_get_wtime();} +#define TB(id) do{ if(PROF_ON){double _n=wt(); PROF[id]+=_n-_pt; _pt=_n;} }while(0) + +// ── fast vectorizable transcendentals (pure float, no libm/branches -> AVX-512 auto-vec) ── +static inline float vexp(float x){ + x = x<-87.0f?-87.0f:(x>88.0f?88.0f:x); + float z = x*1.442695041f; + float n = std::floor(z+0.5f); + float f = z-n; + float p = 1.0f+f*(0.6931472f+f*(0.2402265f+f*(0.0555041f+f*(0.0096181f+f*0.0013333f)))); + uint32_t bits=(uint32_t)(((int)n+127)<<23); float s; std::memcpy(&s,&bits,4); + return p*s; +} +static inline float vtanh(float x){ return 1.0f-2.0f/(vexp(2.0f*x)+1.0f); } // == Triton DyT kernel +static inline float vsilu(float x){ return x/(1.0f+vexp(-x)); } +static inline int cdiv(int a,int b){return (a+b-1)/b;} +static inline int8_t q127(float v){ v=std::rintf(v); return (int8_t)(v>127.0f?127.0f:(v<-127.0f?-127.0f:v)); } + +// ------------------------- mmap weights.bin + manifest ------------------------- +struct Ten{void* p; std::string dt; long n; std::vector shp;}; +static std::map TEN; +static char* BASE=nullptr; + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string WBASE = std::string(sa3_home()) + "/same_s_int8fused_cpu_amx/weights"; +static void load_weights(){ + std::string bin=WBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(WBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("[sames_i8fused] weights mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static int8_t* I8 (const std::string&k){return (int8_t*)TEN.at(k).p;} + +// ------------------------- oneDNN int8 matmul (s8 x s8 -> s32): primitive+handle cache +static dnnl::engine* ENG=nullptr; +static void onednn_init(){ ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); } +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M mm; dnnl::stream* strm=nullptr; }; +static std::vector MMC; +static void mmcache_init(int nworkers){ + MMC.clear(); MMC.resize(std::max(1,nworkers)+1); + for(auto& c:MMC) c.strm=new dnnl::stream(*ENG); +} +static void gemm_i8(int slot,const int8_t*A,const int8_t*B,int M,int N,int K,int32_t* C){ + using dt=dnnl::memory::data_type; + MMCache& c=MMC[slot]; + MMKey key{M,N,K}; auto it=c.mm.find(key); + if(it==c.mm.end()){ + #pragma omp critical(mmcreate) + { + it=c.mm.find(key); + if(it==c.mm.end()){ + dnnl::memory::desc a_md({M,K},dt::s8,{K,1}); + dnnl::memory::desc b_md({K,N},dt::s8,{N,1}); + dnnl::memory::desc c_md({M,N},dt::s32,{N,1}); + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)A), + dnnl::memory(pd.weights_desc(),*ENG,(void*)B), + dnnl::memory(pd.dst_desc(),*ENG,(void*)C)}; + it=c.mm.emplace(key,std::move(e)).first; + } + } + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)A); e.bm.set_data_handle((void*)B); e.cm.set_data_handle((void*)C); + e.prim.execute(*c.strm,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + c.strm->wait(); +} + +// ------------------------- RoPE table for one 34-token chunk (positions 0..33) ------------------------- +static float RCOS[ECH*HALF], RSIN[ECH*HALF]; // [pos, i] i=0..15 +static void build_rope(){ + for(int i=0;i writes int8). +struct Arena{ + int cap=0; + std::vector x,xt,h,qkv,ao; // f32 activations (residual stream stays fp32) + std::vector srcq; // int8 GEMM-src staging (>= max K = DIM*3 = FF) + std::vector acc; // int32 GEMM accumulator (>= max N = GLU2) + std::vector ascl; // per-row activation scale [M] + void ensure(int M){ + if(M<=cap) return; cap=M; + x.assign((size_t)M*DIM,0); xt.assign((size_t)M*DIM,0); h.assign((size_t)M*DIM,0); + qkv.assign((size_t)M*QKV,0); ao.assign((size_t)M*DIM,0); + srcq.assign((size_t)M*GLU2,0); // headroom (widest src is DIM*3=FF AR; + +// ------------------------- fp32 attention-prep elementwise (identical to bf16/naive engine) ------------------------- +static void dyt_slice(float*base,int M,int W,int coloff,float alpha,const float*g,const float*b,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mi8 boundary + project_in/conv) +static void quant_rows(const float*x,int M,int K,int8_t* q,float* s,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mamax)amax=av; } + float sc=amax>0.0f? amax/127.0f : 1e-12f; s[m]=sc; float inv=1.0f/sc; + #pragma omp simd + for(int k=0;k f32 (the ONE dequant into the fp32 attention island): o=acc*a_scale[m]*w_scale[n] +static void deq(const int32_t*acc,const float*as,const float*bs,int M,int N,float* o,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mfp32 h) + quant_rows(h). Bit-identical (t is fp32 either way), 1 pass not 2. +static void dyt_q(const float*x,int M,int K,float alpha,const float*g,const float*b, + int8_t* q,float* s,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mamax)amax=au; } + float sc=amax>0.0f? amax/127.0f : 1e-12f; s[m]=sc; float inv=1.0f/sc; + #pragma omp simd + for(int j=0;jfp32 [M,2F]) + addbias + glu-silu(->fp32 [M,F]) + quant. 1 pass, no fp32 glu buffer. +static void deqglu_q(const int32_t*acc,const float*as,const float*bs,const float*bias, + int M,int Fdim,int8_t* q,float* s,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mamax)amax=au; + } + float sc=amax>0.0f? amax/127.0f : 1e-12f; s[m]=sc; float inv=1.0f/sc; + #pragma omp simd + for(int j=0;jfp32 h) + addbias + resadd. 1 pass, no fp32 h buffer. Residual stays fp32. +static void deq_res(const int32_t*acc,const float*as,const float*bs,const float*bias, + float* x,int M,int N,bool par){ + #pragma omp parallel for schedule(static) if(par) + for(int m=0;mmm)mm=dm; if(dd>md)md=dd; + } + float zm=0,zd=0; + #pragma omp simd reduction(+:zm,zd) + for(int j=0;j patches[512,16T]) --------- +static void decode(int slot,bool par,const float* latent,int T,float* out_patches){ + Arena& A=AR[slot]; + int iT=SUB*T; // internal tokens after new-token expansion (17T) + int M2=iT+ECH; // padded token count for the shifted 2nd half + A.ensure(M2); + float* x=A.x.data(); float* xt=A.xt.data(); float* h=A.h.data(); + float* qkv=A.qkv.data(); float* ao=A.ao.data(); + int8_t* srcq=A.srcq.data(); float* ascl=A.ascl.data(); int32_t* acc=A.acc.data(); + float rstd=F32("running_std")[0]; + + // project_in (w8a8): build f32 (latent^T*running_std) rows, per-row quant, GEMM, deq+bias -> xt[T,768] + #pragma omp parallel for schedule(static) if(par) + for(int t=0;tamax)amax=av; } + float sc=amax>0.0f? amax/127.0f : 1e-12f; ascl[t]=sc; float inv=1.0f/sc; + int8_t* qo=srcq+(size_t)t*LAT; + for(int c=0;c int8 (was: dyt->fp32 h ; quant_rows h) + dyt_q(x,M,DIM,F32(W("pre.alpha"))[0],F32(W("pre.gamma")),F32(W("pre.beta")),srcq,ascl,par); TB(2); + gemm_i8(slot,srcq,I8(W("qkv.q")),M,QKV,DIM,acc); TB(0); + deq(acc,ascl,F32(W("qkv.scale")),M,QKV,qkv,par); TB(6); // the ONE dequant into the fp32 island + dyt_slice(qkv,M,QKV,0*DIM,F32(W("qn.alpha"))[0],F32(W("qn.gamma")),F32(W("qn.beta")),par); + dyt_slice(qkv,M,QKV,3*DIM,F32(W("qn.alpha"))[0],F32(W("qn.gamma")),F32(W("qn.beta")),par); + dyt_slice(qkv,M,QKV,1*DIM,F32(W("kn.alpha"))[0],F32(W("kn.gamma")),F32(W("kn.beta")),par); + dyt_slice(qkv,M,QKV,4*DIM,F32(W("kn.alpha"))[0],F32(W("kn.gamma")),F32(W("kn.beta")),par); TB(2); + rope_slice(qkv,M,QKV,0*DIM,par); rope_slice(qkv,M,QKV,1*DIM,par); + rope_slice(qkv,M,QKV,3*DIM,par); rope_slice(qkv,M,QKV,4*DIM,par); TB(3); + diff_attn(qkv,M,ao,par); TB(1); // fp32 attention island + quant_rows(ao,M,DIM,srcq,ascl,par); TB(5); // the ONE quant out of the island + gemm_i8(slot,srcq,I8(W("out.q")),M,DIM,DIM,acc); TB(0); + deq_res(acc,ascl,F32(W("out.scale")),nullptr,x,M,DIM,par); TB(6); // FUSED deq + residual (x += .) + // ---- FFN ---- + dyt_q(x,M,DIM,F32(W("ff.alpha"))[0],F32(W("ff.gamma")),F32(W("ff.beta")),srcq,ascl,par); TB(2); + gemm_i8(slot,srcq,I8(W("glu.q")),M,GLU2,DIM,acc); TB(0); + deqglu_q(acc,ascl,F32(W("glu.scale")),F32(W("glu.b")),M,FF,srcq,ascl,par); TB(4); // deq+bias+SiLU-GLU+quant->i8 (STAYS i8) + gemm_i8(slot,srcq,I8(W("proj.q")),M,DIM,FF,acc); TB(0); // int8 GEMM directly on int8 (no fp32 round-trip) + deq_res(acc,ascl,F32(W("proj.scale")),F32(W("proj.b")),x,M,DIM,par); TB(6); // FUSED deq + bias + residual + }; + + // blocks 0..2 over the iT tokens + for(int b=0;b<3;b++) run_block(b,iT); + // midpoint shift by 17 (pure f32 memcpy, identical to bf16/naive engine) + std::memcpy(xt, x, (size_t)SHIFT*DIM*sizeof(float)); + std::memcpy(xt+(size_t)SHIFT*DIM, x, (size_t)iT*DIM*sizeof(float)); + std::memcpy(xt+(size_t)(SHIFT+iT)*DIM, x+(size_t)(iT-SHIFT)*DIM, (size_t)SHIFT*DIM*sizeof(float)); + std::memcpy(x, xt, (size_t)M2*DIM*sizeof(float)); + // blocks 3..5 over M2 tokens + for(int b=3;b<6;b++) run_block(b,M2); + // strip shift + drop latent slot -> y[16T, DIM] + int L=SIN*T; + float* y=A.h.data(); // reuse h as [16T,DIM] staging + #pragma omp parallel for schedule(static) if(par) + for(int t=0;t512 k3 pad1 via im2col (w8a8): fused im2col + per-row quant -> s8s8 GEMM -> fused deq+bias+transpose + #pragma omp parallel for schedule(static) if(par) + for(int l=0;l=0&&jamax)amax=av; + } + float sc=amax>0.0f? amax/127.0f : 1e-12f; ascl[l]=sc; float inv=1.0f/sc; + int8_t* qo=srcq+(size_t)l*(DIM*3); + for(int i=0;i out_patches[512,L] + { const float* bs=F32("conv.scale"); const float* bias=F32("conv.b"); + #pragma omp parallel for schedule(static) if(par) + for(int l=0;l0?100*PROF[i]/tot:0); + printf(" total=%.1fms\n",tot*1e3); + for(int i=0;i<8;i++) PROF[i]=0; PROFN=0; fflush(stdout); +} + +} // extern "C" + +// ---- chunk-parallel decode appended below via include (shares all statics) ---- +#include "same_s_int8fused_chunk.inc" diff --git a/optimized/cpu-amx/build/t5gemma/build.sh b/optimized/cpu-amx/build/t5gemma/build.sh new file mode 100644 index 00000000..854a0a21 --- /dev/null +++ b/optimized/cpu-amx/build/t5gemma/build.sh @@ -0,0 +1,8 @@ +#!/usr/bin/env bash +# Build the torch-free C++ AMX-BF16 T5Gemma encoder engine -> t5gemma_cpu_amx.so +set -euo pipefail +ONE=${ONEDNN_HOME:?set ONEDNN_HOME to a static oneDNN+OpenMP build} +cd "${SA3_CPUAMX_HOME:?set SA3_CPUAMX_HOME}/t5gemma_cpu_amx" +g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I"$ONE/include" \ + t5gemma_cpu_amx.cpp -o t5gemma_cpu_amx.so "$ONE/lib/libdnnl.a" -ldl -lpthread -lm +echo "built t5gemma_cpu_amx.so ($(stat -c%s t5gemma_cpu_amx.so) bytes)" diff --git a/optimized/cpu-amx/build/t5gemma/dump_weights.py b/optimized/cpu-amx/build/t5gemma/dump_weights.py new file mode 100644 index 00000000..d0d3ec2e --- /dev/null +++ b/optimized/cpu-amx/build/t5gemma/dump_weights.py @@ -0,0 +1,98 @@ +#!/usr/bin/env python3 +"""Dump the T5Gemma-b-b-ul2 encoder npz -> {weights.bin + weights_manifest.txt} +for the torch-free C++ AMX engine. MATCHES the SAME-L dump format exactly so the +mmap loader (load_weights) is reused unchanged. + +Layout choices (mirror same_l_cpu_amx/dump_weights.py): + * The 7 linears/layer + embed table: stored **bf16**. + - Linears (nn.Linear weight [out,in]): stored as W.T = [in=K, out=N] so the + oneDNN gemm_bf16(src[M,K], wei[K,N]) consumes them directly (pre-transposed). + - Embedding [256000,768]: stored as-is (a gather table; row = token id). bf16. + * RMSNorm weights, rope_inv_freq: fp32 (the cancellation-fragile islands stay fp32). + +bf16 conversion is round-to-nearest-even done in pure numpy, BIT-IDENTICAL to the +C++ runtime f2b() (r = x + 0x7fff + ((x>>16)&1); bf16 = r>>16). No torch needed. + +Manifest line: name dtype byte_offset nelem d0 d1 ... +""" + +import os + +# Paths come from the environment so nothing local is baked in. +# SA3_CPUAMX_HOME where the engine dirs live (default ./engines) +# SA3_REPO checkout providing the reference weights to dump +HOME = os.environ.get("SA3_CPUAMX_HOME", os.path.abspath("engines")) +REPO = os.environ.get("SA3_REPO", os.path.abspath(".")) +import os +import numpy as np + +NPZ = os.environ.get("SA3_T5GEMMA_NPZ", + os.path.join(HOME, "t5gemma_cpu_amx", "t5gemma_f16.npz")) +OUT = os.path.join(HOME, "t5gemma_cpu_amx") +NB = 12 + + +def f32_to_bf16_rne(arr): + """f32 -> bf16 (uint16) round-to-nearest-even, bit-identical to C++ f2b().""" + x = np.ascontiguousarray(arr, dtype=np.float32).view(np.uint32).astype(np.uint64) + r = x + np.uint64(0x7FFF) + ((x >> np.uint64(16)) & np.uint64(1)) + return (r >> np.uint64(16)).astype(np.uint16) + + +z = np.load(NPZ, allow_pickle=True) +def a(k): # fp32 numpy + return z[k].astype(np.float32) + +blob = bytearray() +lines = [] + + +def put(name, arr, dt): + """dt in {'bf16','f32'}. arr is numpy fp32 (any shape); stored C-contiguous.""" + global blob + arr = np.ascontiguousarray(arr.astype(np.float32)) + off = len(blob) + if dt == "bf16": + blob += f32_to_bf16_rne(arr).tobytes() + elif dt == "f32": + blob += arr.ravel().tobytes() + else: + raise ValueError(dt) + shp = " ".join(str(s) for s in arr.shape) + lines.append(f"{name} {dt} {off} {arr.size} {shp}") + + +def put_lin(name, w_oi): + """w_oi = npz [out,in]; store bf16 W.T = [in,out] for oneDNN wei[K,N].""" + put(name, np.ascontiguousarray(w_oi.T), "bf16") + + +# ---- top level ---- +put("embed", a("embed_tokens.weight"), "bf16") # [256000,768] gather table, bf16 +put("norm", a("norm.weight"), "f32") # [768] +put("rope_inv", z["rope_inv_freq"].astype(np.float32), "f32") # [32] + +# ---- 12 layers ---- +for i in range(NB): + p = f"layers.{i}." + put(f"L{i}.pre_a", a(p + "pre_self_attn_layernorm.weight"), "f32") + put(f"L{i}.post_a", a(p + "post_self_attn_layernorm.weight"), "f32") + put(f"L{i}.pre_f", a(p + "pre_feedforward_layernorm.weight"), "f32") + put(f"L{i}.post_f", a(p + "post_feedforward_layernorm.weight"), "f32") + put_lin(f"L{i}.q", a(p + "self_attn.q_proj.weight")) # [768,768] -> bf16 [768,768] + put_lin(f"L{i}.k", a(p + "self_attn.k_proj.weight")) + put_lin(f"L{i}.v", a(p + "self_attn.v_proj.weight")) + put_lin(f"L{i}.o", a(p + "self_attn.o_proj.weight")) + put_lin(f"L{i}.gate", a(p + "mlp.gate_proj.weight")) # [2048,768] -> bf16 [768,2048] + put_lin(f"L{i}.up", a(p + "mlp.up_proj.weight")) # [2048,768] -> bf16 [768,2048] + put_lin(f"L{i}.down", a(p + "mlp.down_proj.weight")) # [768,2048] -> bf16 [2048,768] + +with open(os.path.join(OUT, "weights.bin"), "wb") as f: + f.write(blob) +with open(os.path.join(OUT, "weights_manifest.txt"), "w") as f: + f.write("\n".join(lines) + "\n") + +print(f"wrote weights.bin ({len(blob)/1e6:.1f} MB), {len(lines)} arrays") +print("first/last few manifest lines:") +for l in lines[:6] + [" ..."] + lines[-3:]: + print(" ", l) diff --git a/optimized/cpu-amx/build/t5gemma/t5gemma_cpu_amx.cpp b/optimized/cpu-amx/build/t5gemma/t5gemma_cpu_amx.cpp new file mode 100644 index 00000000..924dd155 --- /dev/null +++ b/optimized/cpu-amx/build/t5gemma/t5gemma_cpu_amx.cpp @@ -0,0 +1,426 @@ +// t5gemma_cpu_amx.cpp — torch-free C++ AMX-BF16 T5Gemma ENCODER as a callable .so. +// +// google/t5gemma-b-b-ul2 encoder half (Gemma2-style): 12 layers, dim=768, 12 heads, +// head_dim=64, GeGLU(2048), RMSNorm(1+w) sandwich, RoPE theta=10000 half-half, +// attn logit softcap=50, embed x sqrt(768). Seq FIXED at 256 (pad token id 0). +// STANDARD softmax attention (not differential/chunked) — simpler than the decoders. +// +// Mirrors the proven SAME-L / SAME-S engines: oneDNN AMX-BF16 GEMMs for every linear +// (q,k,v,o,gate,up,down + the embed gather table are bf16). RMSNorm, RoPE, softcap-tanh, +// softmax, GeLU and residuals stay fp32 (AVX-512 intrinsics for exp/tanh). The attention +// QK^T and P@V are ALSO bf16-AMX batched matmuls (per-head) — T5Gemma uses STANDARD softmax +// attention (not SAME-L's cancellation-fragile differential attn), so the softcap+softmax +// fp32 island absorbs the bf16 rounding: 12-token gate prompts stay 62-67 dB / cos>=0.9997. +// (Scores + softmax are computed in fp32 between the two bf16 matmuls.) +// +// t5g_init(weights_base, threads) -> mmap bf16 weights + AMX/omp/oneDNN +// t5g_forward(ids[256] i32, mask[256] i32, out[256*768]) -> last_hidden_state fp32 +// +// build: g++ -O3 -march=native -std=c++17 -fopenmp -shared -fPIC -I$ONEINC \ +// t5gemma_cpu_amx.cpp -o t5gemma_cpu_amx.so $ONELIB/libdnnl.a -ldl -lpthread -lm +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#if defined(__AVX512F__) +#include +#endif +#include "oneapi/dnnl/dnnl.hpp" +#include "oneapi/dnnl/dnnl_debug.h" + +// ── architecture constants (t5gemma-b-b-ul2 encoder) ── +static const int S=256, DIM=768, H=12, HD=64, HALF=32, NB=12, FF=2048; +static const float SOFTCAP=50.0f, SCALE=0.125f, EPS=1e-6f; // SCALE=64**-0.5 +static float EMBED_SCALE=0.0f; // sqrt(768), set in init + +// ── optional phase profiler (T5G_PROF=1) ── +static double PROF[8]={0}; +static const char* PROFLBL[8]={"gemm","attn","norm","rope","glu","cast","misc",""}; +static bool PROF_ON=false; +static inline double wt(){return omp_get_wtime();} +#define TB(id) do{ if(PROF_ON){double _n=wt(); PROF[id]+=_n-_pt; _pt=_n;} }while(0) + +typedef uint16_t bf16; +static inline bf16 f2b(float f){ // round-to-nearest-even f32->bf16 + uint32_t x; std::memcpy(&x,&f,4); + uint32_t r=x+0x7fff+((x>>16)&1); return (bf16)(r>>16); +} +static inline float b2f(bf16 h){ uint32_t x=(uint32_t)h<<16; float f; std::memcpy(&f,&x,4); return f; } + +// ── fast vectorizable transcendentals (pure float -> AVX-512 auto-vec; err ~2e-6 << bf16) ── +static inline float vexp(float x){ + x = x<-87.0f?-87.0f:(x>88.0f?88.0f:x); + float z = x*1.442695041f; + float n = std::floor(z+0.5f); + float f = z-n; + float p = 1.0f+f*(0.6931472f+f*(0.2402265f+f*(0.0555041f+f*(0.0096181f+f*0.0013333f)))); + uint32_t bits=(uint32_t)(((int)n+127)<<23); float s; std::memcpy(&s,&bits,4); + return p*s; +} +static inline float vtanh(float x){ return 1.0f-2.0f/(vexp(2.0f*x)+1.0f); } +// gelu_pytorch_tanh: 0.5*x*(1+tanh(sqrt(2/pi)*(x+0.044715 x^3))) +static inline float vgelu(float x){ + float x3=x*x*x; + return 0.5f*x*(1.0f+vtanh(0.7978845608028654f*(x+0.044715f*x3))); +} + +#if defined(__AVX512F__) +// AVX-512 exp/tanh — SAME polynomial as scalar vexp() (numerically identical), but the +// 16-lane transcendentals actually vectorize (the scalar vexp/vtanh do NOT auto-vec inside +// omp-simd because of floor/bitcast). This is the attention hot-path lever. +static inline __m512 exp512(__m512 x){ + x=_mm512_max_ps(x,_mm512_set1_ps(-87.0f)); x=_mm512_min_ps(x,_mm512_set1_ps(88.0f)); + __m512 z=_mm512_mul_ps(x,_mm512_set1_ps(1.442695041f)); + __m512 n=_mm512_roundscale_ps(z,_MM_FROUND_TO_NEAREST_INT|_MM_FROUND_NO_EXC); + __m512 f=_mm512_sub_ps(z,n); + __m512 p=_mm512_set1_ps(0.0013333f); + p=_mm512_fmadd_ps(p,f,_mm512_set1_ps(0.0096181f)); + p=_mm512_fmadd_ps(p,f,_mm512_set1_ps(0.0555041f)); + p=_mm512_fmadd_ps(p,f,_mm512_set1_ps(0.2402265f)); + p=_mm512_fmadd_ps(p,f,_mm512_set1_ps(0.6931472f)); + p=_mm512_fmadd_ps(p,f,_mm512_set1_ps(1.0f)); + __m512i ni=_mm512_add_epi32(_mm512_cvtps_epi32(n),_mm512_set1_epi32(127)); + __m512 s=_mm512_castsi512_ps(_mm512_slli_epi32(ni,23)); + return _mm512_mul_ps(p,s); +} +static inline __m512 tanh512(__m512 x){ + __m512 e=exp512(_mm512_mul_ps(x,_mm512_set1_ps(2.0f))); + return _mm512_sub_ps(_mm512_set1_ps(1.0f), + _mm512_div_ps(_mm512_set1_ps(2.0f),_mm512_add_ps(e,_mm512_set1_ps(1.0f)))); +} +#endif + +// ------------------------- mmap weights.bin + manifest (verbatim from SAME-L) ------------------------- +struct Ten{void* p; std::string dt; long n; std::vector shp;}; +static std::map TEN; +static char* BASE=nullptr; + +// Engine paths resolve from $SA3_CPUAMX_HOME (same base the Python side uses), so nothing +// absolute is baked into the binary. Falls back to the current directory. +static const char* sa3_home() { + const char* v = getenv("SA3_CPUAMX_HOME"); + return (v && *v) ? v : "."; +} +static std::string WBASE = std::string(sa3_home()) + "/t5gemma_cpu_amx/weights"; +static void load_weights(){ + std::string bin=WBASE+".bin"; + int fd=open(bin.c_str(),O_RDONLY); struct stat st; fstat(fd,&st); + BASE=(char*)mmap(nullptr,st.st_size,PROT_READ,MAP_PRIVATE,fd,0); + if(BASE==MAP_FAILED){perror("mmap");exit(1);} close(fd); + std::ifstream mf(WBASE+"_manifest.txt"); std::string line; + while(std::getline(mf,line)){ + std::istringstream ss(line); Ten t; std::string name; long off; + ss>>name>>t.dt>>off>>t.n; long d; while(ss>>d)t.shp.push_back(d); + t.p=(void*)(BASE+off); TEN[name]=t; + } + printf("[t5g] weights mmap'd: %ld arrays (%s)\n",(long)TEN.size(),bin.c_str()); +} +static float* F32(const std::string&k){return (float*)TEN.at(k).p;} +static bf16* BF (const std::string&k){return (bf16*)TEN.at(k).p;} + +// ------------------------- oneDNN bf16 matmul (bf16 x bf16 -> f32): primitive+handle cache +static dnnl::engine* ENG=nullptr; +static void onednn_init(){ ENG=new dnnl::engine(dnnl::engine::kind::cpu,0); } +struct MMKey{int M,N,K; bool operator<(const MMKey&o)const{ + return M!=o.M?M mm; dnnl::stream* strm=nullptr; }; +static std::vector MMC; +static void mmcache_init(int nworkers){ + MMC.clear(); MMC.resize(std::max(1,nworkers)+1); + for(auto& c:MMC) c.strm=new dnnl::stream(*ENG); +} +// src A[M,K] bf16, wei B[K,N] bf16 (row-major), dst C[M,N] f32. cache slot `slot`. +static void gemm_bf16(int slot,const bf16*A,const bf16*B,int M,int N,int K,float* C){ + using dt=dnnl::memory::data_type; + MMCache& c=MMC[slot]; + MMKey key{M,N,K}; auto it=c.mm.find(key); + if(it==c.mm.end()){ + #pragma omp critical(mmcreate) + { + it=c.mm.find(key); + if(it==c.mm.end()){ + dnnl::memory::desc a_md({M,K},dt::bf16,{K,1}); + dnnl::memory::desc b_md({K,N},dt::bf16,{N,1}); + dnnl::memory::desc c_md({M,N},dt::f32, {N,1}); + dnnl::matmul::primitive_desc pd(*ENG,a_md,b_md,c_md); + MMEnt e{dnnl::matmul(pd), + dnnl::memory(pd.src_desc(),*ENG,(void*)A), + dnnl::memory(pd.weights_desc(),*ENG,(void*)B), + dnnl::memory(pd.dst_desc(),*ENG,(void*)C)}; + it=c.mm.emplace(key,std::move(e)).first; + } + } + } + MMEnt& e=it->second; + e.am.set_data_handle((void*)A); e.bm.set_data_handle((void*)B); e.cm.set_data_handle((void*)C); + e.prim.execute(*c.strm,{{DNNL_ARG_SRC,e.am},{DNNL_ARG_WEIGHTS,e.bm},{DNNL_ARG_DST,e.cm}}); + c.strm->wait(); +} + +// ------------------------- RoPE table (half-half, positions 0..S-1, 32 freqs) ------------------------- +static std::vector RCOS, RSIN; // [S*HALF] +static void rope_build(){ + const float* inv=F32("rope_inv"); // (32,) fp32 from npz + RCOS.resize((size_t)S*HALF); RSIN.resize((size_t)S*HALF); + for(int p=0;p x,h,q,k,v,ao,gate,up,scores,outh; // f32 activations + std::vector srcb, qg,kgT,vg,pb; // bf16 staging (GEMM src + packed attn operands) + void init(){ + x.assign((size_t)S*DIM,0); h.assign((size_t)S*DIM,0); + q.assign((size_t)S*DIM,0); k.assign((size_t)S*DIM,0); v.assign((size_t)S*DIM,0); + ao.assign((size_t)S*DIM,0); gate.assign((size_t)S*FF,0); up.assign((size_t)S*FF,0); + scores.assign((size_t)H*S*S,0); outh.assign((size_t)H*S*HD,0); + srcb.assign((size_t)S*FF,0); + qg.assign((size_t)H*S*HD,0); kgT.assign((size_t)H*HD*S,0); + vg.assign((size_t)H*S*HD,0); pb.assign((size_t)H*S*S,0); + } +}; +static Arena A; + +// ------------------------- fp32 elementwise kernels ------------------------- +// Gemma RMSNorm: n = x/sqrt(mean(x^2)+eps); out = n*(1+w). Accumulate in double (matches numpy ref). +static void rmsnorm(const float* x,float* o,const float* w,int M){ + #pragma omp parallel for schedule(static) + for(int m=0;mbf16 round-to-nearest-even, 16-lane -> pb + __m512i bits=_mm512_castps_si512(p); + __m512i lsb=_mm512_and_si512(_mm512_srli_epi32(bits,16),_mm512_set1_epi32(1)); + bits=_mm512_add_epi32(_mm512_add_epi32(bits,_mm512_set1_epi32(0x7fff)),lsb); + _mm256_storeu_si256((__m256i*)(pr+j),_mm512_cvtepi32_epi16(_mm512_srli_epi32(bits,16))); + } +#else + for(int j=0;jmx?sr[j]:mx; + z=0; for(int j=0;j ao[S, h*HD] ---- + #pragma omp parallel for collapse(2) schedule(static) + for(int h=0;h bf16 for down GEMM + #pragma omp parallel for schedule(static) + for(int m=0;m out last_hidden_state[256*768] f32 (caller-allocated) +void t5g_forward(const int32_t* ids,const int32_t* mask,float* out){ + // embedding gather (bf16 table) + x sqrt(768) + bf16* emb=BF("embed"); + #pragma omp parallel for schedule(static) + for(int s=0;s caller buffer +} + +int t5g_DIM(){return DIM;} int t5g_S(){return S;} + +void t5g_prof_dump(){ + double tot=0; for(int i=0;i<7;i++) tot+=PROF[i]; + printf("[prof] "); + for(int i=0;i<7;i++) printf("%s=%.1fms(%.0f%%) ",PROFLBL[i],PROF[i]*1e3,tot>0?100*PROF[i]/tot:0); + printf(" total=%.1fms\n",tot*1e3); + for(int i=0;i<8;i++) PROF[i]=0; fflush(stdout); +} + +} // extern "C" diff --git a/optimized/cpu-amx/requirements-gradio.txt b/optimized/cpu-amx/requirements-gradio.txt new file mode 100644 index 00000000..efaab0c8 --- /dev/null +++ b/optimized/cpu-amx/requirements-gradio.txt @@ -0,0 +1,4 @@ +# Web UI extras (on top of requirements.txt). +gradio>=4.0 +pillow>=10.0 +soundfile>=0.12 diff --git a/optimized/cpu-amx/requirements.txt b/optimized/cpu-amx/requirements.txt new file mode 100644 index 00000000..b570a122 --- /dev/null +++ b/optimized/cpu-amx/requirements.txt @@ -0,0 +1,6 @@ +numpy>=1.24 +sentencepiece>=0.2 +# Non-WAV / non-44.1k init audio for --init-audio is decoded via ffmpeg if present. +soundfile>=0.12 +# No torch. Text-to-audio, CFG, audio-to-audio and inpainting all run on the C++ AMX +# engines + numpy; the a2a/inpaint init-encoder is the C++ SAME-S/SAME-L encoder. diff --git a/optimized/cpu-amx/sa3 b/optimized/cpu-amx/sa3 new file mode 100755 index 00000000..37dd0052 --- /dev/null +++ b/optimized/cpu-amx/sa3 @@ -0,0 +1,33 @@ +#!/usr/bin/env bash +# +# Wrapper around scripts/sa3_cpu_amx.py — the cpu-amx (torch-free C++ AMX) SA3 CLI. +# +# Interpreter resolution (first that works): +# 1. $SA3_PYTHON (export it to pin a specific python) +# 2. $SCRIPT_DIR/.venv/bin/python +# 3. python3 on PATH +# +# The runtime needs numpy + sentencepiece (+ soundfile/ffmpeg for non-wav input; +# + torch & CUDA only for --init-audio / --inpaint-range). See requirements.txt. +# +# Usage: ./sa3 --prompt "lofi house" --dit medium --decoder same-l --out a.wav +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +pick_python() { + if [[ -n "${SA3_PYTHON:-}" ]]; then echo "$SA3_PYTHON"; return; fi + if [[ -x "$SCRIPT_DIR/.venv/bin/python" ]]; then echo "$SCRIPT_DIR/.venv/bin/python"; return; fi + command -v python3 || command -v python +} +PY="$(pick_python)" + +if [[ -z "$PY" ]] || ! "$PY" -c "import numpy" >/dev/null 2>&1; then + echo "error: no usable python with numpy found." >&2 + echo " set SA3_PYTHON to an interpreter that has numpy+sentencepiece," >&2 + echo " or create a .venv (see requirements.txt)." >&2 + exit 1 +fi + +cd "$SCRIPT_DIR" +exec "$PY" scripts/sa3_cpu_amx.py "$@" diff --git a/optimized/cpu-amx/sa3-gradio b/optimized/cpu-amx/sa3-gradio new file mode 100755 index 00000000..48ed1a73 --- /dev/null +++ b/optimized/cpu-amx/sa3-gradio @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +# Shortcut for the cpu-amx gradio web UI (scripts/sa3_gradio.py). Args pass through. +# +# ./sa3-gradio # same-l default, public share link +# ./sa3-gradio --decoder same-s +# ./sa3-gradio --no-share # local-only (http://127.0.0.1:7860) +# +# Interpreter: $SA3_PYTHON, else .venv/bin/python, else python3. +# The UI needs a few extra packages (gradio, pillow, soundfile) — see requirements-gradio.txt. +set -euo pipefail +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +pick_python() { + if [[ -n "${SA3_PYTHON:-}" ]]; then echo "$SA3_PYTHON"; return; fi + if [[ -x "$SCRIPT_DIR/.venv/bin/python" ]]; then echo "$SCRIPT_DIR/.venv/bin/python"; return; fi + command -v python3 || command -v python +} +PY="$(pick_python)" + +if [[ -z "$PY" ]] || ! "$PY" -c "import gradio, PIL, numpy" >/dev/null 2>&1; then + echo "The gradio UI needs extra packages (gradio, pillow, soundfile, numpy)." >&2 + echo "Install them into your interpreter: $PY -m pip install -r requirements-gradio.txt" >&2 + exit 1 +fi + +cd "$SCRIPT_DIR" +exec "$PY" scripts/sa3_gradio.py "$@" diff --git a/optimized/cpu-amx/scripts/backends.py b/optimized/cpu-amx/scripts/backends.py new file mode 100644 index 00000000..0216afae --- /dev/null +++ b/optimized/cpu-amx/scripts/backends.py @@ -0,0 +1,216 @@ +"""Backend loaders for the cpu-amx release. + +Every heavy component is a torch-free C++ AMX engine, loaded by inserting its +source directory on sys.path and importing the ctypes shim (the .so + weights +live in those directories; we do not copy them). The a2a/inpaint init-encoder is +now ALSO a torch-free C++ AMX engine (SAME-S / SAME-L, matched to the decoder), +so the whole release is C++/numpy — no torch anywhere. + +Components + T5Gemma : t5gemma_cpu_amx.so (bf16 AMX) -> [1,256,768] + DiT : dit_cpu_amx.so (int8 AMX, MEDIUM) cond_backend(x,t,cross,gcond)->v + decoders : same_{s,l}_cpu_amx.so (bf16, default) or *_int8fused (int8, --decoder-precision int8) + encoder : same_{s,l}_encoder_cpu_amx.so (bf16 AMX) audio[1,2,N] -> latent[1,256,T] (a2a/inpaint) + +DiT stability note: the dit .so heap-corrupts if called at more than one +sequence length in a process, and double-frees at teardown. The CLI runs one +generation (one T_lat) per process and os._exit(0)s before teardown; the gradio +isolates each generation in its own subprocess. DiT threads default to 1 (the +verified-stable setting; it is fast enough for short clips). +""" +from __future__ import annotations + +import os +import sys + +import numpy as np + +# ── C++ backend engine directories ─────────────────────────────────────────────────────── +# All resolved from ONE base so this file and weights.py (which downloads into it) can never +# disagree. Override with SA3_CPUAMX_HOME; see weights.py for the DiT .so path caveat. +try: + from weights import HOME as _HOME +except ImportError: # standalone import of this module + _HOME = os.environ.get("SA3_CPUAMX_HOME", + os.path.expanduser("~/.cache/stable-audio-3/cpu-amx")) + +_D = lambda name: os.path.join(_HOME, name) + +DIR_T5 = _D("t5gemma_cpu_amx") +DIR_DIT = _D("dit_medium_cpu_amx") # cpu_amx_backend.py (DiTCppAmx) +DIR_SAMES = _D("same_s_cpu_amx") +DIR_SAMEL = _D("same_l_cpu_amx") +DIR_SAMES_INT8 = _D("same_s_int8fused_cpu_amx") +DIR_SAMEL_INT8 = _D("same_l_int8fused_cpu_amx") +DIR_SAMES_ENC = _D("same_s_encoder_cpu_amx") # C++ AMX SAME-S encoder (bf16) +DIR_SAMEL_ENC = _D("same_l_encoder_cpu_amx") # C++ AMX SAME-L encoder (bf16) +DIR_SAMES_ENC_INT8 = _D("same_s_encoder_int8fused_cpu_amx") # SAME-S encoder (int8) +DIR_SAMEL_ENC_INT8 = _D("same_l_encoder_int8fused_cpu_amx") # SAME-L encoder (int8) + +DIT_THREADS = 1 # verified-stable; the .so heap-races at higher thread counts + + +def _add_path(p): + if p not in sys.path: + sys.path.insert(0, p) + + +try: + import weights as _hfw # pulls each engine's .so + weight blob from HF on first use +except Exception: + _hfw = None + +def _ensure(group): + """Download an engine's binaries from HF if not present locally (no-op on a local build).""" + if _hfw is not None: + try: + _hfw.ensure(group) + except Exception as e: + print(f"[cpu-amx] HF fetch '{group}' failed ({e}); assuming local build", flush=True) + + +# ── T5Gemma text encoder ──────────────────────────────────────────────────── +def load_t5gemma(threads=16): + """Returns a callable enc(ids[1,256] int, mask[1,256] int) -> [1,256,768] fp32.""" + _ensure("t5gemma") + _add_path(DIR_T5) + from t5gemma_cpu_backend import T5GemmaCPU + return T5GemmaCPU(threads=int(threads)) + + +# ── DiT (medium) — the cond_backend for sample()/sample_cfg ───────────────── +def load_dit(precision="int8", threads=None): + """Returns a DiTCppAmx: __call__(x[1,256,T], t, cross[1,257,768], gcond[1,768]) -> v[1,256,T]. + + precision 'int8' (default, shipped): the int8 C++ core, pinned to 1 thread (its .so + heap-races above 1) — fast, ~40 dB. 'bf16': the near-lossless fp32-RoPE/RMSNorm-islands + core (dit_cpu_amx_bf16.so), ~59/54 dB, runs at `threads` (~1.24× the int8 latency).""" + assert precision in ("int8", "bf16"), precision + if precision == "int8": + _ensure("dit") + else: + _ensure("dit_bf16") # HF: dit_cpu_amx_bf16.so + core_bf16 + pin_fp32 + bf16 flash + dit_threads = threads # BOTH precisions now run at `threads` (int8 stress-tested at 16) + _add_path(DIR_DIT) + from cpu_amx_backend import DiTCppAmx + return DiTCppAmx(precision=precision, threads=dit_threads) + + +# ── Decoders (bf16 default, int8 optional) ────────────────────────────────── +class Decoder: + """Wraps a C++ SAME-S / SAME-L engine as decode(latent[1,256,T]) -> audio[2,N]. + + Even-length requirement (PAD_MODULO=2): SAME-S (both precisions) AND the + SAME-L *int8-fused* engine require an even latent length — an odd T is + edge-padded by one column, decoded, and the extra 4096 samples trimmed. + SAME-L *bf16* takes any length (its band attention is linear; forward_pcm + chunks C=64/overlap=8).""" + + def __init__(self, name: str, precision: str = "bf16", threads: int = 16): + assert name in ("same-s", "same-l"), name + assert precision in ("bf16", "int8"), precision + self.name, self.precision = name, precision + self.is_sames = name == "same-s" + # SAME-S (any prec) and SAME-L int8 assert even T; SAME-L bf16 does not. + self.needs_even = self.is_sames or (name == "same-l" and precision == "int8") + if name == "same-s" and precision == "bf16": + _ensure("same_s_decoder_bf16"); _add_path(DIR_SAMES) + from same_s_cpu_backend import SamesCPU + self.m = SamesCPU(threads=threads) + elif name == "same-s" and precision == "int8": + _ensure("same_s_decoder_int8"); _add_path(DIR_SAMES_INT8) + from same_s_int8fused_backend import SamesInt8FusedCPU + self.m = SamesInt8FusedCPU(threads=threads) + elif name == "same-l" and precision == "bf16": + _ensure("same_l_decoder_bf16"); _add_path(DIR_SAMEL) + from same_l_cpu_backend import SamelCPU + self.m = SamelCPU(threads=threads) + else: # same-l int8 + _ensure("same_l_decoder_int8"); _add_path(DIR_SAMEL_INT8) + from same_l_int8fused_backend import SamelInt8FusedCPU + self.m = SamelInt8FusedCPU(threads=threads) + + def decode(self, latent: np.ndarray) -> np.ndarray: + """latent (1,256,T) fp32 -> audio (2, T*4096) fp32.""" + latent = np.ascontiguousarray(latent, np.float32) + T = latent.shape[-1] + if self.needs_even and (T % 2 == 1): + padded = np.concatenate([latent, latent[..., -1:]], axis=-1) # -> even + pcm = self.m.forward_pcm(padded)[0] # (2,(T+1)*4096) + return np.ascontiguousarray(pcm[:, : T * SAMPLES_PER_LATENT]) + return self.m.forward_pcm(latent)[0] # (2, T*4096) + + +from pipeline import SAMPLES_PER_LATENT # noqa: E402 (after Decoder so import is cheap) + + +def load_decoder(name: str, precision: str = "bf16", threads: int = 16) -> Decoder: + return Decoder(name, precision, threads) + + +# ── Audio-to-audio / inpaint init-encoder (torch-free C++ AMX SAME-{S,L} encoder) ── +class CppEncoder: + """Torch-free C++ AMX-BF16 autoencoder ENCODER for a2a / inpaint. Matches the + chosen decoder (same-s decoder -> SAME-S encoder, same-l -> SAME-L encoder), so + encode/decode share the same autoencoder. numpy + ctypes only — no torch, no 2 GB + checkpoint load. + + encode(audio (2,N) fp32, T_lat) -> latent (1,256,T_lat) fp32. + The encoder downsamples 4096 audio samples per latent token, so the audio is + trimmed/zero-padded to exactly T_lat*4096 samples before encoding.""" + + def __init__(self, name: str, precision: str = "bf16", threads: int = 16): + assert name in ("same-s", "same-l"), name + assert precision in ("bf16", "int8"), precision + self.name = name; self.precision = precision + self.device = "cpu-amx" + if name == "same-s" and precision == "bf16": + _ensure("same_s_encoder_bf16"); _add_path(DIR_SAMES_ENC) + from same_s_encoder_backend import SamesEncoderCPU + self.m = SamesEncoderCPU(threads=threads) + elif name == "same-s": + _ensure("same_s_encoder_int8"); _add_path(DIR_SAMES_ENC_INT8) + from same_s_encoder_int8fused_backend import SameSEncoderInt8FusedCPU + self.m = SameSEncoderInt8FusedCPU(threads=threads) + elif name == "same-l" and precision == "bf16": + _ensure("same_l_encoder_bf16"); _add_path(DIR_SAMEL_ENC) + from same_l_encoder_backend import SamelEncoderCPU + self.m = SamelEncoderCPU(threads=threads) + else: + _ensure("same_l_encoder_int8"); _add_path(DIR_SAMEL_ENC_INT8) + from same_l_encoder_int8fused_backend import SameLEncoderInt8FusedCPU + self.m = SameLEncoderInt8FusedCPU(threads=threads) + + def encode(self, audio: np.ndarray, T_lat: int) -> np.ndarray: + return np.ascontiguousarray(self.m.encode(audio, T_lat), np.float32) + + +def load_encoder(name: str = "same-l", precision: str = "bf16", threads: int = 16) -> CppEncoder: + """C++ AMX encoder matching the decoder; precision follows --decoder-precision. name in {'same-s','same-l'}.""" + return CppEncoder(name, precision=precision, threads=threads) + + +# ── (fallback) fp32 torch AE encoder — retained for reference; not used by the release ── +def _pick_free_gpu() -> str | None: + """Return the index (as str) of the CUDA device with the most free memory, + or None if CUDA/nvidia-smi is unavailable (then the AE runs on CPU).""" + try: + import subprocess + out = subprocess.run( + ["nvidia-smi", "--query-gpu=memory.free", "--format=csv,noheader,nounits"], + capture_output=True, text=True, timeout=15) + if out.returncode != 0: + return None + free = [int(x) for x in out.stdout.split()] + if not free: + return None + return str(int(np.argmax(free))) + except Exception: + return None + + +# NOTE: a legacy torch fp32 SAME-L encoder (AEEncoder / load_ae_encoder_torch) lived here. +# Nothing called it — a2a/inpaint use the C++ AMX encoders via load_encoder() — and it +# pulled in torch + stable_audio_tools + an unpublished loader, against this release being +# torch-free. Removed. + diff --git a/optimized/cpu-amx/scripts/examples.py b/optimized/cpu-amx/scripts/examples.py new file mode 100644 index 00000000..c057703f --- /dev/null +++ b/optimized/cpu-amx/scripts/examples.py @@ -0,0 +1,83 @@ +"""Shared, colored "Try these commands" block for the cpu-amx release. + +Appended to `./sa3 --help`. The cpu-amx runtime ships the MEDIUM DiT only (the +int8 C++ AMX core), so every example uses `--dit medium`; the decoder toggles +between the native `same-l` and the faster distilled `same-s`. +""" +from __future__ import annotations + +import os +import sys +from pathlib import Path + +SCRIPT_DIR = Path(__file__).resolve().parent.parent + + +def _c(code: str) -> str: + return code if sys.stdout.isatty() else "" + +BOLD, CYAN, GREEN = _c("\033[1m"), _c("\033[1;36m"), _c("\033[1;32m") +YELLOW, DIM, RESET = _c("\033[1;33m"), _c("\033[2m"), _c("\033[0m") + + +def _prefix() -> str: + wrapper = SCRIPT_DIR / "sa3" + if wrapper.exists() and os.access(wrapper, os.X_OK): + return "./sa3" + return "python scripts/sa3_cpu_amx.py" + + +def print_example_commands(header: str | None = None) -> None: + prefix = _prefix() + + def hdr(text): + print(f"\n {CYAN}{text}{RESET}") + + def cmd(args, comment=""): + line = f"{prefix} {args}" + if comment: + print(f" {GREEN}$ {line}{RESET} {DIM}# {comment}{RESET}") + else: + print(f" {GREEN}$ {line}{RESET}") + + if header is None: + header = f"{BOLD}Examples:{RESET}" + print(f"\n{BOLD}━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━{RESET}") + print(f" {header}") + print(f"{BOLD}━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━{RESET}") + + hdr("🎵 Generate audio from a prompt") + cmd('--prompt "A beautiful piano arpeggio grows into a cinematic climax" \\\n' + ' --dit medium --decoder same-l --seconds 30 --out piano.wav', + "native codec, best fidelity") + cmd('--prompt "lofi house loop, 120 BPM" \\\n' + ' --dit medium --decoder same-s --seconds 15 --out lofi.wav', + "distilled same-s decoder — faster") + + hdr("▶ Play immediately after generation") + cmd('--prompt "ambient drone" --dit medium --decoder same-l \\\n' + ' --seconds 10 --out drone.wav --play', + "writes WAV + plays (ffplay/aplay/paplay/afplay)") + + hdr("🎚️ Audio-to-audio & inpainting (needs an input WAV; C++ AMX SAME-{S,L} encoder)") + cmd('--prompt "jazz fusion with electric piano" --dit medium --decoder same-l \\\n' + ' --init-audio funk.wav --init-noise-level 0.7 --out funk_jazz.wav', + "variation: 0.4-0.8 typical") + cmd('--prompt "explosive drum break" --dit medium --decoder same-l \\\n' + ' --init-audio funk.wav --inpaint-range "4,7" --out funk_drums.wav', + "regenerate seconds 4-7, keep the rest") + + hdr("🎯 Steer with CFG + negative prompts") + cmd('--prompt "ambient drone" --cfg 3.0 \\\n' + ' --negative-prompt "drums, vocals, distortion" \\\n' + ' --dit medium --decoder same-l --out clean_drone.wav', + "cfg > 1.0 toward prompt, neg pushes away") + + hdr("⚙️ Precision / speed dials") + cmd('--prompt "techno beat" --dit medium --decoder same-s \\\n' + ' --decoder-precision int8 --threads 32 --out techno.wav', + "int8 fused decoder (smaller/faster); more threads") + + print(f"\n {YELLOW}note:{RESET} cpu-amx ships the MEDIUM DiT only. For sm-music / sm-sfx use " + f"optimized/tflite or optimized/mlx.") + print(f"\n{BOLD}━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━{RESET}\n") diff --git a/optimized/cpu-amx/scripts/pipeline.py b/optimized/cpu-amx/scripts/pipeline.py new file mode 100644 index 00000000..0614619d --- /dev/null +++ b/optimized/cpu-amx/scripts/pipeline.py @@ -0,0 +1,263 @@ +"""SA3 text-to-audio pipeline — numpy host side, self-contained (torch-free). + +This is the cpu-amx sibling of optimized/mlx/models/defs/sa3_pipeline.py and the +`tflite_pipeline` module: the fp32 pre/post and the pingpong rectified-flow +sampler are pure numpy, so the heavy compute (T5Gemma, DiT, decoder) can be any +pluggable backend. Here the backends are the torch-free C++ AMX engines +(see backends.py); this module never imports torch, mlx or tflite. + +Pipeline: + prompt -> SentencePiece -> T5Gemma -> conditioning (cross_attn + global) + -> DiT pingpong (rectified-flow) -> SAME-S / SAME-L decoder -> WAV + +The DiT is a plug-in callable ``cond_backend(x, t, cross, gcond) -> velocity``. +The C++ ``DiTCppAmx`` instance is exactly such a callable, so wiring it into +``sample`` / ``sample_cfg`` gives every feature (CFG, APG, negative prompt, +audio-to-audio, inpaint paste-back) for free. + +Numerics are a verbatim port of tflite_pipeline.py (which matches sa3_mlx.py / +the TensorRT release to ~90 dB), so a cpu-amx generation is the same music as +the other releases for a given prompt/seed (up to backend precision). +""" +from __future__ import annotations + +import math +import subprocess +import wave +from pathlib import Path + +import numpy as np + +# ── constants (shared across the whole SA3 family) ────────────────────────── +SAMPLE_RATE = 44100 +SAMPLES_PER_LATENT = 4096 # decoder upsample (256 patch x 16) +COND_TOKENS = 256 # T5Gemma sequence length + +ASSETS = Path(__file__).resolve().parent.parent / "assets" +TOKENIZER_NPZ = ASSETS / "t5gemma_f16.npz" # SentencePiece model lives in here +COND_NPZ = ASSETS / "cond_medium.npz" # learned padding + seconds embedder + + +# ── WAV I/O ───────────────────────────────────────────────────────────────── +def save_wav(path, audio, sample_rate: int = SAMPLE_RATE): + """audio: (channels, T) float32 in [-1, 1]. Writes 16-bit PCM WAV.""" + audio = np.asarray(audio, np.float32) + if not np.isfinite(audio).all(): + n_bad = int((~np.isfinite(audio)).sum()) + raise RuntimeError(f"refusing to write WAV — {n_bad} non-finite samples (NaN/Inf)") + audio = np.clip(audio, -1.0, 1.0) + pcm = (audio * 32767.0).astype(np.int16).T # (T, channels) interleaved + with wave.open(str(path), "wb") as w: + w.setnchannels(audio.shape[0]) + w.setsampwidth(2) + w.setframerate(sample_rate) + w.writeframes(pcm.tobytes()) + + +def read_wav(path) -> np.ndarray: + """Read a WAV. Returns (2, T) float32 in [-1, 1]. + + 16-bit PCM @ 44.1 kHz is read natively; any other format (24/32-bit, + 48 kHz, mp3/flac) falls back to ffmpeg. Mono is duplicated to stereo.""" + path = str(path) + try: + with wave.open(path, "rb") as w: + nch, sw, sr, nframes = (w.getnchannels(), w.getsampwidth(), + w.getframerate(), w.getnframes()) + if sr == SAMPLE_RATE and sw == 2: + raw = np.frombuffer(w.readframes(nframes), np.int16).astype(np.float32) / 32767.0 + if nch == 1: + return np.stack([raw, raw], axis=0) + return raw.reshape(-1, nch).T[:2] + except wave.Error: + pass + try: + result = subprocess.run( + ["ffmpeg", "-v", "error", "-i", path, + "-f", "s16le", "-ar", str(SAMPLE_RATE), "-ac", "2", "-"], + capture_output=True, check=True) + except FileNotFoundError: + raise RuntimeError( + f"{path}: unsupported WAV format and ffmpeg is not installed.\n" + f"Install ffmpeg, or convert to 16-bit/44.1 kHz stereo first.") + except subprocess.CalledProcessError as e: + raise RuntimeError(f"{path}: ffmpeg failed — {e.stderr.decode().strip()}") + raw = np.frombuffer(result.stdout, np.int16).astype(np.float32) / 32767.0 + return raw.reshape(-1, 2).T + + +# ── Tokenizer (SentencePiece from the bundled npz) ────────────────────────── +class Tokenizer: + def __init__(self, npz_path=TOKENIZER_NPZ): + import sentencepiece as spm + arrs = np.load(npz_path) + if "TOKENIZER_MODEL" not in arrs.files: + raise ValueError(f"{npz_path} missing TOKENIZER_MODEL") + self.sp = spm.SentencePieceProcessor() + self.sp.LoadFromSerializedProto(arrs["TOKENIZER_MODEL"].tobytes()) + self.pad = 0 + + def __call__(self, prompt: str, max_len: int = COND_TOKENS): + ids = np.full((1, max_len), self.pad, np.int32) + mask = np.zeros((1, max_len), np.int32) + toks = self.sp.Encode(prompt or "")[:max_len] + ids[0, :len(toks)] = toks + mask[0, :len(toks)] = 1 + return ids, mask + + +# ── Conditioner (numpy port of sa3_pipeline) ──────────────────────────────── +def _expo_fourier(norm, dim=256, min_freq=0.5, max_freq=10000.0): + norm = np.asarray(norm, np.float32).reshape(-1, 1) + half = dim // 2 + ramp = np.arange(half, dtype=np.float32) / max(half - 1, 1) + freqs = np.exp(ramp * (math.log(max_freq) - math.log(min_freq)) + math.log(min_freq)) + args = norm * freqs * 2 * math.pi + return np.concatenate([np.cos(args), np.sin(args)], axis=-1).astype(np.float32) + + +class Conditioner: + """Loads cond.{padding_embedding,seconds_total_weight,seconds_total_bias}.""" + def __init__(self, npz=COND_NPZ): + z = np.load(npz) + self.pad = z["cond.padding_embedding"].astype(np.float32) # (768,) + self.W = z["cond.seconds_total_weight"].astype(np.float32) # (768, 256) + self.b = z["cond.seconds_total_bias"].astype(np.float32) # (768,) + + def seconds_embed(self, seconds, min_val=0.0, max_val=384.0): + s = np.clip(np.float32(seconds), min_val, max_val) + norm = (s - min_val) / (max_val - min_val) + ff = _expo_fourier([norm], dim=256) + return (ff @ self.W.T + self.b)[:, None, :].astype(np.float32) # (1,1,768) + + def build(self, last_hidden, mask, seconds): + """last_hidden (1,256,768), mask (1,256) -> cross (1,257,768), global (1,768).""" + m = mask.astype(np.float32)[..., None] + padded = last_hidden * m + self.pad.reshape(1, 1, -1) * (1 - m) + se = self.seconds_embed(seconds) # (1,1,768) + cross = np.concatenate([padded, se], axis=1).astype(np.float32) + gcond = se[:, 0, :].astype(np.float32) + return cross, gcond + + +# ── Pingpong schedule (numpy port; monotonic a2a rebuild for sigma_max<1) ──── +def _logsnr_shift(t, anchor=-6.2, end=2.0): + t = t.astype(np.float32) + logsnr = end - t * (end - anchor) + out = 1.0 / (1.0 + np.exp(logsnr)) + out = np.where(t <= 0, 0.0, out) + out = np.where(t >= 1, 1.0, out) + return out.astype(np.float32) + + +def build_pingpong_schedule(steps, sigma_max=1.0): + """(steps+1) sigmas from sigma_max down to 0. Warp the normalized [1->0] grid + through the logSNR shift, then scale by sigma_max — monotonic, first step + exactly at sigma_max. sigma_max=1.0 is plain text-to-audio (bit-identical to + upstream); sigma_max<1.0 is the audio-to-audio start.""" + t = _logsnr_shift(np.linspace(1.0, 0.0, steps + 1).astype(np.float32)) * np.float32(sigma_max) + t[0] = np.float32(sigma_max) + return t + + +# ── patch/unpatch (encoder patch grid) ────────────────────────────────────── +def patch_audio(audio: np.ndarray, patch_size: int = 256) -> np.ndarray: + """Patched-pretransform encode: (B, 2, T_audio) -> (B, 512, T_audio/256).""" + B, C, T = audio.shape + assert T % patch_size == 0, f"audio length {T} not a multiple of {patch_size}" + L = T // patch_size + x = audio.reshape(B, C, L, patch_size).transpose(0, 1, 3, 2) + return x.reshape(B, C * patch_size, L) + + +# ── Sampler (shared, numpy) ───────────────────────────────────────────────── +def make_noise(T_lat, steps, seed): + rng = np.random.default_rng(seed) + x0 = rng.standard_normal((1, 256, T_lat)).astype(np.float32) + step_noise = [rng.standard_normal((1, 256, T_lat)).astype(np.float32) for _ in range(steps)] + return x0, step_noise + + +def _cfg_velocity(x, tc, cond_v, uncond_v, cfg_scale, apg=0.0): + """Combine cond/uncond velocities in denoised space (RF), guide, map back. + Mirrors sa3_mlx.model_fn. apg>0 = Adaptive Projected Guidance (project the + cond-uncond diff orthogonal to cond_denoised). All fp32.""" + x = x.astype(np.float32) + sigma = np.float32(tc) + cond_d = x - cond_v.astype(np.float32) * sigma + uncond_d = x - uncond_v.astype(np.float32) * sigma + diff = cond_d - uncond_d + if apg <= 0.0 or cfg_scale < 1.0: + # cfg<1 is the interpolation regime (0 = pure uncond/neg branch); APG's + # orthogonal projection only applies to the extrapolation regime cfg>1. + cfg_diff = diff + else: + norm = np.sqrt((cond_d * cond_d).sum(axis=(-2, -1), keepdims=True)) + unit = cond_d / np.maximum(norm, 1e-8) + parallel = (diff * unit).sum(axis=(-2, -1), keepdims=True) * unit + diff_orth = diff - parallel + cfg_diff = diff_orth if apg >= 1.0 else (apg * diff_orth + (1.0 - apg) * diff) + cfg_d = cond_d + (cfg_scale - 1.0) * cfg_diff + if sigma == 0: + return np.zeros_like(x) + return ((x - cfg_d) / sigma).astype(np.float32) + + +def sample(dit_forward, x0, step_noise, sigmas, cross, gcond, + on_step=None, paste_back=None): + """Rectified-flow pingpong. dit_forward(x,t,cross,gcond)->v. + + paste_back=(init_lat, keep_mask): after every step restore the preserved + region (keep_mask 1=keep init, 0=regenerate) so inpainting leaves untouched + regions bit-exact.""" + steps = len(sigmas) - 1 + x = x0.copy() + for i in range(steps): + tc, tn = float(sigmas[i]), float(sigmas[i + 1]) + v = dit_forward(x, tc, cross, gcond) + denoised = x - tc * v + if i < steps - 1 and tn > 0: + x = (1 - tn) * denoised + tn * step_noise[i] + else: + x = denoised + if paste_back is not None: + init_lat, keep_mask = paste_back + x = init_lat * keep_mask + x * (1.0 - keep_mask) + if on_step: + on_step(i + 1, steps) + return x + + +def sample_cfg(cond_backend, x0, step_noise, sigmas, cross, gcond, null_cross, + cfg_scale, apg=0.0, batched=False, on_step=None, paste_back=None): + """Rectified-flow pingpong WITH classifier-free guidance. + + cfg_scale == 1.0 -> no uncond branch (identical to sample()). + cfg_scale != 1.0 -> per step, evaluate cond and uncond velocities and guide. + + batched is accepted for signature parity but must be False here — there is no + batch-2 C++ DiT, so CFG is a sequential dual-pass (two batch=1 forwards).""" + if cfg_scale == 1.0: + return sample(cond_backend, x0, step_noise, sigmas, cross, gcond, + on_step=on_step, paste_back=paste_back) + if batched: + raise ValueError("cpu-amx has no batch-2 DiT; call sample_cfg(..., batched=False)") + + steps = len(sigmas) - 1 + x = x0.copy() + for i in range(steps): + tc, tn = float(sigmas[i]), float(sigmas[i + 1]) + cond_v = cond_backend(x, tc, cross, gcond) + uncond_v = cond_backend(x, tc, null_cross, gcond) + v = _cfg_velocity(x, tc, cond_v, uncond_v, cfg_scale, apg) + denoised = x - tc * v + if i < steps - 1 and tn > 0: + x = (1 - tn) * denoised + tn * step_noise[i] + else: + x = denoised + if paste_back is not None: + init_lat, keep_mask = paste_back + x = init_lat * keep_mask + x * (1.0 - keep_mask) + if on_step: + on_step(i + 1, steps) + return x diff --git a/optimized/cpu-amx/scripts/sa3_cpu_amx.py b/optimized/cpu-amx/scripts/sa3_cpu_amx.py new file mode 100644 index 00000000..8c117ff7 --- /dev/null +++ b/optimized/cpu-amx/scripts/sa3_cpu_amx.py @@ -0,0 +1,476 @@ +"""SA3 text-to-audio inference on CPU via torch-free C++ AMX engines. + +The cpu-amx sibling of optimized/mlx/scripts/sa3_mlx.py and optimized/tensorRT. +Same CLI, same modes; the compute runs on Xeon AMX C++ engines instead of MLX: + + prompt -> T5Gemma (C++ AMX) -> numpy conditioner + -> DiT pingpong (C++ AMX int8 | bf16, MEDIUM) -> SAME-S/SAME-L decoder (C++ AMX) + -> WAV + +Modes (identical flags to the MLX / TensorRT releases): + text-to-audio --prompt P + audio-to-audio --prompt P --init-audio IN.wav [--init-noise-level sigma] + inpainting --prompt P --init-audio IN.wav --inpaint-range START,END + negative CFG --prompt P --cfg N [--negative-prompt P_NEG] [--apg S] + +cpu-amx specifics vs MLX: + * DiT is MEDIUM only, in int8 (default) or bf16 (--dit-precision). --dit sm-music / + sm-sfx are rejected with a pointer to optimized/tflite or optimized/mlx. + * MLX's per-model --dit-dtype splits into --dit-precision {int8,bf16} and + --decoder-precision {bf16,int8}. The int8 DiT is pinned to 1 thread (its .so + heap-races higher); the bf16 DiT (near-lossless, fp32 RoPE/RMSNorm islands) runs + at --threads. T5Gemma is bf16 regardless. + * audio-to-audio / inpainting init-encode is the torch-free C++ AMX SAME encoder; + the whole stack is 100% C++/numpy. +""" +from __future__ import annotations + +import argparse +import math +import os +import random +import subprocess +import sys +import time +from pathlib import Path +from shutil import which + +import numpy as np + +SCRIPTS = Path(__file__).resolve().parent +sys.path.insert(0, str(SCRIPTS)) + +import pipeline as P # noqa: E402 (vendored numpy pipeline) +import backends as B # noqa: E402 (C++ AMX engine loaders) + +SAMPLE_RATE = P.SAMPLE_RATE +SAMPLES_PER_LATENT = P.SAMPLES_PER_LATENT + +DIT_CHOICES = ["medium"] # cpu-amx: MEDIUM int8 core only +DIT_UNAVAILABLE = ("sm-music", "sm-sfx") # accepted for a clear rejection message +DECODER_CHOICES = ["same-s", "same-l"] +DEFAULT_DECODER = "same-l" # medium's native codec (mirrors MLX) + + +# ─── display helpers (ANSI colour when stdout is a TTY) — copied from sa3_mlx ─ +_USE_COLOR = sys.stdout.isatty() +_RULE_W = 64 + +def _c(code, s): return f"\x1b[{code}m{s}\x1b[0m" if _USE_COLOR else s +def bold(s): return _c("1", s) +def dim(s): return _c("2", s) +def cyan(s): return _c("36", s) +def yellow(s): return _c("33", s) +def green(s): return _c("32", s) +def magenta(s): return _c("35", s) + +def rule(char="━", color=cyan): + print(color(char * _RULE_W)) + +def banner(title): + rule(); print(f" {bold(title)}"); rule() + +def stage(idx_total, label, ms=None): + head = f" {cyan(idx_total)} {bold(label)}" + if ms is None: + print(head); return + visible = len(f" {idx_total} {label}") + fill = max(2, _RULE_W - visible - 9) + print(f"{head} {dim('·' * fill)} {yellow(f'{ms:>5.0f} ms')}") + +def sub(text): + print(f" {dim(text)}") + +def _peak_rss_mb(): + try: + import resource + return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1024.0 # KB->MB on Linux + except Exception: + return 0.0 + + +# ─── arrow-key picker (posix termios; numeric fallback off-TTY) — from sa3_mlx ─ +def _arrow_pick(prompt, options, default=None): + if not sys.stdin.isatty(): + print(prompt) + for i, o in enumerate(options): + print(f" {'*' if o == default else ' '} [{i}] {o}") + s = input(f"Choose [0-{len(options)-1}] (Enter for default): ").strip() + if s == "": + return default or options[0] + if s.isdigit() and 0 <= int(s) < len(options): + return options[int(s)] + return s if s in options else (default or options[0]) + import termios, tty + idx = options.index(default) if default in options else 0 + fd = sys.stdin.fileno() + old = termios.tcgetattr(fd) + print(prompt) + for _ in options: + print() + try: + tty.setcbreak(fd) + while True: + sys.stdout.write(f"\x1b[{len(options)}A") + for i, o in enumerate(options): + sys.stdout.write(f"\x1b[2K\x1b[36m▶ {o}\x1b[0m\n" if i == idx + else f"\x1b[2K {o}\n") + sys.stdout.flush() + ch = sys.stdin.read(1) + if ch == "\x1b": + seq = sys.stdin.read(2) + if seq == "[A": idx = (idx - 1) % len(options) + elif seq == "[B": idx = (idx + 1) % len(options) + elif ch in ("\n", "\r"): + return options[idx] + elif ch == "\x03": + raise KeyboardInterrupt + finally: + termios.tcsetattr(fd, termios.TCSADRAIN, old) + + +class _HelpfulParser(argparse.ArgumentParser): + def error(self, message): + sys.stderr.write(f"\nerror: {message}\n\n") + self.print_help(sys.stderr) + sys.exit(2) + def print_help(self, file=None): + super().print_help(file) + try: + from examples import print_example_commands + print_example_commands() + except Exception: + pass + + +def _play(path): + for player in ("ffplay", "aplay", "paplay", "play", "afplay"): + if which(player): + args = [player, "-autoexit", "-nodisp", path] if player == "ffplay" else [player, path] + try: + print(f" {bold('▶ playing')} {path} {dim('(Ctrl-C to stop)')}") + subprocess.run(args, check=False, + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) + except KeyboardInterrupt: + print() + return + print(f" {dim('(--play: no audio player found — install ffmpeg/alsa-utils to hear it)')}") + + +def main(): + """Top-level wrapper: run, then ALWAYS exit via os._exit so the DiT .so's + teardown double-free can never fire (it would mask the real error as a + confusing SIGABRT). Real errors are printed here and exit non-zero.""" + code = 0 + try: + _run() + except SystemExit as e: + c = e.code + if isinstance(c, str): # sys.exit("message") + print(c, file=sys.stderr); code = 1 + else: + code = int(c) if c is not None else 0 + except KeyboardInterrupt: + code = 130 + except Exception: + import traceback + traceback.print_exc(); code = 1 + sys.stdout.flush(); sys.stderr.flush() + os._exit(code) + + +def _run(): + ap = _HelpfulParser( + description="SA3 text-to-audio (+ audio-to-audio + inpainting) on CPU via C++ AMX engines", + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=( + "modes\n" + " text-to-audio --prompt P\n" + " audio-to-audio --prompt P --init-audio IN.wav [--init-noise-level σ]\n" + " inpainting --prompt P --init-audio IN.wav --inpaint-range START,END\n" + " negative CFG --prompt P --cfg N --negative-prompt P_NEG\n" + ), + ) + # Inputs + ap.add_argument("--prompt", default=None, + help="Text prompt. Empty string is valid (unconditional). " + "If omitted, asked interactively via stdin.") + ap.add_argument("--negative-prompt", default=None, + help="Negative prompt for CFG's uncond branch. No effect at --cfg=1.0. " + "When unset and --cfg≠1.0, the uncond branch uses the learned " + "padding embedding (all-zero prompt).") + ap.add_argument("--init-audio", default=None, + help="WAV (44.1 kHz stereo/mono) starting point. With --init-noise-level " + "= audio-to-audio; with --inpaint-range = inpainting. Init-encode " + "uses the torch-free C++ AMX SAME-{S,L} encoder (matched to --decoder). " + "Trimmed/padded to --seconds.") + ap.add_argument("--inpaint-range", default=None, + help="Inpaint span 'START,END' in seconds (needs --init-audio). The span is " + "regenerated; the rest is kept bit-exact (per-step paste-back).") + # Models + ap.add_argument("--dit", choices=DIT_CHOICES + list(DIT_UNAVAILABLE), default=None, + help="DiT model. cpu-amx ships MEDIUM only (the int8 C++ core). " + "sm-music / sm-sfx are not available here — use optimized/tflite or mlx.") + ap.add_argument("--decoder", choices=DECODER_CHOICES, default=None, + help="Audio decoder. 'same-l' = native 426M medium codec (default). " + "'same-s' = distilled 50M (faster, shares the medium latent space).") + ap.add_argument("--dit-precision", choices=["int8", "bf16"], default="int8", + help="DiT C++ engine precision. int8 (default, ~40 dB) or bf16 (near-lossless " + "fp32 RoPE/RMSNorm islands, ~59/54 dB @L1292/L4096, ~1.24× the int8 latency). " + "Both run at --threads. The cpu-amx analogue of MLX's per-model dtype.") + ap.add_argument("--decoder-precision", choices=["bf16", "int8"], default="bf16", + help="Decoder C++ engine precision. bf16 (default, best fidelity) or int8 " + "(SmoothQuant+GPTQ fused w8a8, smaller/faster). T5Gemma is bf16 regardless.") + ap.add_argument("--threads", type=int, default=16, + help="Thread count for T5Gemma + the decoder + the DiT (default 16). Both DiT " + "precisions run multi-threaded — at 1 thread the DiT is ~7.6× slower.") + ap.add_argument("--lora", action="append", nargs="+", default=None, metavar="ADAPTER", + help="(not supported in cpu-amx — accepted and ignored with a note; the int8 " + "C++ DiT core has no runtime LoRA merge. Use optimized/mlx for LoRA.)") + ap.add_argument("--lora-strength", type=float, default=1.0, help="(ignored in cpu-amx)") + # Sampling + ap.add_argument("--seconds", type=float, default=30.0, + help="Output length in seconds. T_lat = ceil(seconds*44100/4096) (decoder-" + "independent). Final WAV trimmed to exactly --seconds.") + ap.add_argument("--steps", type=int, default=8, + help="Pingpong sampling steps. Minimum 1 (single forward). The rf_denoiser is " + "distilled for 8 (default); >8 gives diminishing returns.") + ap.add_argument("--seed", type=int, default=None, + help="Random seed (any int). If omitted, chosen randomly and printed at the end.") + ap.add_argument("--init-noise-level", type=float, default=1.0, + help="σmax — the schedule's starting noise level. With --init-audio: 0.4–0.8 " + "typical for variation, 1.0 = full regeneration. Min 0.01.") + ap.add_argument("--cfg", type=float, default=1.0, + help="Classifier-Free Guidance scale. 1.0 = off (single forward). >1 pushes " + "toward the prompt; [0,1) toward the uncond/negative branch. Any value " + "≠1.0 costs ~2× per step (sequential cond + uncond forward).") + ap.add_argument("--apg", type=float, default=1.0, + help="Adaptive Projected Guidance [0..1], only when --cfg≠1.0. 1.0 = full APG " + "(project cond−uncond orthogonal to cond_denoised); 0.0 = vanilla CFG.") + # Runtime / output + ap.add_argument("--free-models", action=argparse.BooleanOptionalAction, default=True, + help="Free each model after its last use to lower peak RAM (default on).") + ap.add_argument("--out", "-o", default=None, + help="Output WAV path. Relative paths land in output/; absolute as-is. " + "16-bit PCM stereo @ 44.1 kHz, trimmed to --seconds. Auto-named if omitted.") + ap.add_argument("--play", action="store_true", + help="After writing, play the WAV (ffplay/aplay/paplay/afplay if present).") + args = ap.parse_args() + + if args.steps < 1: + ap.error(f"--steps must be ≥ 1 (got {args.steps})") + + # DiT selection — MEDIUM only. + if args.dit in DIT_UNAVAILABLE: + sys.exit(f"error: --dit {args.dit} is not available in cpu-amx (MEDIUM int8 C++ core only).\n" + f" Use optimized/tflite or optimized/mlx for sm-music / sm-sfx.") + if args.dit is None: + # Only one DiT, so pick the decoder interactively (matches MLX's picker feel). + args.dit = "medium" + if args.decoder is None: + if sys.stdin.isatty() and sys.stdout.isatty() and args.prompt is not None: + args.decoder = _arrow_pick("Choose audio decoder:", DECODER_CHOICES, default=DEFAULT_DECODER) + print(f" → {args.decoder}") + else: + args.decoder = DEFAULT_DECODER + if args.seed is None: + args.seed = random.randint(0, 2**31 - 1) + if args.prompt is None: + args.prompt = input("Prompt: ").strip() + if args.lora: + print(dim(" note: --lora is not supported in cpu-amx (int8 C++ DiT core has no runtime " + "LoRA merge) — ignoring. Use optimized/mlx for LoRA.")) + + # Output path + if args.out is None: + import re + slug = re.sub(r'[^a-z0-9]+', '_', args.prompt.lower()).strip('_')[:48] + args.out = f"{slug}_{args.seed}.wav" if slug else f"out_{args.seed}.wav" + out_path = Path(args.out) + if not out_path.is_absolute(): + out_path = SCRIPTS.parent / "output" / out_path + out_path.parent.mkdir(parents=True, exist_ok=True) + args.out = str(out_path) + + # T_lat: natural ceil (decoder-independent), matches MLX / TRT. + T_lat = max(1, math.ceil(args.seconds * SAMPLE_RATE / SAMPLES_PER_LATENT)) + target_dur = T_lat * SAMPLES_PER_LATENT / SAMPLE_RATE + + # Inpaint validation + latent range + inpaint_range = None + if args.inpaint_range is not None: + if args.init_audio is None: + sys.exit("error: --inpaint-range requires --init-audio") + try: + s_str, e_str = args.inpaint_range.split(",") + inp_start_sec, inp_end_sec = float(s_str), float(e_str) + except ValueError: + sys.exit(f"error: --inpaint-range must be 'START,END' seconds; got {args.inpaint_range!r}") + if not (0 <= inp_start_sec < inp_end_sec <= args.seconds): + sys.exit(f"error: invalid inpaint range {inp_start_sec}-{inp_end_sec}s " + f"(need 0 ≤ start < end ≤ {args.seconds}s)") + s0 = max(0, int(round(inp_start_sec * SAMPLE_RATE / SAMPLES_PER_LATENT))) + s1 = min(T_lat, int(round(inp_end_sec * SAMPLE_RATE / SAMPLES_PER_LATENT))) + inpaint_range = (s0, s1) + + sigma_max = float(args.init_noise_level) + mode = ("inpaint" if inpaint_range else + "audio-to-audio" if args.init_audio else "text-to-audio") + MIN_SIGMA = 0.01 + if sigma_max < MIN_SIGMA: + sys.exit(f"error: --init-noise-level={sigma_max} too low (min {MIN_SIGMA}); " + f"the model is undefined at t≈0.") + + t_wall = time.time() + print() + banner(f"SA3 → CPU-AMX {mode}") + k = lambda s: dim(f"{s:>12}") + v = lambda s, w=10: f"{s:<{w}}" + print(f" {k('prompt')} {bold(repr(args.prompt))}") + if args.negative_prompt: + suffix = "" if args.cfg != 1.0 else dim(" (ignored: --cfg=1.0)") + print(f" {k('neg prompt')} {bold(repr(args.negative_prompt))}{suffix}") + line = f" {k('dit')} {magenta(v(args.dit))} {k('decoder')} {magenta(v(args.decoder))}" + if args.init_audio: + line += f" {k('encoder')} {magenta(v(f'C++ {args.decoder}'))}" + print(line) + if args.init_audio: + print(f" {k('init audio')} {bold(args.init_audio)}") + if inpaint_range: + print(f" {k('inpaint')} {bold(f'{inp_start_sec:.2f}s..{inp_end_sec:.2f}s')} " + f"{dim(f'(latent {inpaint_range[0]}..{inpaint_range[1]} of {T_lat})')}") + print(f" {k('σmax')} {bold(f'{sigma_max:.2f}')}") + print(f" {k('seconds')} {v(f'{args.seconds}s')} {k('steps')} {v(args.steps)} {k('seed')} {args.seed}") + cfg_label = f"{args.cfg}" + (f" (apg={args.apg})" if args.cfg != 1.0 else "") + print(f" {k('dit prec')} {v(args.dit_precision)} {k('dec prec')} {v(args.decoder_precision)} " + f"{k('threads')} {v(args.threads)} {k('cfg')} {cfg_label}") + print(f" {k('T_lat')} {T_lat} {dim(f'({target_dur:.2f}s → trimmed to {args.seconds}s)')}") + print() + + # ── 1. T5Gemma encode ── + t0 = time.time() + tok = P.Tokenizer() + ids, mask = tok(args.prompt) + t5 = B.load_t5gemma(threads=args.threads) + last_hidden = t5(ids.astype(np.int32), mask.astype(np.int32)) # (1,256,768) + stage("[1/5]", "T5Gemma encode", (time.time() - t0) * 1000) + sub(f"last_hidden {last_hidden.shape} nnz(mask)={int(mask.sum())}") + + # ── 2. Conditioning ── + t0 = time.time() + cond = P.Conditioner() + cross, gcond = cond.build(last_hidden, mask, args.seconds) # (1,257,768),(1,768) + null_cross = None + if args.cfg != 1.0: + if args.negative_prompt: + n_ids, n_mask = tok(args.negative_prompt) + neg_hidden = t5(n_ids.astype(np.int32), n_mask.astype(np.int32)) + null_cross, _ = cond.build(neg_hidden, n_mask, args.seconds) + else: + # learned-padding uncond: conditioner on all-zero hidden+mask + null_cross, _ = cond.build(np.zeros((1, 256, 768), np.float32), + np.zeros((1, 256), np.int32), args.seconds) + stage("[2/5]", "Conditioning", (time.time() - t0) * 1000) + sub(f"cross {cross.shape} global {gcond.shape}" + + (f" null_cross ({'neg prompt' if args.negative_prompt else 'learned padding'})" + if null_cross is not None else "")) + if args.free_models: + del t5 # T5Gemma no longer needed + + # ── 3a. (a2a / inpaint) encode init audio → init_latents (C++ AMX SAME-{S,L} encoder) ── + init_latents = None + if args.init_audio: + stage("[3a]", f"Encode init audio → latents (C++ AMX {args.decoder} encoder)") + t0 = time.time() + enc = B.load_encoder(args.decoder, precision=args.decoder_precision, threads=args.threads) + audio_in = P.read_wav(args.init_audio) # (2, N) + init_latents = enc.encode(audio_in, T_lat) # (1,256,T_lat) + sub(f"device={enc.device} {(time.time()-t0)*1000:.0f} ms latents {init_latents.shape}") + if args.free_models: + del enc + + # ── 3b. DiT load + pingpong sample ── + stage("[3/5]", f"DiT — load + sample ({args.dit_precision}, {args.steps} steps, σmax={sigma_max:.2f})") + t0 = time.time() + dit = B.load_dit(precision=args.dit_precision, threads=args.threads) + sub(f"load {time.time()-t0:.1f}s (medium {args.dit_precision} C++ core, {args.threads} thread{'s' if args.threads != 1 else ''})") + + sigmas = P.build_pingpong_schedule(args.steps, sigma_max=sigma_max) + sub("schedule " + " · ".join(f"{float(x):.3f}" for x in sigmas)) + + x0, step_noise = P.make_noise(T_lat, args.steps, args.seed) + if init_latents is not None and inpaint_range is None: + x0 = init_latents * (1.0 - sigma_max) + x0 * sigma_max # a2a init mix + sub(f"init: latent * {1-sigma_max:.2f} + noise * {sigma_max:.2f}") + + paste_back = None + if inpaint_range is not None: + s0, s1 = inpaint_range + keep = np.ones((1, 1, T_lat), np.float32); keep[:, :, s0:s1] = 0.0 + paste_back = (init_latents.astype(np.float32), keep) + sub(f"inpaint mask {s0}..{s1} of {T_lat} ({(s1-s0)/max(T_lat,1)*100:.0f}% regenerated); " + f"paste-back keeps the rest bit-exact") + + def _on_step(i, n): + if not _USE_COLOR: + return + bar_w = 20; filled = int(round(bar_w * i / n)) + bar = cyan("█" * filled) + dim("·" * (bar_w - filled)) + sys.stdout.write(f"\r\x1b[K {dim('sampling')} {bar} {bold(f'step {i}/{n}')}") + sys.stdout.flush() + + t0 = time.time() + if args.cfg == 1.0: + latents = P.sample(dit, x0, step_noise, sigmas, cross, gcond, + on_step=_on_step, paste_back=paste_back) + else: + latents = P.sample_cfg(dit, x0, step_noise, sigmas, cross, gcond, null_cross, + cfg_scale=args.cfg, apg=args.apg, batched=False, + on_step=_on_step, paste_back=paste_back) + sample_ms = (time.time() - t0) * 1000 + if _USE_COLOR: + sys.stdout.write("\r\x1b[K") + if not np.isfinite(latents).all(): + sys.exit("error: DiT produced non-finite latents (try a different seed or σmax)") + sub(f"sample {sample_ms:.0f} ms ({sample_ms/max(args.steps,1):.0f} ms/step) " + f"latent {latents.shape}") + if args.free_models: + del dit # the DiT crashes on teardown anyway; os._exit(0) skips it + + # ── 4. Decode → audio ── + stage("[4/5]", f"Decoder ({args.decoder}, {args.decoder_precision} C++ AMX)") + t0 = time.time() + dec = B.load_decoder(args.decoder, args.decoder_precision, threads=args.threads) + audio_np = dec.decode(latents) # (2, T_lat*4096) + stage("[4/5]", "decode", (time.time() - t0) * 1000) + sub(f"audio {audio_np.shape}") + + # ── 5. Trim + write WAV ── + t0 = time.time() + requested = int(round(args.seconds * SAMPLE_RATE)) + if audio_np.shape[-1] > requested: + audio_np = audio_np[..., :requested] + P.save_wav(args.out, audio_np) + peak = float(np.abs(audio_np).max()); rms = float(np.sqrt((audio_np ** 2).mean())) + stage("[5/5]", "write WAV", (time.time() - t0) * 1000) + sub(f"audio {audio_np.shape} peak {peak:.3f} rms {rms:.3f}") + + total = time.time() - t_wall + audio_dur = audio_np.shape[-1] / SAMPLE_RATE + print() + rule() + print(f" {bold(green('done'))} {bold(f'{total:.2f}s')} wall → {audio_dur:.1f}s audio → " + f"{bold(yellow(f'{audio_dur/max(total,1e-9):.2f}× realtime'))} " + f"{dim(f'peak RSS {_peak_rss_mb()/1024:.2f} GB')} {dim(f'seed {args.seed}')}") + abs_out = os.path.abspath(args.out) + print(f" {bold(green('▸ saved'))} {bold(abs_out)}") + rule() + + if args.play: + _play(args.out) + # main() exits via os._exit — the DiT .so double-frees at teardown, so we + # must never let the interpreter run global destructors. + + +if __name__ == "__main__": + main() diff --git a/optimized/cpu-amx/scripts/sa3_gradio.py b/optimized/cpu-amx/scripts/sa3_gradio.py new file mode 100644 index 00000000..15ec400c --- /dev/null +++ b/optimized/cpu-amx/scripts/sa3_gradio.py @@ -0,0 +1,314 @@ +"""SA3 cpu-amx — gradio web UI (torch-free C++ AMX engines on CPU). + +The cpu-amx sibling of optimized/mlx/scripts/sa3_gradio.py, with every generation +mode wired: text-to-audio, CFG + negative prompt + APG, audio-to-audio, and +inpainting. Each clip renders as a 3-band tinted stereo mel spectrogram (numpy +port — no torch) with a click-to-seek playhead; only one clip plays at a time. + +Layout mirrors the MLX app (model/decoder/precision row, prompt+seed, +seconds/steps/cfg, Advanced = apg/σmax/negative prompt, Audio-to-audio and +Inpainting accordions, Output options) minus the Apple-only bits (LoRA, +Infinite-Radio/Hotswap) that don't apply here. + +Robustness: the int8 C++ DiT core heap-corrupts if it is invoked at more than one +sequence length in a single process (and double-frees at teardown). So each +generation runs the tested CLI (scripts/sa3_cpu_amx.py) in a FRESH subprocess — +one T_lat per process, os._exit(0) before teardown — which makes arbitrary +seconds/steps changes safe. Models reload per generation (mmap, a few seconds). + +Launch: + ./sa3-gradio # same-l default, public share link + ./sa3-gradio --decoder same-s + ./sa3-gradio --no-share # local only +""" +from __future__ import annotations + +import argparse +import base64 +import html as html_lib +import math +import subprocess +import sys +import tempfile +import time +import urllib.parse +import wave +from pathlib import Path + +import numpy as np + +SCRIPTS = Path(__file__).resolve().parent +ROOT = SCRIPTS.parent +sys.path.insert(0, str(SCRIPTS)) + +from spec import render_spectrogram_png # noqa: E402 + +CLI = str(SCRIPTS / "sa3_cpu_amx.py") +SAMPLE_RATE = 44100 +SAMPLES_PER_LATENT = 4096 +OUTPUT_DIR = ROOT / "output" / "gradio" +OUTPUT_DIR.mkdir(parents=True, exist_ok=True) + +DECODER_CHOICES = ["same-l", "same-s"] +PRECISION_CHOICES = ["bf16", "int8"] # decoder quantization +DIT_PRECISION_CHOICES = ["int8", "bf16"] # DiT quantization (int8 default/shipped, bf16 near-lossless) +MAX_SECONDS = 380 # medium's trained max +MIN_SIGMA = 0.01 + + +# ── run one generation via the CLI subprocess (DiT-crash isolation) ───────── +def run_generation(decoder, precision, dit_precision, threads, prompt, negative_prompt, seconds, + steps, seed, cfg, apg, sigma_max, a2a_path, inpaint_path, + inp_start, inp_end): + """Returns (audio_np (2,T) float32, info dict). Raises on failure.""" + out = Path(tempfile.gettempdir()) / f"sa3cpuamx_gr_{time.time_ns()}.wav" + cmd = [sys.executable, CLI, "--dit", "medium", + "--dit-precision", str(dit_precision), + "--decoder", str(decoder), "--decoder-precision", str(precision), + "--threads", str(int(threads)), "--prompt", prompt or "", + "--seconds", str(float(seconds)), "--steps", str(int(steps)), + "--seed", str(int(seed)), "--cfg", str(float(cfg)), "--apg", str(float(apg)), + "--out", str(out)] + # init audio: inpaint reference takes priority for the encode; a2a guide otherwise + init_audio = inpaint_path or a2a_path + if a2a_path and not inpaint_path: + cmd += ["--init-audio", a2a_path, "--init-noise-level", str(float(sigma_max))] + elif init_audio: + cmd += ["--init-audio", init_audio] + if cfg != 1.0 and negative_prompt and negative_prompt.strip(): + cmd += ["--negative-prompt", negative_prompt.strip()] + if inpaint_path and inp_end > inp_start: + cmd += ["--inpaint-range", f"{float(inp_start)},{float(min(inp_end, seconds))}"] + if a2a_path: + cmd += ["--init-noise-level", str(float(sigma_max))] + elif not a2a_path and sigma_max != 1.0: + cmd += ["--init-noise-level", str(float(sigma_max))] + + t0 = time.time() + proc = subprocess.run(cmd, capture_output=True, text=True, timeout=1200, cwd=str(ROOT)) + wall = time.time() - t0 + if proc.returncode != 0 or not out.exists(): + last = (proc.stderr.strip() or proc.stdout.strip() or "(no output)").splitlines()[-1] + raise RuntimeError(last[:200]) + + with wave.open(str(out), "rb") as w: + nch, sr, n = w.getnchannels(), w.getframerate(), w.getnframes() + pcm = np.frombuffer(w.readframes(n), np.int16).reshape(-1, nch).T.astype(np.float32) / 32767.0 + out.unlink(missing_ok=True) + T_lat = max(1, math.ceil(seconds * SAMPLE_RATE / SAMPLES_PER_LATENT)) + info = {"wall": wall, "T_lat": T_lat, "samples": pcm.shape[-1], + "realtime": (pcm.shape[-1] / SAMPLE_RATE) / max(wall, 1e-9), "seed": int(seed)} + return pcm, info + + +# ── HTML player (spectrogram + seekable audio, one-at-a-time playback) ─────── +# On play, pause every other