Skip to content

vulkan: shader and sequencing efficiency backlog (attention, reductions, staging reuse, first-submit latency) — unmeasured #475

Description

@geisten

Findings of a four-angle review of src/backends/vulkan/ (efficiency), left over after the cleanup PR. Nothing was measured — every cost is an estimate from reading the shaders and the sequencing code; measure (GEIST_VK_PROFILE=1, bench_perf_sweep) before and after each item. Complements #467 (Q4_0/Q5_K/TQ2_0/DeltaNet kernels) and #469 (pipeline cache, registry lookup, staging ring, act-quant fusion).

Decode attention (grows with context, every layer)

  1. V pass uses only hd/4 threads (attn_part_f16.comp:113-123, attention_f16.comp:127-137): 32 of 128 threads at hd = 128, each walking up to 128–256 positions with one dependent a += chain. → slice by lid % hd4, position group lid / hd4, 2–4 independent accumulators, merge through shared memory once per chunk. Est. 2–4× fewer serialized loads in that phase — probably the largest single shader gain.
  2. QK dot reads uncoalesced K rows (attn_part_f16.comp:80-85, attention_f16.comp:91-96): one thread per KV position loops over hd4 with adjacent lanes n_kv_heads·hd·2 bytes apart (2 KiB at 4 kv-heads, hd 256) → 25 % sector utilisation. → 8 lanes per position, 16-byte uvec4 loads, subgroupAdd tree; a warp covers 4 positions × 128 contiguous bytes. Est. 1.5–3× on the QK phase at long context (if L1-bound — unconfirmed).
  3. attn_comb (attn_comb.comp:27-59): thread 0 alone does two serial passes over n_chunks, then every thread recomputes exp(mc − m) for all chunks (hd · n_chunks redundant exps). → wc = exp(mc − m)·inv_l once per chunk into shared (max/sum by subgroupMax/Add), the d loop is FMAs only. Est. 50–150 µs/token at 4k context over 30 layers.

Reductions

  1. Shared-memory tree reductions with workgroup barriers where a subgroup reduction fits: attention softmax (~16–20 barriers per chunk), qkv_prep_* (9), matvec_q4k/q6k (6, inside a single 32-lane warp but with workgroup-wide barrier()), ffn_norm_gate_up_q4k, matvec_f32, ple_gate, argmax, act_quant. rmsnorm_f32.comp already shows the pattern (subgroupAdd + one tiny shared pass). Keep the tree as fallback when gl_SubgroupSize != 32 (see vulkan: GEMM kernels assume 32-lane subgroups — per-row matvec fallback on RADV/Intel #471). Est. 1–2 µs saved per short kernel, 10–30 % of the softmax phases, low single digit % of matvec time.
  2. matvec_q6k static shared memory ≈ 25.6 KiB (xsh[1536] vec4 + reduction) caps occupancy at about two 256-thread workgroups per SM, for the big ffn_down/output projections. The 8 warps read identical x addresses (L1 broadcasts) → drop the staging or size it with a specialization constant. Up to 2× resident warps; 10–30 % or nothing if latency is already hidden.

Sequencing

  1. linear_t/pair re-stage the same x (ops.c vk_linear_t, vk_linear_t_pair, resources.c vk_xring_stage): on the non-BAR path every consumer copies t_x into a new ring slot (copy, barrier, matvec, copy, barrier, matvec). → remember the last staged (buffer, offset, m, n_in, ring offset), reuse until a write to that range (vk_seq_hazard already knows writes). Q/K/V and gate/up share x: 1–2 copies + full-pipeline barriers saved per layer per token (est. 2–5 µs each, 30–100 per token on a 32-layer model).
  2. First submit only after 64 dispatches (VK_SEQ_ROTATE = 64, sequence.c:307): the GPU idles while the CPU encodes the first 64 dispatches of every token (est. 60–130 µs, 0.5–1 % of a 10 ms token). → geometric rotation (8, 16, 32, 64). Also vk_seq_dispatch_acc hard-flushes synchronously at 4000 dispatches (very deep models stall mid-token) → recycle command buffers behind the fence.
  3. Redundant per-dispatch calls: CmdBindPipeline even when the pipe is unchanged (last_pipe), ResetDescriptorPool(seq_pool) on every flush although the pool is only used on cache overflow, and read-only weight buffers (vk_acc_all(false)) occupy dirty[] slots (96) and force scans/conservative barriers. Tens of µs per token in total.

Prefill / other kernels

  1. matmul_q4k_cm staging (matmul_q4k_cm.comp:76-100): 16 scalar float16_t shared stores per k-step, ~2-way bank conflicts (ASTRIDE = 40 halves, row = lid>>1), Q4_K header re-decoded every BK = 32 step (8× per superblock), B converted f32→f16 in every one of the n_out/64 workgroups → packHalf2x16/uvec4 shared stores, BK = 64 or hoisted headers, convert x to f16 once in the xring copy. Est. 5–15 % on prefill GEMM; matmul_q4k_cm32/matmul_q6k_cm likely share it (not read).
  2. argmax is one workgroup over the whole vocabulary (argmax_f32.comp, gx = 1): ~600–1000 dependent compare iterations per thread on the per-token critical path → 32–64 workgroups with vec4 loads and subgroupMax + a second pass, or fold into the logits matvec epilogue. Est. 20–50 µs/token (<1 % of decode).
  3. Integer div/mod by uniform runtime values in hadamard_f32.comp (p / len, p % len per stage; 4 div/mod per element in the permutation), add/mul_f32.comp (3 div + 3 mod per strided element), rope_f32.comp (4 per pair) → ((p & ~(len − 1)) << 1) | (p & (len − 1)), one r = i / cols reused, and subgroupShuffleXor for the first 5 Hadamard stages (removes ~5 of 10 barriers). ALU-bound, maybe 20–40 % of these tiny kernels.
  4. matvec_f32 / ple_gate use scalar loads and one workgroup per row (matvec_f32.comp:29-33, ple_gate_f32.comp:29-33): vec4 loads (guard n_in % 4, offsets % 4, as matvec_q4k does) plus subgroupAdd; 10–25 % on the F32-weight and Gemma-3n PLE paths.

Acceptance

Each item: before/after numbers (kernel time from GEIST_VK_PROFILE=1 and pp/tg from bench_perf_sweep) in the PR, test_backend_vulkan_* green on RTX 2080 Ti + RADV.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    enhancementNew feature or request

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions