Repository navigation
vulkan : GDN conv straight from x and the state, and a short-row RMS norm with two fusions - #161
Merged
Merged
Conversation
…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
approved these changes
Oct 9, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Two prefill changes for the gated delta net layers (Qwen3.5/3.6, Flash-Next):
[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=0turns it off.GGML_VK_NORM_SMALL=0turns it off;GGML_VK_FUSE_NORM_EXTRA=0keeps 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: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_batchcases (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