---
name: flash-attn-training
description: How HOT-Step's custom flash-attention training ops (GGML_OP_FLASH_ATTN_TRAIN/_BACK) work, what the AS1.5 DiT trainer campaign proved and disproved, and the exact contract for porting flash mode to the other trainers (AS1.5 LM, MM3 LM, MM3 DiT). Use when adding --attn flash to any ace-train subcommand, touching engine/ggml/src/ggml-cuda/fattn-train.*, changing a trainer's VRAM model, debugging "flash is slower/uses more VRAM than expected", or interpreting any flash-vs-exact measurement.
---

# Flash-attention training (the fused backward) — adoption playbook

Written 2026-09-02 from the AS1.5 DiT campaign (commits 28ca16d3 → 10c37556).
Everything here was measured on an RTX 5090 (32 GB, sm_120) unless it says otherwise.
The deep docs are **gitignored, local-only** (`docs/plans/2026-09-01-flash-attn-backward.md`,
`fattn-train-spec.md`, `fattn-train-tf32-design.md`); this skill is the committed distillation.

**Context for a reader with zero prior exposure:** ggml's autodiff had no attention backward,
so every trainer built attention as `mul_mat → soft_max_ext → mul_mat` and retained the
`[S,S,Nh]` softmax per layer for the backward — the O(S²) term that capped DiT training crops
at ~50 s of audio on 32 GB. We wrote our own fused forward+backward ops (CPU reference + CUDA
TF32 kernels), carried as commits on the HOT-ggml fork that `engine/ggml` pins (see `docs/dev/ggml-fork.md`). Attention memory is now
linear in S; the DiT auto-fit picks full-song crops. Rob ear-validated the first flash-trained
adapter as "fantastic".

## 1. What exists

| Piece | Where | Notes |
|---|---|---|
| Ops `GGML_OP_FLASH_ATTN_TRAIN` / `_BACK` | `engine/ggml/include/ggml.h`, `src/ggml.c` (constructors, view getters, **autodiff case**), `src/ggml-cpu/ops.cpp` (f32 reference), `src/ggml-backend-meta.cpp` | Appended at the END of the op enum. Forward output is ONE packed tensor: O `[D,Nh,S,B]` then LSE `[Nh,S,B]`; `ggml_flash_attn_train_get_o()` views O out. Backward packs dQ\|dK\|dV. |
| CUDA kernels | `engine/ggml/src/ggml-cuda/fattn-train.cu/.cuh` (NEW files — never touch the inference `fattn-*.cu/.cuh`) | Scalar f32 v1 kernels kept as strict mode + pre-sm_80 fallback; TF32 mma (m16n8k8) kernels are the default. Bitwise-deterministic in every mode: no fp atomics, fixed schedules. |
| Precision knob | `ggml_flash_attn_train_set_prec/get_prec` (op_params slot 3) | `GGML_PREC_DEFAULT` (= 0 = zero-init!) → TF32 on sm_80+; `GGML_PREC_F32` → v1 scalar. Autodiff copies the forward's prec onto the backward node. |
| Where it lives | HOT-ggml `hot-step-neutral` commits `flash-attn-train` (+ `alloc-free-blocks`), pinned by `engine/ggml` ([docs/dev/ggml-fork.md](../../../docs/dev/ggml-fork.md)) | Change the kernels on the fork, then move the gitlink. `verify-hooks.ps1` Hook 12/13 grep the markers; Hook 17 refuses a dirty `engine/ggml`. |
| DiT trainer surface | `ace-train train-dit --attn exact\|flash\|flash-f32` (default `exact` in the CLI; the Training Studio form defaults to `flash`) | `dit_attn_flash()` in `engine/src/train/dit-train-graph.h` beside the untouched `dit_attn_f32()`. Both self- and cross-attention route through it. |
| Parity harness | `engine/tools/fattn-train-test.cpp`, target `fattn-train-test` | `--backend cpu\|cuda\|vulkan` (default `cpu`), resolved by registry device name — never "first GPU found", and a missing requested backend is a hard failure, never a silent CPU fallback. `--prec f32\|tf32` (tf32 is CUDA-only; Vulkan runs f32 at 1e-4 same as CPU), `--extra`, `--large`, `--bench`, `--bench-tr` (CUDA-only; rejects `--backend vulkan`). |
| Profilers | `--profile-step N` (coarse buckets); `DIT_PROFILE_NODES=1` per-node with site attribution (`engine/src/train/dit-node-profile.h`) | Node profiler is env-gated, zero cost when off. |
| Server/UI | `attnBackend: 'exact'\|'flash'\|'flash-f32'` through `types.ts` → `routes/training.ts` → `aceTrain.ts`; Training Studio checkbox | `cropMax 0` = "no pin" end to end (see trap 6). |

## 2. The adoption contract (non-negotiable, proven necessary)

Every trainer that gains flash mode must keep all of these. Each one exists because its
absence bit us.

1. **Per-trainer mode flag, default `exact`, and exact means byte-identical.** With the flag
   off the emitted graph must be the pre-flash graph to the byte — the DiT proves it with T3
   (`0.00e+00` on 17 named taps) and SC1–SC3 (`0.000e+00` grad delta) against a
   reverted-tree baseline. Gate the mode at the attention call sites only; restructure nothing else.
2. **A `supports_op` probe at trainer init, hard error on false.** `ggml_backend_supports_op`
   returning false is NOT an error in this engine: `backend_sched_new` registers the CPU backend
   alongside CUDA, so the scheduler silently splits attention onto the CPU — correct, unusably
   slow, low VRAM, tripwire silent, i.e. indistinguishable from a pass on every number the run
   reports. Build a scratch no_alloc node pair at the run's REAL shapes (both attention sites,
   effective Nkv) and abort with a named error. See DiT `dit-train-run.h` "spec 9.8 probe".
3. **A parity/selftest rung, exact vs flash, on CPU f32.** Gate on the CPU backend where both
   arms are f32; CUDA exact-vs-flash deltas (~3e-3) size the *reference's* cuBLAS TF32 rounding,
   not the fused op. Also gate the CUDA supports_op result so a silently-CPU flash arm can't pass.
4. **A measured drift class, documented like `--bwd mm`.** DiT: over 200 same-seed epochs flash
   drifted *less* than `--bwd mm`. Not identity — never claim identity.
5. **That trainer's VRAM model taught the flash branch** — otherwise the auto-fit keeps pricing
   the retained softmax and the flag buys nothing. See §5.
6. **Record the RESOLVED precision** (`attn_prec`) in the run log, not just the requested mode.
   Reason: op_params zero-init == `GGML_PREC_DEFAULT`, so every `--attn flash` run on Ampere+
   was ALREADY TF32 before the knob existed and said nothing about it.

## 3. Per-trainer porting checklist

Adoption is call-site wiring, not kernel work. The ops take any additive F16 mask
(`[S_kv,S]` or `[S_kv,S,1,B]` broadcast), GQA (Nkv < Nh at B=1), S_kv ≠ S, and
non-contiguous q/k/v views (only `nb[0]==4` is required — do not `ggml_cont` them, that
gives back the VRAM win).

| Trainer | Files | Specifics |
|---|---|---|
| **AS1.5 LM** (R2) — **DONE 2026-09-02** | `engine/src/train/lm-graph.h`, `lm-train-run.h`, `lm-vram.h`, `lm-selftest.h`, `flash-prec.h` | Causal = one triangular −INF mask; the kernel skips all-−INF tiles, so causal gets ~half its compute skipped free. Qwen GQA at B=1 is the tested path. Ships **off by default** (CLI and Training Studio checkbox); 4B low-VRAM is 5.5% faster than the shipped head-blocked arm and 1.2% faster at equal graph shape (opposite split from the DiT — there the fused kernel is the whole win, here the head-block copies are); naive 0.6B roughly doubles auto-fit `maxLen`, 1.7B only 1.27×. Not yet ear-validated — see `project-flash-attn-backward.md` in memory and §7/§8 below for the full numbers and open items. |
| **MM3 LM** (R3) — **DONE 2026-09-05** | `mm3-lm-train-run.h`, `mm3-lm-adapter.h`, `mm3-lm-graph.h` | The "sequence term was quadratic all along" retained softmax is what goes — but refused rather than composed with a frozen/trained KV prefix: `--attn flash` is rejected together with `--prefix-frames > 0` or `--prefix-n > 0` (the fused kernel doesn't take the rectangular mask a prefix needs), so the no-dK/dV-for-frozen-columns idea above was never built. **Default `exact`.** Measured (RTX 5090, `mm3-lm-f16`/`mm3-lm-q8_0`, `oasis_morningglory`, rank 256, checkpointed): flash is within noise of exact up to ~1500 frames (checkpointing already hides the small softmax in allocator slack), then saves VRAM growing to ~9 GB by 5000 frames; the usable crop ceiling moves from ~4300 frames (exact, before it starts spilling past ~29 GB used) to at least 11,178 frames (flash — this corpus's longest track, no OOM reached). Paired 20-step run at the shipped recipe's crop (750): 2118 ms/step flash vs 2215 ms exact, max loss drift 1.9e-4. Resolves to tf32 on this card for `--attn flash`, f32 for `--attn flash-f32`. **Not ear-validated** — the shipped recipe still trains at crop 750, where flash measures no benefit. Full numbers: `docs/dev/training-internals.md` MM3 section. |
| **MM3 DiT** (R4) | `mm3-dit-train-*.h` | Bidirectional like the AS DiT; smallest win (shorter sequences). |
| **YuE2 joint (AITK)** — **already flash, always** | `engine/src/train/yue2-aitk-graph.h` `block()` | Every AR and NAR attention call is `ggml_flash_attn_train` (TF32; `YUE2_AITK_STRICT_F32` = F32). The quadratic chain exists only under `YUE2_AITK_DIAGNOSTIC_MATH_ATTENTION`. The log's `attention_forward: "tf32"` is this kernel. No port needed; the decoder window is `--nar-crop-frames` (2026-09-24). Don't mistake it for an exact-attention trainer (an agent did). |

For each: (a) sibling `xxx_attn_flash()` returning exactly the shape the manual chain returned;
(b) flag + log fields; (c) probe; (d) selftest rung; (e) VRAM branch; (f) drift A/B;
(g) `--bench-tr`-style measurement at that trainer's REAL geometries (see §4).

## 4. Measurement discipline (where every wrong conclusion came from)

- **Pair arms at equal graph shape.** "Exact vs flash at crop 1250" once compared exact
  auto-fit crop 820 against flash's 1250; pinning both to 1250 forced exact into 2 checkpoint
  segments. Paired properly (same S, same segments, back to back): **flash is ~8.5% SLOWER than
  exact per token at equal shape on the DiT.** Flash's win is the CROP it affords, not per-token
  speed. Any claim otherwise needs a paired, interleaved measurement.
- **Interleave and repeat.** This box drifts ~7% between invocations; run-to-run contention is
  ±10%. Only within-invocation paired ratios are readable. Use 3 runs per arm.
- **Bench the real geometries.** `--bench` (window mask only) flattered fused. The trainer has
  three: windowed self (fused 0.89× cuBLAS), full self with NO mask (1.18×), cross at
  S_kv = enc_S (1.52×). `--bench-tr` covers all three. Half the DiT layers are full attention
  (`layer_type = i % 2`) and get no tile skip.
- **Attribute before fixing.** `DIT_PROFILE_NODES=1` found the whole flash deficit is the
  cross-attention BACKWARD (67.8 vs 34.8 ms/step); self-attention is a wash, cross forward is
  2× faster. Root cause: both TF32 backward kernels split warps by output d-range and recompute
  the shared S/dP tiles (dK/dV 1.5×, dQ 1.67× the needed mma). A dQ role split measured −2.7%
  end to end → **reverted under a 3% bar**. dK/dV split is blocked by ~128 B of static shared
  memory at the 3-blocks/SM occupancy cliff. Recorded in the plan doc; not worked around.
- **A 3% end-to-end bar for kernel churn.** Isolated-kernel wins of 15–25% can be 1% of a step.
- **Loss-to-target speed ≠ quality.** The overnight sweep's fastest-to-0.5 config (LoRA r128,
  pinned short crop) is a step-cost win that inverts at long crops; the 0.5 proxy's leader
  changed three times between ma5 0.8 and 0.5. Ear tests decide; nothing trained in flash mode
  after the first adapter has been heard.

## 5. VRAM model rules

- **Estimate must over-predict, never under** (target +5–15%); the NVML tripwire and the
  high-water probe are the backstop, never the plan.
- The exact-mode arena polynomial hides an **enc_S** dependence in its linear coefficient. The
  flash branch (`dit_vram_arena_bytes_flash`) takes enc_S explicitly; cross-attention scales
  with `enc_S×S` and at crop 1250 exceeds self-attention's S² — "enc_S is small" was refuted.
- **Read the arena log line as the TOTAL.** A "4319 est vs 7824 measured" line that omitted the
  LoKR-apply term sent a whole refit chasing a non-existent under-prediction; the total was
  over-predicting 73%. The line now prints both terms — keep it that way in every trainer.
- Fits are per-adapter-graph: `DIT_FLASH_LOKR_RETENTION` (0.62) was fitted before the LoKR
  apply reorder and now over-predicts +16.5% (safe direction, ~one crop step unspent). **Owed
  refit**; the batch>1 term is B=1-fitted and over-conservative.
- The flash lift raises `crop_max` to the dataset's longest track ONLY when the user passed
  no `--crop-max`; `a.crop_max_user` is the pin flag. See trap 6.

## 6. Trap list

1. **ggml.h edits invalidate ~141 CUDA objects — ~1 h rebuild.** Batch header changes.
   New `.cu` files need a cmake re-configure (the ggml-cuda CMake globs `*.cu`).
2. **DLL locks.** A running `ace-server` holds `ggml-base.dll`/`ggml-cuda.dll`; any ggml change
   needs the app down (`/api/shutdown` or `dev-rebuild.bat`). `ace-train.exe` is NOT held, so
   trainer-only edits build with the app up. Never kill ace-server (Node respawns it).
3. **Packed-output alignment gap.** Autodiff builds the packed gradient as
   `ggml_scale(packed, 0)` + `ggml_acc(dO)`; garbage in the O→LSE alignment gap becomes NaN.
   Both CUDA and CPU forwards zero the gap explicitly. Zero-width at every tested geometry, so
   tests never see it — keep the memset.
4. **In-place SCALE hazard.** `ggml_scale` is in `ggml_op_can_inplace`; it is safe only because
   the packed tensor always has a view child. The backward asserts `dst->data != fwd->data`.
5. **GQA at B>1 cannot be parity-tested against the manual chain** (ggml MUL_MAT backward
   asserts on broadcast src0) — that is why `dit_expand_heads` exists. Flash mode skips the
   expansion (native GQA), which also disarms the CUDA `REPEAT_BACK` cap on `Nkv·max(S,enc_S)·B`.
   Measured: **batch 1 still wins** on throughput and loss.
6. **The server always emitted `--crop-max`**, which the engine treats as a user pin → the flash
   lift never fired from the UI. `cropMax 0` now means "omit the flag". Quality presets must not
   re-pin it in flash mode. Any new trainer flag with an engine-side "user set it" sentinel has
   this exact failure mode — check the arg emitter.
7. **The parity tool must seed the loss gradient with 1.0** (`ggml_set_loss` only allocates)
   and assert a non-zero reference gradient, or both arms compare 0 vs 0 and pass vacuously.
8. **`dit_sa_mask` never produces a fully-masked key column** (pad columns stay open for padded
   query rows) — use `dit_ca_mask` for the exactly-zero-gradient assertion.
9. **Fully-masked query rows**: the fused op defines O=0, LSE=0; `soft_max_ext` produces NaN.
   Exclude them from reference diffs, check them directly.
10. **TF32 A-operand lane map ≠ accumulator map.** mma.cuh's `tile<16,8,float>` is the C/D map;
    using it as the tf32 A operand gives deterministic garbage. Derive with a probe kernel.
11. **Patch files are LF; a scratch tree extracted under `core.autocrlf=true` is CRLF** and every
    hunk fails. Replay with `git -c core.autocrlf=false -c core.eol=lf archive`. Export patches
    hunk-filtered: several patches share `ggml.c` and `ggml-cuda.cu`.
12. **rocm-build.yml did not apply patches** until 1b7e50d5 — every workflow that builds the
    engine needs the apply loop now that the trainer references patch-provided symbols.
13. **`MAX_FREE_BLOCKS`** (ggml-alloc) was 256; LoKR dim 256 (19k-node graph) overflowed it.
    Now 1024 via `alloc-free-blocks.patch`. Inference-shared → smoke generation after touching.
14. **Workflows die with the VSCode/Claude process.** Long unattended runs need the window open;
    machine sleep is "never" on this box (checked).
15. **`--mirror bf16` means bf16 COMPUTE, not just bf16 storage** — it rounds activations and
    gradients at every trainable-layer GEMM, and the adapters it trains are audibly coarse
    ("bitty", Rob 2026-09-02). Use **`--mirror bf16-f32`**: same BF16 residency, an in-graph
    `ggml_cast` to F32 at each `mul_mat` site, and over 12 same-seed epochs on mika it is
    bit-identical to `--mirror f32` while `bf16` drifts to 7.8e-3. It costs ~180 MB of transient
    arena and ~25% step time against f32 at equal crop, and buys 2.5× the flash auto-fit crop
    (1542 vs 610). Only `--bwd mm` carries it — the `out_prod` fallback arm keeps the forward
    cast alive and silently spends the ~8 GB back.
16. **Disk.** Probe runs write adapters; a campaign filled D: to 2.4 GB free and artifacts were
    deleted for space. Clean scratch dirs between grid cells.
17. **A bench tool's own reference arm can be non-contiguous where the trainer's never is.**
    `fattn-train-test --bench-lm`'s `blocked` arm fed a `ggml_cont(view)` straight into the
    reference attention chain, whose backward hands back a transposed (non-contiguous) gradient
    — `GGML_OP_CONT`'s backward asserts on that and the tool produced no table at all. The
    trainer never hits it because a `ggml_reshape` always sits between the cont and the chain,
    and RESHAPE's backward re-conts. Fix: wrap each bench-arm tensor in a shape-preserving
    `ggml_reshape` too, so the bench pays the same backward copy the trainer pays. Any bench
    harness that hand-builds a reference graph needs to mirror the trainer's node shapes, not
    just its op sequence.
18. **`--max-len` filters, it does not truncate.** Songs longer than it are skipped outright, so
    `alloc_seq = min(max_len, longest SURVIVING sample)` — pinning a value above the whole
    corpus's longest track yields an empty dataset (`no-samples`), and a VRAM-model cell "at
    S=1024" is really whatever the longest surviving song happens to be. Pick the dataset for
    the S you want, then report the actual S; don't trust the flag to hit a number.
19. **The exact-mode naive auto-fit can pick a `maxLen` whose own `estMb` already exceeds free
    VRAM**, then die on `cudaMalloc` with a hard access violation (0xC0000005) instead of a
    clean `lm_fatal` — reproduces identically on a pre-flash binary, so it is not new. Root
    cause is the same non-attention polynomial (`c2f`/`c2h`) the flash branch's
    `naive_nonattn_scale` now corrects around; the exact-mode fix is owed (see §8) and needs its
    own gate since it moves every shipped run's `estMb`.

## 7. Numbers worth remembering (5090)

| Measurement | Value |
|---|---|
| Fused TF32 vs cuBLAS per site, fwd+bwd, window mask | 0.94× / 0.64× / 0.49× at S=625/1250/3000 |
| Same at the trainer's real geometries | windowed 0.89×, full-self 1.18×, cross(S_kv 1877) 1.52× |
| Attention VRAM per site at S=3000 | 487 MB fused vs 4.9 GB manual |
| Parity worst rel err | f32 3.5e-6 (bar 1e-4); tf32 4.7e-4 (bar 5e-3, floor 1e-5) |
| Flash vs exact drift, 200 same-seed epochs | smaller than `--bwd mm` |
| Done-gate auto-fit, production LoKR, unpinned | albumJ 1498 (enc_S 1877), album D 1616 (enc_S 640); LoRA r16 ~3400 |
| LoKR apply reorder | −10% step, LoKR:LoRA 1.35→1.21; the two copies are unavoidable, ~7% of step |
| 12 GB emulated card, flash+bf16+LoRA r16 | full 32-layer depth, crop 410, 4 segments |
| **LM, 4B low-VRAM, flash vs shipped (`exact --attn-head-block 8`)** | **5.5% faster/micro-step, 3.8% lower peak VRAM** (paired, interleaved, albumF substitute) |
| LM, 4B low-VRAM, flash vs equal-shape (`exact --attn-head-block 0`) | 1.2% faster — the head-block copies are almost the whole DiT-vs-LM difference |
| LM attention-only bound (`fattn-train-test --bench-lm` vs blocked) | 0.74×/0.79×/0.80× at S=1024/2113/3500 |
| LM naive auto-fit `maxLen` lift, flash vs exact | 0.6B ~2.0× (3136→6208 tok); 1.7B ~1.27× (2624→3328 tok) |
| LM 50-epoch same-seed drift, flash vs exact | same class as `--weights bf16`; smaller on 2/3 measures, ~20% larger on final CE (1 seed, no error bar) |
| **MM3 LM, usable crop ceiling, flash vs exact** | **~4300 frames exact -> >=11,178 frames flash** (this corpus's longest track; RTX 5090, `oasis_morningglory`, rank 256) |
| MM3 LM, paired step time at the shipped recipe's crop (750 frames) | 2118 ms/step flash vs 2215 ms exact |

## 8. Open items (as of 2026-09-02)

- R4 (MM3 DiT) port remains (this skill is its brief) — R2 (AS1.5 LM) and R3 (MM3 LM) are DONE, both off by default pending ear tests.
- Cross-attention backward kernel: dK/dV role split blocked by smem; a dQ split exists in
  the plan doc (reverted, −2.7%).
- `DIT_FLASH_LOKR_RETENTION` refit after the apply reorder; batch>1 VRAM term.
- Exact-mode arena polynomial under-predicts 13–18% (masked by LoKR over-count; fix gated to flash).
- **LM exact-mode `c2f`/`c2h` non-attention polynomial is ~2.2× light on the naive path**
  (−11.9% to −12.9% measured, same class as the DiT's exact-mode item above); the flash branch's
  `naive_nonattn_scale` corrects around it but the exact-mode fix itself is owed and needs its
  own gate, since it would move every shipped run's `estMb`/auto-fit `maxLen`.
- LM G5/G6 ran on `albumF`, not album I — the box has no `albumI*` tensor dir, and
  the plan's ear pair (G7) is specified on album I/E3 lineage. album I codes need Preprocess +
  Extract via the Training Studio batch pipeline before G7 can run as written.
- LM G7 ear test (twin album I adapters, staged in `_experiments/_LISTENING`) — not run, needs
  Rob; the flash checkbox stays off until it lands.
- Pre-existing bugs surfaced while porting R2, neither fixed (both reproduce on a pre-flash
  binary): `mm3-lm-train` crashes at export with a `ggml-backend.cpp` tensor-write-out-of-bounds
  assert; the LM exact-mode naive auto-fit can pick a `maxLen` that OOMs via access violation
  instead of a clean fatal (trap 19).
- Ear validation of anything trained since the first flash adapter, and of the LoKR reorder.
- Low-VRAM training profiles for users (B1) — deferred by Rob until the 32 GB path is nailed.
