SIGN IN SIGN UP

[ET-VK][sdpa] Reuse shared V cache across GQA query heads in AV coop-GEMV

Pull Request resolved: https://github.com/pytorch/executorch/pull/21063

The LLM decode AV coop-GEMV reloads the shared V cache once per query head. In
grouped-query attention Hq = G * Hkv query heads share each KV head (Llama G=4,
Phi G=3, Qwen G=2), and out[q_h, d] = sum_c attn[c, q_h] * V[c, kv_h, d] reads
the SAME V texel for every query head in a group. The per-query-head coop shader
gives each of the Hq heads its own workgroup, so V -- the dominant traffic
(head_dim-wide per context texel, vs a scalar attn weight) -- is read G times.

This adds a GQA-reuse AV variant that assigns ONE workgroup per (d4, kv_h): it
loads each V texel once and reuses it across all G query heads in the group,
producing G output texels. For this bandwidth-bound kernel that cuts V-cache
traffic ~Gx.

Implementation:
- The variant is a codegen flag (`GQA`) on the existing
  `sdpa_compute_out_coop.glsl` template, not a separate file: one shared header
  plus two `#ifdef GQA` `main()`s (per-head and GQA-reuse), so the shared setup
  lives in one place while each algorithm reads end-to-end. It emits the shader
  `sdpa_compute_out_gqa_coop`.
- Reduction reuses the per-head coop shader's shared-memory tree reduction (no
  subgroup arithmetic), so the variant runs on any Vulkan device -- Adreno and
  Mali alike -- with no capability gate.
- Each thread holds G output accumulators; the array is sized to a compile-time
  `MAX_GROUP_SIZE` = 8 and the group loop is bounded by the `group_size` =
  Hq/Hkv spec constant, so the driver fully unrolls it at pipeline creation.
- Dispatch (`pick_sdpa_av_shader` + global-wg picker + spec-const wiring in
  `add_sdpa_compute_out_node`): the GQA variant is selected on the LLM decode
  coop path when Hq > Hkv, evenly divisible, and G <= 8 (`use_gqa_av_coop`); it
  sets `group_size` and changes the global workgroup z-dim from Hq to Hkv.
  Everything else -- MHA (Hq == Hkv), groups exceeding the cap (G > 8, e.g. MQA
  with Hq > 8), and non-divisible shapes -- falls back to the unchanged per-head
  `sdpa_compute_out_coop`. (Low-ratio MQA -- Hkv == 1 with Hq <= 8 -- is
  eligible and takes the GQA path.)
- A test-only `gqa_override` knob is threaded through `add_sdpa_compute_out_node`
  (declared in the new `SDPA.h`): -1 auto-select, 0 force per-head, 1 force GQA,
  so a benchmark can exercise both AV shaders on the same shape; forcing GQA is
  VK_CHECK'd against shape eligibility.
ghstack-source-id: 405400503
@exported-using-ghexport

Differential Revision: [D112906311](https://our.internmc.facebook.com/intern/diff/D112906311/)
S
Stephen Jia committed
bd7991f4f505e30827d9d8d103525dc75b51fc39
Parent: ff9b4c5