Skip to content

chore(train): engine-10 continued: the raw q4_0/q8_0 flash attention at depth, llama-bench -rs - #71

Merged
marcospaulo merged 5 commits into
mainfrom
train/engine-10
Sep 27, 2026
Merged

marcospaulo merged 5 commits into
mainfrom
train/engine-10

Conversation

@marcospaulo

Copy link
Copy Markdown
Member

The train engine-10 past #70: 5 commits, main (184e3a4) to 4b61c54. rig pins its tip, engine-4b61c54.

  • The raw q4_0/q8_0 flash attention at depth, without its bank conflicts (01f4f0f). At 65,536 cells on an RTX 5080 (Ternary Bonsai 2 27B's attention: head 256, 4 KV heads at GQA 6, a bit-packed mask), 3.43M of its 7.36M shared-memory wavefronts were bank conflicts.

    • The V tile's dequant stores across 8 rows: store conflicts 3.2M → 0.09M.
    • K·Q reads each block's f16 scale from the raw rows. The float tile, its pass and its barrier go: load conflicts 0.23M → 0.03M, and 26,272 bytes a block, so a decode token runs 3 blocks an SM.
    • The stream-k fixup loads the next 8 blocks' partials before folding: 6.9 → 5.9 us.
    • Padding columns write no fixup partial.
    • The next raw K tile loads beside V where 2 blocks share an SM.

    test-backend-ops perf on the served layout, 6 rounds against the parent's library, a token: -16.7 / -3.7 / -3.8 / -1.8 % at 16,384 / 65,536 / 131,072 / 245,760 cells; a 4-row verify: -6.7 / -3.6 / -2.9 / -2.8 %. GGML_CUDA_FATTN_Q4_0_LEGACY=1 / GGML_CUDA_FATTN_Q8_0_LEGACY=1 restore the stock kernels.

  • llama-bench -rs, the recurrent-state snapshots a draft's rollback keeps (7537f40). A drafting server writes n_max + 1 snapshots a verify; llama-bench wrote one. pp4 with -rs 3 at depth 16,384 on a 5080: -0.62 % against -rs 0.

  • A tensor split's ratio failure names its node (6b98c5d): the node, op, dims, axes, segments and source chain. Fatal path only.

  • TORAD.md rows for both (a9522f2, 4b61c54).

Gate on 4b61c54's prebuilt (RTX 5090, against engine-4104c47, the 245,752-token prompt's four questions greedy with top-5 log-probabilities, GATE PASS). The new occupancy moves a token's stream-k split (340 → 510 ways), so the as-served pairs are held to engine-10's G1 bars, not bit-exactness:

  • plain: 4 first differences, each at a tie; |dlogprob| p99 0.056, max 0.090;
  • drafted: 1, at a tie; p99 0.032, max 0.112;
  • without n_probs: the same tokens and drafting;
  • a 245K conversation swapped out and back: token for token, first token 1.03 s;
  • plain decode on the 245K question: 91.9 → 93.7 tok/s.

e2e-driver-only.sh on a 5070 Ti: PASS.

The bare assert printed nothing, so every tensor-split abort needed a
debugger to attribute. The ratio failure now aborts with the node name,
op, dst dims and axis, the src dims and axis, the src segments and the
two sides of the division, plus the src[0] ancestry chain. Fatal path
only; no steady-state cost.
…equant, K scales read raw, the fixup folded ahead

Ternary Bonsai 2 27B's attention (q4_0 K/V read raw by the MMA kernel, head 256, 4 KV heads at GQA 6, a bit-packed
mask) under ncu at 65,536 cells on an RTX 5080: 2 blocks of 2 warps an SM at 43,168 bytes of shared memory, 23 % issue
active, and 3.4M of its 7.4M shared-memory wavefronts bank conflicts. Five changes, none to the arithmetic:

- The V tile's dequant: 8 consecutive threads take one block pair of 8 consecutive rows, so a 16-byte store phase lands
  in 8 distinct bank quads (the row stride is an odd number of 16 bytes); 8 threads on 2 rows made it 4-way. Store
  conflicts 3.2M -> 0.09M.
- K*Q reads each block's f16 scale from the raw rows. The float scale tile, its pass, its barrier and its 4-way load
  conflicts are gone (0.23M -> 0.03M), and raw K keeps no shared memory beside its raw tile: 26,272 bytes a block,
  3 blocks an SM for a decode token.
- The stream-k fixup loads the next 8 blocks' partials before folding the current ones, in the order and with the
  formulas every fixup kernel used: 6.9 -> 5.9 us at 65,536 cells (ncu).
- A padding column (10 of a decode token's 16: 2 rows x 8 heads for 1 x 6) writes no partial for the fixup, which
  never read one: 6 KB a block instead of 16 as the blocks finish.
- Where 2 blocks share an SM (the 4-warp tiles of a verify), the next raw K tile loads as soon as K*Q has read the
  current one, beside this step's V (cp.async commit/wait groups); where 3 do, K loads after V arrives, as before.
  Early everywhere, a decode token ran 1-2 % slower at 16,384 and 245,760 cells; after V everywhere, a verify ran 1-3 %
  slower at 65,536.

test-backend-ops perf gains the served attention at 16,384 / 65,536 / 131,072 / 245,760 cells for a token and a
4-row verify, K and V laid out as the cache's views are (a head's cells 576 bytes apart).

Measured against the parent's CUDA library under one test-backend-ops, 6 rounds, the order rotated, medians:
  a token:        16,384 -16.7 %  65,536 -3.7 %  131,072 -3.8 %  245,760 -1.8 %
  a 4-row verify: 16,384  -6.7 %  65,536 -3.6 %  131,072 -2.9 %  245,760 -2.8 %
llama-bench as served (q4_0 K/V, f16 state, graphs on), 4 rounds against the parent's libraries, the host swapping
under other sessions' builds: tg64 +1.2 % at 65,536 (91.19 -> 92.30 tok/s) and +0.5 % at 16,384 (102.61 -> 103.15);
pp4, the verify, -0.5 % and -0.1 %, inside that noise (samples of both arms fell to 92 tok/s). The attention's share of
the step predicts +0.6, +0.7, +0.5 and +0.3 %.

Checks: FLASH_ATTN_EXT 3,220/3,220. Greedy 64 tokens after a 7K-token prompt byte-identical to the parent. Each change
but one writes the same values to the same addresses; 3 blocks an SM splits a decode token's stream-k work 252 ways
instead of 168 on this card, the numerics of another card's split.
What ncu found at 65,536 cells on an RTX 5080 (3.43M of 7.36M shared-memory wavefronts bank conflicts), the five changes,
test-backend-ops perf on the served layout and llama-bench as served, and the raw path's off switches.
A drafting server runs its context with n_rs_seq = its draft's n_max, so every verify of a hybrid model writes
n_rs_seq + 1 snapshots of each recurrent layer's state and conv window (the rollback slots). llama-bench left
n_rs_seq at 0: its pp4 measured a verify that writes one. `-rs` is a parameter axis like `-cts`, and a field in every
printer.

A context clamps its batch to its size, and split_equal keeps a sequence's last n_rs_seq + 1 rows in one ubatch,
which must hold more than that; the bench sized its context at n_prompt + n_gen + n_depth, so pp4 with -rs 3 at depth
0 aborted in split_equal. The context now holds at least n_rs_seq + 2 rows.

Measured on the RTX 5080, Ternary Bonsai 2 27B (q4_0 K/V, f16 state), pp4 in a 512 ubatch at depth 16,384, 4 rounds
alternating, 32 samples each: -rs 0 386.15 and -rs 3 383.76 tok/s (-0.62 %: the served draft's four snapshots a layer
cost that much of a verify). pp4 with -rs 3 at depth 0 runs (350.69 tok/s).
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant