Skip to content

perf(webgpu): subgroup cooperative shaders for MatMul, BatchMatMul, Softmax (#141) - #151

Merged
kolkov merged 3 commits into
mainfrom
perf/subgroup-shaders-sprint3
Aug 4, 2026
Merged

kolkov merged 3 commits into
mainfrom
perf/subgroup-shaders-sprint3

Conversation

@kolkov

@kolkov kolkov commented Aug 4, 2026

Copy link
Copy Markdown
Contributor

Summary

Implements P0+P1 items from #141 — subgroup cooperative reductions for GPU compute shaders.

MatMul + BatchMatMul (P0)

  • matmulSubgroupShader: workgroup_size(32), lanes stride over K dimension, subgroupAdd(partial) reduces across lanes, thread 0 writes result
  • batchMatMulSubgroupShader: same pattern with batch index via wid.z
  • Runtime feature gate: adapter.Features().Contains(FeatureSubgroupOperations) → request feature on device → subgroupsEnabled=true
  • Scalar fallback always available (CI software renderer uses scalar path)

Softmax (P1)

  • softmaxSubgroupShader: three-phase cooperative — subgroupMax for max, subgroupAdd for exp-sum, per-lane normalize
  • One workgroup(32) per row instead of one thread per row

Flash Attention (P1 — deferred)

Current Flash Attention uses per-thread independent accumulators (64 threads × 1 query each). Subgroup optimization would require restructuring the Q@K dot product loop for per-query cooperative reduction — a full algorithm rewrite, not an incremental change. Tracked as follow-up.

Tests

  • 9 new tests: MatMul correctness (8 edge cases), scalar fallback, batch, WGSL syntax
  • 5 new tests: Softmax correctness (10 cases), scalar fallback, explicit subgroup, numerical stability, WGSL syntax

Test plan

  • go build ./...
  • GOOS=js GOARCH=wasm go build ./...
  • golangci-lint run — 0 issues
  • All tests pass (subgroup tests skip on software renderer)

Closes #141

kolkov added 3 commits August 4, 2026 14:52
…#141)

Add matmulSubgroupShader and batchMatMulSubgroupShader using subgroupAdd
for cooperative K-reduction across lanes. Workgroup_size(32) maps to GPU
warp/wavefront. Runtime feature gate: subgroup shaders only when device
supports FeatureSubgroupOperations, scalar fallback otherwise.

Backend requests subgroup feature from adapter when available. Software
renderer (CI) always uses scalar path — zero regression risk.

4 new tests: correctness (8 edge cases), scalar fallback, batch, WGSL syntax.
Three-phase subgroup Softmax: subgroupMax for max, subgroupAdd for
exp-sum, per-lane normalize. One workgroup(32) per row instead of
one thread per row. Runtime feature gate + scalar fallback.

Flash Attention subgroup deferred: requires full algorithm redesign
(per-query cooperative dot product), not an incremental change.

5 tests: correctness (10 cases), scalar fallback, explicit subgroup
path, numerical stability (large values), WGSL syntax validation.
@codecov

codecov Bot commented Aug 4, 2026

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 18.01802% with 91 lines in your changes missing coverage. Please review.

Files with missing lines Patch % Lines
internal/backend/webgpu/compute.go 0.00% 49 Missing ⚠️
internal/backend/webgpu/lazy_compute.go 40.81% 25 Missing and 4 partials ⚠️
internal/backend/webgpu/backend.go 0.00% 13 Missing ⚠️

📢 Thoughts on this report? Let us know!

@kolkov
kolkov merged commit 96ccb8e into main Aug 4, 2026
11 checks passed
@kolkov
kolkov deleted the perf/subgroup-shaders-sprint3 branch August 4, 2026 12:27
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.

perf(webgpu): use subgroupShuffleXor for parallel reduction in compute shaders

1 participant