Skip to content

vulkan : GDN conv straight from x and the state, and a short-row RMS norm with two fusions - #161

Merged
dzannotti merged 2 commits into
halo-box:masterfrom
Nathanw1014:hb/vk-gdn-conv-norm
Oct 9, 2026
Merged

dzannotti merged 2 commits into
halo-box:masterfrom
Nathanw1014:hb/vk-gdn-conv-norm

Conversation

@Nathanw1014

Copy link
Copy Markdown

Two prefill changes for the gated delta net layers (Qwen3.5/3.6, Flash-Next):

  • CONCAT + SSM_CONV + SILU, direct. The concat transposes x into a [T+3, C] tensor that the conv reads once, which is 84 MB written and read per layer on Flash-Next at a 2048-token ubatch. This runs the conv and SiLU straight from x and the state, and writes only the conv-state tail columns. It follows the CUDA backend's gdn_conv match, with no llama change. GGML_VK_SSMCONV_DIRECT=0 turns it off.
  • rms_norm_small. One subgroup per row for rows of 256 or fewer (the 128-wide GDN head norms), instead of a 512-thread workgroup with 3/4 of the lanes idle. It adds two fusions: RMS_NORM_MUL_MUL (the gated norm) and RMS_NORM_SCALE (the GDN l2 norm). GGML_VK_NORM_SMALL=0 turns it off; GGML_VK_FUSE_NORM_EXTRA=0 keeps the kernel without the fusions.

Strix Halo (gfx1151, RADV, Mesa 26.2.4), pp2048 -ub 2048, same binary with both switches off as the off arm, two runs per arm in alternating order:

model depth off on change
Qwen3.6-35B-A3B UD-Q4_K_XL 0 2151.1 2272.3 +5.6%
Qwen3.6-35B-A3B UD-Q4_K_XL 8k 1702.1 1779.8 +4.6%
Qwen3.8-Flash-Next REAP-320 0 785.8 821.1 +4.5%
Qwen3.8-Flash-Next REAP-320 8k 740.8 762.7 +3.0%

New CONCAT_SSM_CONV_SILU cases cover n_t 1..300, ragged channel and token tiles, two sequences, 1/4/8 rollback slots, 10240 channels, and an extra concat reader that must not fuse. The full test-backend-ops suite gives 34323/34326 with the changes on and with them off. The three failures are test_mul_mat_exact_batch cases (bf16 n=9 and 16, q8_0 n=4) that also fail on current master, so they are not from this PR.

Written with Claude, tested by me.

Assisted-by: Claude Opus 5.5

…CONV + SILU)

Gated DeltaNet layers build CONCAT(state, transpose(x)) -> SSM_CONV -> SILU,
plus copies of the last 3 concat columns into the recurrent conv state (one
per rollback slot). The concat transposes x into a [T+3, C] tensor that the
conv reads once: 84 MB written and read per layer for Flash-Next at a
2048-token ubatch.

This follows the CUDA backend's gdn_conv match (same graph, no llama change):
- at the CONCAT, when its only readers are one SSM_CONV (+ SILU right after)
  and conv-state tail copies, write only the tail columns those copies read;
- at that SSM_CONV, run a direct conv + SiLU: one thread per channel and 8
  tokens, the 4-tap window sliding in registers, lanes along channels, the same
  dot() as the existing nc == 4 path (identical output).

The decision is made once at the concat and consumed at its conv. The match
checks the allocated addresses (the SiLU output overlaps nothing the conv
reads, and nothing between the concat and the conv writes x or the state), and
graph_optimize asks the allocator to keep x and the state alive through the
SiLU. x and the state are also added to the barrier tracking. Ubatches under
8 tokens per sequence keep the concat (decode/verify unchanged).
GGML_VK_SSMCONV_DIRECT=0 restores the concat + SSM_CONV_SILU pair;
GGML_VK_FUSION_DEBUG=1 prints why a concat did not match.

test-backend-ops: new CONCAT_SSM_CONV_SILU cases (scheduler-allocated, verify
the SiLU output and every state copy): n_t 1..300, ragged channels and token
tiles, two sequences, 1/4/8 rollback slots (8 reaches into the state columns),
10240 channels, and a whole-concat extra reader that must not fuse.
Doubling the direct-conv store fails the fused cases + GDN_CONV_PREFILL;
doubling the tail store fails the fused cases' state copies; both pass with
GGML_VK_SSMCONV_DIRECT=0. Only tested on gfx1151 (RADV).

llama-bench pp2048 -ub 2048, same binary with the gate off/on, ABBA x2, -r 2
(mean of 8): Flash-Next REAP-320 778.4 -> 790.3 t/s (+1.5%), Qwen3.6-35B-A3B
2101.8 -> 2161.0 (+2.8%). Every gate-on sample beat every gate-off sample.
Flash-Next PPL c4096 (2 ubatches per chunk) identical per chunk.

Assisted-by: Claude Opus 5.5
… fusions

rms_norm.comp runs one 512-thread workgroup per row, so on 128-wide rows (the GDN
head norms) 3/4 of the lanes idle and every row pays a shared-memory tree reduction.
rms_norm_small.comp gives each row one subgroup (subgroupAdd), for ne00 <= 256,
f32 in/out, contiguous rows. It handles plain RMS_NORM and RMS_NORM_MUL, plus two
fusions only it has:

- RMS_NORM_MUL_MUL: the gated norm, normalised * gamma * gate (qwen3.5/qwen4exp
  build_norm_gated). The gate chain (sigmoid/silu of z) sits between the two MULs
  in the graph, so graph_optimize now schedules it ahead of the norm and keeps the
  three nodes adjacent. The products are the same f32 multiplies as the separate ops.
- RMS_NORM_SCALE: build_gdn_l2_norm, scale(rms_norm(x, eps/n), 1/sqrt(n)). The
  factor is applied on the store after rounding the normalised value to f32, the
  same two roundings as the SCALE op. graph_optimize keeps the SCALE behind its norm.

The kernel strides over rows, so the grid can be capped at 65535 workgroups and a
subgroup size other than the one the host assumed still covers every row.

GGML_VK_NORM_SMALL=0 restores the old path (and disables both fusions);
GGML_VK_FUSE_NORM_EXTRA=0 keeps the kernel but disables the two fusions.

gfx1151 (RADV), same binary with GGML_VK_NORM_SMALL off/on, ABBA x2, -r 2
(mean of 8): Flash-Next REAP-320 pp2048 -ub 2048 773.3 -> 785.0 t/s (+1.5%),
Qwen3.6-35B-A3B pp2048 -ub 2048 2100.4 -> 2145.5 (+2.1%), Qwen3.8-27B pp512
-ub 256 430.5 -> 431.6 (flat). Norm ops 640/640 in all three gate settings.
PPL moves only by the reduction order (subgroupAdd vs the shared-memory tree).
Only tested on gfx1151.

Assisted-by: Claude Opus 5.5
@dzannotti
dzannotti merged commit 5bca18d into halo-box:master Oct 9, 2026
7 of 8 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants