Xenova HF Staff commited on
Commit
1dd2fce
·
verified ·
1 Parent(s): 64496a6

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -65,14 +65,14 @@ Attributes and default values (overridable per request):
65
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
66
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
67
  - [`test.json`](build/webgpu/test.json) — correctness cases
68
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
69
  - [`matmul-nbits-fused-rms-norm.wgsl.jinja`](build/webgpu/matmul-nbits-fused-rms-norm.wgsl.jinja)
70
  - [`qkv-projection.wgsl.jinja`](build/webgpu/qkv-projection.wgsl.jinja)
71
 
72
  ## Use with `@huggingface/kernels`
73
 
74
  ```sh
75
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
76
  ```
77
 
78
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
65
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
66
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
67
  - [`test.json`](build/webgpu/test.json) — correctness cases
68
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
69
  - [`matmul-nbits-fused-rms-norm.wgsl.jinja`](build/webgpu/matmul-nbits-fused-rms-norm.wgsl.jinja)
70
  - [`qkv-projection.wgsl.jinja`](build/webgpu/qkv-projection.wgsl.jinja)
71
 
72
  ## Use with `@huggingface/kernels`
73
 
74
  ```sh
75
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
76
  ```
77
 
78
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/manifest.json CHANGED
@@ -54,20 +54,35 @@
54
  "aRows": "numel(shapes.aT) / max(1, attrs.K)",
55
  "rowTile": "1 if aRows <= 1 else min(aRows, tunables.ROW_TILE)",
56
  "rowGroups": "ceilDiv(aRows, rowTile)",
57
- "kBlocks": "dim(shapes.qBT, 1)",
58
- "blobSize": "dim(shapes.qBT, 2)",
59
  "codesPerByte": "8 / attrs.bits",
60
  "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
61
- "pairSharesWord": "codesPerByte >= 2",
62
  "epsilonValue": "attrs.epsilon",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
63
  "weightShapeOk": "ranks.qBT == 3 and ranks.kBT == 3 and ranks.vBT == 3 and dim(shapes.qBT, 0) == attrs.Nq and dim(shapes.kBT, 0) == attrs.Nkv and dim(shapes.vBT, 0) == attrs.Nkv and dim(shapes.kBT, 1) == kBlocks and dim(shapes.vBT, 1) == kBlocks and dim(shapes.kBT, 2) == blobSize and dim(shapes.vBT, 2) == blobSize and kBlocks == ceilDiv(attrs.K, attrs.block_size) and blobSize * 8 == attrs.block_size * attrs.bits",
64
  "scaleShapeOk": "ranks.qScalesT == 2 and ranks.kScalesT == 2 and ranks.vScalesT == 2 and dim(shapes.qScalesT, 0) == attrs.Nq and dim(shapes.qScalesT, 1) == kBlocks and dim(shapes.kScalesT, 0) == attrs.Nkv and dim(shapes.kScalesT, 1) == kBlocks and dim(shapes.vScalesT, 0) == attrs.Nkv and dim(shapes.vScalesT, 1) == kBlocks",
65
  "ioShapeOk": "(ranks.aT == 2 or ranks.aT == 3) and dim(shapes.aT, ranks.aT - 1) == attrs.K and ranks.qT == ranks.aT and ranks.kT == ranks.aT and ranks.vT == ranks.aT and dim(shapes.qT, ranks.qT - 1) == attrs.Nq and dim(shapes.kT, ranks.kT - 1) == attrs.Nkv and dim(shapes.vT, ranks.vT - 1) == attrs.Nkv and sameShape(prefix(shapes.qT, ranks.qT - 1), prefix(shapes.aT, ranks.aT - 1)) and sameShape(prefix(shapes.kT, ranks.kT - 1), prefix(shapes.aT, ranks.aT - 1)) and sameShape(prefix(shapes.vT, ranks.vT - 1), prefix(shapes.aT, ranks.aT - 1))",
66
  "dtypeOk": "tensorDtypes.qBT == \"uint8\" and tensorDtypes.kBT == \"uint8\" and tensorDtypes.vBT == \"uint8\" and tensorDtypes.qScalesT == tensorDtypes.aT and tensorDtypes.kScalesT == tensorDtypes.aT and tensorDtypes.vScalesT == tensorDtypes.aT and tensorDtypes.qT == tensorDtypes.aT and tensorDtypes.kT == tensorDtypes.aT and tensorDtypes.vT == tensorDtypes.aT and tensorDtypes.normScaleT == tensorDtypes.aT and f16Ok(tensorDtypes.aT)",
67
- "lanesPow2": "tunables.LANES == pow2ceil(tunables.LANES)",
68
  "normContractOk": "ranks.normScaleT == 1 and dim(shapes.normScaleT, 0) == attrs.K and (sameShape(shapes.skipT, shapes.aT) and tensorDtypes.skipT == tensorDtypes.aT if present.skipT else true) and (sameShape(shapes.residualT, shapes.aT) and tensorDtypes.residualT == tensorDtypes.aT and present.skipT if present.residualT else true)",
69
  "qkvShapeOk": "weightShapeOk and scaleShapeOk and ioShapeOk and dtypeOk and lanesPow2 and normContractOk and pairSharesWord and attrs.K > 0 and attrs.Nq > 0 and attrs.Nkv > 0",
70
- "decodeWalk": "aRows <= 1",
71
  "decodeCols": "4",
72
  "gemvWalk": "rowTile == 1 and blobSize % 16 == 0",
73
  "decodeActVec4": "gemvWalk and attrs.K % attrs.block_size == 0",
@@ -75,59 +90,44 @@
75
  "decodeWorkgroupOk": "tunables.DECODE_WORKGROUP_SIZE >= 4 and pow2ceil(tunables.DECODE_WORKGROUP_SIZE) == tunables.DECODE_WORKGROUP_SIZE and tunables.DECODE_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.DECODE_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX",
76
  "projectionTiles": "ceilDiv(attrs.Nq, tileCols) + 2 * ceilDiv(attrs.Nkv, tileCols)",
77
  "dispatchFits": "decodeWorkgroupOk and projectionTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and aRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and tunables.TILE_N * tunables.LANES <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.TILE_N * tunables.LANES <= device.limits.maxComputeWorkgroupSizeX and tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX",
78
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
79
- "scalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
80
- "K": "attrs.K",
81
  "nq": "attrs.Nq",
82
  "nkv": "attrs.Nkv",
83
- "blockSize": "attrs.block_size",
84
- "bits": "attrs.bits",
85
  "defaultZero": "\"8.0\"",
86
- "tileN": "tunables.TILE_N",
87
- "lanes": "tunables.LANES",
88
- "hidden": "attrs.K",
89
- "workgroupSize": "tunables.NORM_WORKGROUP_SIZE",
90
- "epsilon": "epsilonValue",
91
- "hasSkip": "present.skipT",
92
- "writeResidual": "present.residualT",
93
- "K_LEN": "attrs.K",
94
- "rowCount": "aRows",
95
  "decodeNCols": "decodeCols",
96
  "actVec4": "decodeActVec4",
97
  "weightElement": "\"vec4<u32>\" if gemvWalk else \"u32\"",
98
  "normedElement": "\"vec4<f32>\" if decodeActVec4 else \"f32\"",
99
- "decodeWorkgroupSize": "tunables.DECODE_WORKGROUP_SIZE",
100
- "useSubgroups": "device.features.has(\"subgroups\")"
101
  },
102
  "when": ["dispatchFits", "qkvShapeOk"],
103
  "bindings": {
104
- "a": { "arg": "aT", "buffer": "read-only-storage", "elementType": "$aScalar" },
105
- "norm_scale": { "arg": "normScaleT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$K_LEN" },
106
- "normed": { "scratch": "normedA", "buffer": "storage", "elementType": "f32" },
107
- "params": { "buffer": "uniform", "struct": [{ "name": "rows", "type": "u32", "value": "aRows" }] },
108
- "skip": { "arg": "skipT", "buffer": "read-only-storage", "elementType": "$aScalar" },
109
- "residual": { "arg": "residualT", "buffer": "storage", "elementType": "$aScalar" },
110
- "normed_2": {
111
  "scratch": "normedA",
112
  "name": "normed",
113
  "buffer": "read-only-storage",
114
  "elementType": "$normedElement"
115
  },
116
- "q_b": { "arg": "qBT", "buffer": "read-only-storage", "elementType": "$weightElement" },
117
- "q_scales": { "arg": "qScalesT", "buffer": "read-only-storage", "elementType": "$aScalar" },
118
- "k_b": { "arg": "kBT", "buffer": "read-only-storage", "elementType": "$weightElement" },
119
- "k_scales": { "arg": "kScalesT", "buffer": "read-only-storage", "elementType": "$aScalar" },
120
- "v_b": { "arg": "vBT", "buffer": "read-only-storage", "elementType": "$weightElement" },
121
- "v_scales": { "arg": "vScalesT", "buffer": "read-only-storage", "elementType": "$aScalar" },
122
- "q": { "arg": "qT", "buffer": "storage", "elementType": "$aScalar" },
123
- "k": { "arg": "kT", "buffer": "storage", "elementType": "$aScalar" },
124
- "v": { "arg": "vT", "buffer": "storage", "elementType": "$aScalar" }
125
  },
126
  "variants": [
127
  {
128
  "id": "norm",
129
  "priority": 20,
130
- "when": ["not present.skipT", "not present.residualT"],
131
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
132
  "passes": [
133
  {
@@ -142,16 +142,15 @@
142
  "name": "MatMulNBitsQkv.Projection",
143
  "shader": "qkv-projection.wgsl.jinja",
144
  "derive": { "singleProjection": "\"\"" },
145
- "bindings": ["normed_2", "q_b", "q_scales", "k_b", "k_scales", "v_b", "v_scales", "q", "k", "v"],
146
- "dispatch": { "x": "projectionTiles", "y": "rowGroups" },
147
- "subgroupCollectivesWidth": "portable"
148
  }
149
  ]
150
  },
151
  {
152
  "id": "split_norm",
153
  "priority": 10,
154
- "when": ["not present.skipT", "not present.residualT"],
155
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
156
  "passes": [
157
  {
@@ -166,34 +165,31 @@
166
  "name": "MatMulNBitsQkv.ProjectionQ",
167
  "shader": "qkv-projection.wgsl.jinja",
168
  "derive": { "singleProjection": "\"q\"" },
169
- "bindings": ["normed_2", "q_b", "q_scales", "q"],
170
- "dispatch": { "x": "ceilDiv(attrs.Nq, tileCols)", "y": "rowGroups" },
171
- "subgroupCollectivesWidth": "portable"
172
  },
173
  {
174
  "id": "k",
175
  "name": "MatMulNBitsQkv.ProjectionK",
176
  "shader": "qkv-projection.wgsl.jinja",
177
  "derive": { "singleProjection": "\"k\"" },
178
- "bindings": ["normed_2", "k_b", "k_scales", "k"],
179
- "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" },
180
- "subgroupCollectivesWidth": "portable"
181
  },
182
  {
183
  "id": "v",
184
  "name": "MatMulNBitsQkv.ProjectionV",
185
  "shader": "qkv-projection.wgsl.jinja",
186
  "derive": { "singleProjection": "\"v\"" },
187
- "bindings": ["normed_2", "v_b", "v_scales", "v"],
188
- "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" },
189
- "subgroupCollectivesWidth": "portable"
190
  }
191
  ]
192
  },
193
  {
194
  "id": "skip",
195
  "priority": 20,
196
- "when": ["present.skipT", "not present.residualT"],
197
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
198
  "passes": [
199
  {
@@ -208,16 +204,15 @@
208
  "name": "MatMulNBitsQkv.Projection",
209
  "shader": "qkv-projection.wgsl.jinja",
210
  "derive": { "singleProjection": "\"\"" },
211
- "bindings": ["normed_2", "q_b", "q_scales", "k_b", "k_scales", "v_b", "v_scales", "q", "k", "v"],
212
- "dispatch": { "x": "projectionTiles", "y": "rowGroups" },
213
- "subgroupCollectivesWidth": "portable"
214
  }
215
  ]
216
  },
217
  {
218
  "id": "split_skip",
219
  "priority": 10,
220
- "when": ["present.skipT", "not present.residualT"],
221
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
222
  "passes": [
223
  {
@@ -232,34 +227,31 @@
232
  "name": "MatMulNBitsQkv.ProjectionQ",
233
  "shader": "qkv-projection.wgsl.jinja",
234
  "derive": { "singleProjection": "\"q\"" },
235
- "bindings": ["normed_2", "q_b", "q_scales", "q"],
236
- "dispatch": { "x": "ceilDiv(attrs.Nq, tileCols)", "y": "rowGroups" },
237
- "subgroupCollectivesWidth": "portable"
238
  },
239
  {
240
  "id": "k",
241
  "name": "MatMulNBitsQkv.ProjectionK",
242
  "shader": "qkv-projection.wgsl.jinja",
243
  "derive": { "singleProjection": "\"k\"" },
244
- "bindings": ["normed_2", "k_b", "k_scales", "k"],
245
- "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" },
246
- "subgroupCollectivesWidth": "portable"
247
  },
248
  {
249
  "id": "v",
250
  "name": "MatMulNBitsQkv.ProjectionV",
251
  "shader": "qkv-projection.wgsl.jinja",
252
  "derive": { "singleProjection": "\"v\"" },
253
- "bindings": ["normed_2", "v_b", "v_scales", "v"],
254
- "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" },
255
- "subgroupCollectivesWidth": "portable"
256
  }
257
  ]
258
  },
259
  {
260
  "id": "skipsum",
261
  "priority": 20,
262
- "when": ["present.skipT", "present.residualT"],
263
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
264
  "passes": [
265
  {
@@ -274,16 +266,15 @@
274
  "name": "MatMulNBitsQkv.Projection",
275
  "shader": "qkv-projection.wgsl.jinja",
276
  "derive": { "singleProjection": "\"\"" },
277
- "bindings": ["normed_2", "q_b", "q_scales", "k_b", "k_scales", "v_b", "v_scales", "q", "k", "v"],
278
- "dispatch": { "x": "projectionTiles", "y": "rowGroups" },
279
- "subgroupCollectivesWidth": "portable"
280
  }
281
  ]
282
  },
283
  {
284
  "id": "split_skipsum",
285
  "priority": 10,
286
- "when": ["present.skipT", "present.residualT"],
287
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
288
  "passes": [
289
  {
@@ -298,27 +289,24 @@
298
  "name": "MatMulNBitsQkv.ProjectionQ",
299
  "shader": "qkv-projection.wgsl.jinja",
300
  "derive": { "singleProjection": "\"q\"" },
301
- "bindings": ["normed_2", "q_b", "q_scales", "q"],
302
- "dispatch": { "x": "ceilDiv(attrs.Nq, tileCols)", "y": "rowGroups" },
303
- "subgroupCollectivesWidth": "portable"
304
  },
305
  {
306
  "id": "k",
307
  "name": "MatMulNBitsQkv.ProjectionK",
308
  "shader": "qkv-projection.wgsl.jinja",
309
  "derive": { "singleProjection": "\"k\"" },
310
- "bindings": ["normed_2", "k_b", "k_scales", "k"],
311
- "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" },
312
- "subgroupCollectivesWidth": "portable"
313
  },
314
  {
315
  "id": "v",
316
  "name": "MatMulNBitsQkv.ProjectionV",
317
  "shader": "qkv-projection.wgsl.jinja",
318
  "derive": { "singleProjection": "\"v\"" },
319
- "bindings": ["normed_2", "v_b", "v_scales", "v"],
320
- "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" },
321
- "subgroupCollectivesWidth": "portable"
322
  }
323
  ]
324
  }
 
54
  "aRows": "numel(shapes.aT) / max(1, attrs.K)",
55
  "rowTile": "1 if aRows <= 1 else min(aRows, tunables.ROW_TILE)",
56
  "rowGroups": "ceilDiv(aRows, rowTile)",
 
 
57
  "codesPerByte": "8 / attrs.bits",
58
  "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
 
59
  "epsilonValue": "attrs.epsilon",
60
+ "lanesPow2": "tunables.LANES == pow2ceil(tunables.LANES)",
61
+ "decodeWalk": "aRows <= 1",
62
+ "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
63
+ "scalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
64
+ "K": "attrs.K",
65
+ "blockSize": "attrs.block_size",
66
+ "bits": "attrs.bits",
67
+ "tileN": "tunables.TILE_N",
68
+ "lanes": "tunables.LANES",
69
+ "rowCount": "aRows",
70
+ "useSubgroups": "device.features.has(\"subgroups\")",
71
+ "hidden": "attrs.K",
72
+ "workgroupSize": "tunables.NORM_WORKGROUP_SIZE",
73
+ "epsilon": "epsilonValue",
74
+ "hasSkip": "present.skipT",
75
+ "writeResidual": "present.residualT",
76
+ "K_LEN": "attrs.K",
77
+ "kBlocks": "dim(shapes.qBT, 1)",
78
+ "blobSize": "dim(shapes.qBT, 2)",
79
+ "pairSharesWord": "codesPerByte >= 2",
80
  "weightShapeOk": "ranks.qBT == 3 and ranks.kBT == 3 and ranks.vBT == 3 and dim(shapes.qBT, 0) == attrs.Nq and dim(shapes.kBT, 0) == attrs.Nkv and dim(shapes.vBT, 0) == attrs.Nkv and dim(shapes.kBT, 1) == kBlocks and dim(shapes.vBT, 1) == kBlocks and dim(shapes.kBT, 2) == blobSize and dim(shapes.vBT, 2) == blobSize and kBlocks == ceilDiv(attrs.K, attrs.block_size) and blobSize * 8 == attrs.block_size * attrs.bits",
81
  "scaleShapeOk": "ranks.qScalesT == 2 and ranks.kScalesT == 2 and ranks.vScalesT == 2 and dim(shapes.qScalesT, 0) == attrs.Nq and dim(shapes.qScalesT, 1) == kBlocks and dim(shapes.kScalesT, 0) == attrs.Nkv and dim(shapes.kScalesT, 1) == kBlocks and dim(shapes.vScalesT, 0) == attrs.Nkv and dim(shapes.vScalesT, 1) == kBlocks",
82
  "ioShapeOk": "(ranks.aT == 2 or ranks.aT == 3) and dim(shapes.aT, ranks.aT - 1) == attrs.K and ranks.qT == ranks.aT and ranks.kT == ranks.aT and ranks.vT == ranks.aT and dim(shapes.qT, ranks.qT - 1) == attrs.Nq and dim(shapes.kT, ranks.kT - 1) == attrs.Nkv and dim(shapes.vT, ranks.vT - 1) == attrs.Nkv and sameShape(prefix(shapes.qT, ranks.qT - 1), prefix(shapes.aT, ranks.aT - 1)) and sameShape(prefix(shapes.kT, ranks.kT - 1), prefix(shapes.aT, ranks.aT - 1)) and sameShape(prefix(shapes.vT, ranks.vT - 1), prefix(shapes.aT, ranks.aT - 1))",
83
  "dtypeOk": "tensorDtypes.qBT == \"uint8\" and tensorDtypes.kBT == \"uint8\" and tensorDtypes.vBT == \"uint8\" and tensorDtypes.qScalesT == tensorDtypes.aT and tensorDtypes.kScalesT == tensorDtypes.aT and tensorDtypes.vScalesT == tensorDtypes.aT and tensorDtypes.qT == tensorDtypes.aT and tensorDtypes.kT == tensorDtypes.aT and tensorDtypes.vT == tensorDtypes.aT and tensorDtypes.normScaleT == tensorDtypes.aT and f16Ok(tensorDtypes.aT)",
 
84
  "normContractOk": "ranks.normScaleT == 1 and dim(shapes.normScaleT, 0) == attrs.K and (sameShape(shapes.skipT, shapes.aT) and tensorDtypes.skipT == tensorDtypes.aT if present.skipT else true) and (sameShape(shapes.residualT, shapes.aT) and tensorDtypes.residualT == tensorDtypes.aT and present.skipT if present.residualT else true)",
85
  "qkvShapeOk": "weightShapeOk and scaleShapeOk and ioShapeOk and dtypeOk and lanesPow2 and normContractOk and pairSharesWord and attrs.K > 0 and attrs.Nq > 0 and attrs.Nkv > 0",
 
86
  "decodeCols": "4",
87
  "gemvWalk": "rowTile == 1 and blobSize % 16 == 0",
88
  "decodeActVec4": "gemvWalk and attrs.K % attrs.block_size == 0",
 
90
  "decodeWorkgroupOk": "tunables.DECODE_WORKGROUP_SIZE >= 4 and pow2ceil(tunables.DECODE_WORKGROUP_SIZE) == tunables.DECODE_WORKGROUP_SIZE and tunables.DECODE_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.DECODE_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX",
91
  "projectionTiles": "ceilDiv(attrs.Nq, tileCols) + 2 * ceilDiv(attrs.Nkv, tileCols)",
92
  "dispatchFits": "decodeWorkgroupOk and projectionTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and aRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and tunables.TILE_N * tunables.LANES <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.TILE_N * tunables.LANES <= device.limits.maxComputeWorkgroupSizeX and tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX",
 
 
 
93
  "nq": "attrs.Nq",
94
  "nkv": "attrs.Nkv",
 
 
95
  "defaultZero": "\"8.0\"",
 
 
 
 
 
 
 
 
 
96
  "decodeNCols": "decodeCols",
97
  "actVec4": "decodeActVec4",
98
  "weightElement": "\"vec4<u32>\" if gemvWalk else \"u32\"",
99
  "normedElement": "\"vec4<f32>\" if decodeActVec4 else \"f32\"",
100
+ "decodeWorkgroupSize": "tunables.DECODE_WORKGROUP_SIZE"
 
101
  },
102
  "when": ["dispatchFits", "qkvShapeOk"],
103
  "bindings": {
104
+ "a": { "arg": "aT", "elementType": "$aScalar" },
105
+ "norm_scale": { "arg": "normScaleT", "elementType": "$aScalar", "length": "$K_LEN" },
106
+ "normed": { "scratch": "normedA", "elementType": "f32" },
107
+ "params": { "struct": [{ "name": "rows", "type": "u32", "value": "aRows" }] },
108
+ "skip": { "arg": "skipT", "elementType": "$aScalar" },
109
+ "residual": { "arg": "residualT", "elementType": "$aScalar" },
110
+ "normed_q": {
111
  "scratch": "normedA",
112
  "name": "normed",
113
  "buffer": "read-only-storage",
114
  "elementType": "$normedElement"
115
  },
116
+ "q_b": { "arg": "qBT", "elementType": "$weightElement" },
117
+ "q_scales": { "arg": "qScalesT", "elementType": "$aScalar" },
118
+ "k_b": { "arg": "kBT", "elementType": "$weightElement" },
119
+ "k_scales": { "arg": "kScalesT", "elementType": "$aScalar" },
120
+ "v_b": { "arg": "vBT", "elementType": "$weightElement" },
121
+ "v_scales": { "arg": "vScalesT", "elementType": "$aScalar" },
122
+ "q": { "arg": "qT", "elementType": "$aScalar" },
123
+ "k": { "arg": "kT", "elementType": "$aScalar" },
124
+ "v": { "arg": "vT", "elementType": "$aScalar" }
125
  },
126
  "variants": [
127
  {
128
  "id": "norm",
129
  "priority": 20,
130
+ "when": ["not hasSkip", "not writeResidual"],
131
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
132
  "passes": [
133
  {
 
142
  "name": "MatMulNBitsQkv.Projection",
143
  "shader": "qkv-projection.wgsl.jinja",
144
  "derive": { "singleProjection": "\"\"" },
145
+ "bindings": ["normed_q", "q_b", "q_scales", "k_b", "k_scales", "v_b", "v_scales", "q", "k", "v"],
146
+ "dispatch": { "x": "projectionTiles", "y": "rowGroups" }
 
147
  }
148
  ]
149
  },
150
  {
151
  "id": "split_norm",
152
  "priority": 10,
153
+ "when": ["not hasSkip", "not writeResidual"],
154
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
155
  "passes": [
156
  {
 
165
  "name": "MatMulNBitsQkv.ProjectionQ",
166
  "shader": "qkv-projection.wgsl.jinja",
167
  "derive": { "singleProjection": "\"q\"" },
168
+ "bindings": ["normed_q", "q_b", "q_scales", "q"],
169
+ "dispatch": { "x": "ceilDiv(attrs.Nq, tileCols)", "y": "rowGroups" }
 
170
  },
171
  {
172
  "id": "k",
173
  "name": "MatMulNBitsQkv.ProjectionK",
174
  "shader": "qkv-projection.wgsl.jinja",
175
  "derive": { "singleProjection": "\"k\"" },
176
+ "bindings": ["normed_q", "k_b", "k_scales", "k"],
177
+ "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" }
 
178
  },
179
  {
180
  "id": "v",
181
  "name": "MatMulNBitsQkv.ProjectionV",
182
  "shader": "qkv-projection.wgsl.jinja",
183
  "derive": { "singleProjection": "\"v\"" },
184
+ "bindings": ["normed_q", "v_b", "v_scales", "v"],
185
+ "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" }
 
186
  }
187
  ]
188
  },
189
  {
190
  "id": "skip",
191
  "priority": 20,
192
+ "when": ["hasSkip", "not writeResidual"],
193
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
194
  "passes": [
195
  {
 
204
  "name": "MatMulNBitsQkv.Projection",
205
  "shader": "qkv-projection.wgsl.jinja",
206
  "derive": { "singleProjection": "\"\"" },
207
+ "bindings": ["normed_q", "q_b", "q_scales", "k_b", "k_scales", "v_b", "v_scales", "q", "k", "v"],
208
+ "dispatch": { "x": "projectionTiles", "y": "rowGroups" }
 
209
  }
210
  ]
211
  },
212
  {
213
  "id": "split_skip",
214
  "priority": 10,
215
+ "when": ["hasSkip", "not writeResidual"],
216
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
217
  "passes": [
218
  {
 
227
  "name": "MatMulNBitsQkv.ProjectionQ",
228
  "shader": "qkv-projection.wgsl.jinja",
229
  "derive": { "singleProjection": "\"q\"" },
230
+ "bindings": ["normed_q", "q_b", "q_scales", "q"],
231
+ "dispatch": { "x": "ceilDiv(attrs.Nq, tileCols)", "y": "rowGroups" }
 
232
  },
233
  {
234
  "id": "k",
235
  "name": "MatMulNBitsQkv.ProjectionK",
236
  "shader": "qkv-projection.wgsl.jinja",
237
  "derive": { "singleProjection": "\"k\"" },
238
+ "bindings": ["normed_q", "k_b", "k_scales", "k"],
239
+ "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" }
 
240
  },
241
  {
242
  "id": "v",
243
  "name": "MatMulNBitsQkv.ProjectionV",
244
  "shader": "qkv-projection.wgsl.jinja",
245
  "derive": { "singleProjection": "\"v\"" },
246
+ "bindings": ["normed_q", "v_b", "v_scales", "v"],
247
+ "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" }
 
248
  }
249
  ]
250
  },
251
  {
252
  "id": "skipsum",
253
  "priority": 20,
254
+ "when": ["hasSkip", "writeResidual"],
255
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
256
  "passes": [
257
  {
 
266
  "name": "MatMulNBitsQkv.Projection",
267
  "shader": "qkv-projection.wgsl.jinja",
268
  "derive": { "singleProjection": "\"\"" },
269
+ "bindings": ["normed_q", "q_b", "q_scales", "k_b", "k_scales", "v_b", "v_scales", "q", "k", "v"],
270
+ "dispatch": { "x": "projectionTiles", "y": "rowGroups" }
 
271
  }
272
  ]
273
  },
274
  {
275
  "id": "split_skipsum",
276
  "priority": 10,
277
+ "when": ["hasSkip", "writeResidual"],
278
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
279
  "passes": [
280
  {
 
289
  "name": "MatMulNBitsQkv.ProjectionQ",
290
  "shader": "qkv-projection.wgsl.jinja",
291
  "derive": { "singleProjection": "\"q\"" },
292
+ "bindings": ["normed_q", "q_b", "q_scales", "q"],
293
+ "dispatch": { "x": "ceilDiv(attrs.Nq, tileCols)", "y": "rowGroups" }
 
294
  },
295
  {
296
  "id": "k",
297
  "name": "MatMulNBitsQkv.ProjectionK",
298
  "shader": "qkv-projection.wgsl.jinja",
299
  "derive": { "singleProjection": "\"k\"" },
300
+ "bindings": ["normed_q", "k_b", "k_scales", "k"],
301
+ "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" }
 
302
  },
303
  {
304
  "id": "v",
305
  "name": "MatMulNBitsQkv.ProjectionV",
306
  "shader": "qkv-projection.wgsl.jinja",
307
  "derive": { "singleProjection": "\"v\"" },
308
+ "bindings": ["normed_q", "v_b", "v_scales", "v"],
309
+ "dispatch": { "x": "ceilDiv(attrs.Nkv, tileCols)", "y": "rowGroups" }
 
310
  }
311
  ]
312
  }
build/webgpu/matmul-nbits-fused-rms-norm.wgsl.jinja CHANGED
@@ -12,68 +12,32 @@ const EPSILON: f32 = {{ epsilon }};
12
  var<workgroup> partial: array<f32, WG>;
13
 
14
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
15
- {% if op == "max" %}
16
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
17
- {%- else %}
18
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
19
- {%- endif %}
20
- {% endmacro %}
21
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
22
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
23
  loop {
24
- {% if form == "head" %}
25
- {% if breakInline %}
26
- if ({{ svar }} == 0u) { break; }
27
- {% else %}
28
  if ({{ svar }} == 0u) {
29
  break;
30
  }
31
- {% endif %}
32
- {% endif %}
33
- {% if bodyInline %}
34
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
35
- {% else %}
36
  if ({{ idx }} < {{ svar }}) {
37
  {% for a in arrays %}
38
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
39
  {% endfor %}
40
  }
41
- {% endif %}
42
- {% if form == "head" %}
43
- {% if barrierFirst %}
44
- workgroupBarrier();
45
  {{ svar }} = {{ svar }} / 2u;
46
- {% else %}
47
- {{ svar }} = {{ svar }} / 2u;
48
- workgroupBarrier();
49
- {% endif %}
50
- {% else %}
51
  workgroupBarrier();
52
- if ({{ svar }} == 1u) {
53
- break;
54
- }
55
- {{ svar }} = {{ svar }} / 2u;
56
- {% endif %}
57
- }
58
- {%- endmacro %}
59
-
60
  // Reusing partial after this reduction requires a barrier between the read of
61
  // partial[0] and the next write, or the next round can race the prior readers.
62
- {% set trailingBarrier = trailingBarrier is defined and trailingBarrier %}
63
  fn reduce_sum(value: f32, tid: u32) -> f32 {
64
  partial[tid] = value;
65
  workgroupBarrier();
66
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
67
- {% if trailingBarrier %}
68
- let total = partial[0];
69
- workgroupBarrier();
70
- return total;
71
- {% else %}
72
  return partial[0];
73
- {% endif %}
74
  }
75
 
76
-
77
  fn row_value(index: u32) -> f32 {
78
  {% if hasSkip %}
79
  return f32(a[index]) + f32(skip[index]);
 
12
  var<workgroup> partial: array<f32, WG>;
13
 
14
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
15
+ {% if op == "max" or op == "min" %}
16
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
17
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
18
+ {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false, reuse=false) %}
 
 
 
19
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
20
  loop {
 
 
 
 
21
  if ({{ svar }} == 0u) {
22
  break;
23
  }
 
 
 
 
 
24
  if ({{ idx }} < {{ svar }}) {
25
  {% for a in arrays %}
26
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
27
  {% endfor %}
28
  }
 
 
 
 
29
  {{ svar }} = {{ svar }} / 2u;
 
 
 
 
 
30
  workgroupBarrier();
31
+ }{% endmacro %}
 
 
 
 
 
 
 
32
  // Reusing partial after this reduction requires a barrier between the read of
33
  // partial[0] and the next write, or the next round can race the prior readers.
 
34
  fn reduce_sum(value: f32, tid: u32) -> f32 {
35
  partial[tid] = value;
36
  workgroupBarrier();
37
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
 
 
 
 
 
38
  return partial[0];
 
39
  }
40
 
 
41
  fn row_value(index: u32) -> f32 {
42
  {% if hasSkip %}
43
  return f32(a[index]) + f32(skip[index]);
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.MatMulNBitsQkv",
3
- "id": "_com_microsoft_matmulnbitsqkv_webgpu_f6dfd6b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,15 +8,15 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "58HIjRCefPnXYJoowt4ER40ZwXZBaytxqPwmLff2bUM=",
11
- "manifest.json": "fAGWh8u4euc77NPIK/A78FwhhxsO0TXSThumgnholyE=",
12
- "matmul-nbits-fused-rms-norm.wgsl.jinja": "4lOdB+RprQh3iv29i6RV8UWkn8S5y1aK6Te5nJpxuEk=",
13
- "qkv-projection.wgsl.jinja": "M0SyxodhZFrp/YOpcBTegsqP0M2Vz98ZwSMxSDDa8ew=",
14
  "test.json": "tpaBr6bhveKXElA96j2p78/ivs6jAtlZyTmSxQbbb5w="
15
  }
16
  },
17
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
18
  "webgpu": {
19
- "manifestSpec": "2.0",
20
  "variants": {
21
  "norm": ["matmul-nbits-fused-rms-norm.wgsl.jinja", "qkv-projection.wgsl.jinja"],
22
  "split_norm": ["matmul-nbits-fused-rms-norm.wgsl.jinja", "qkv-projection.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.MatMulNBitsQkv",
3
+ "id": "_com_microsoft_matmulnbitsqkv_webgpu_191b468",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "58HIjRCefPnXYJoowt4ER40ZwXZBaytxqPwmLff2bUM=",
11
+ "manifest.json": "hW1Sfmmk02eyE/qFiM5wkoVf9NjZcTvA4lTyLfpK8l4=",
12
+ "matmul-nbits-fused-rms-norm.wgsl.jinja": "lbzl6cUXSdORYxwAlSR40vtv2Qq+oNUsoJ7DAzgk0zg=",
13
+ "qkv-projection.wgsl.jinja": "yCOTi0oztXRtq5YnUeX7nvhjAnJAWPj+hfsc2FrGM+U=",
14
  "test.json": "tpaBr6bhveKXElA96j2p78/ivs6jAtlZyTmSxQbbb5w="
15
  }
16
  },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
  "webgpu": {
19
+ "manifestSpec": "2.1",
20
  "variants": {
21
  "norm": ["matmul-nbits-fused-rms-norm.wgsl.jinja", "qkv-projection.wgsl.jinja"],
22
  "split_norm": ["matmul-nbits-fused-rms-norm.wgsl.jinja", "qkv-projection.wgsl.jinja"],
build/webgpu/qkv-projection.wgsl.jinja CHANGED
@@ -12,9 +12,7 @@ fn {{ fn }}(n: u32, block: u32, offset: u32) -> u32 {
12
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
13
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
14
  {% endif %}
15
- }
16
- {%- endmacro %}
17
-
18
  {% if gemvWalk and useSubgroups %}
19
  enable subgroups;
20
  {% endif %}
@@ -95,9 +93,9 @@ const VEC_PER_BLOCK: u32 = BLOB_SIZE / 16u;
95
  const VEC_GROUPS: u32 = KBLOCKS * VEC_PER_BLOCK;
96
  const CODES_PER_VEC: u32 = 16u * CODES_PER_BYTE;
97
  {% endif %}
98
-
99
  {% for stream in (["q", "k", "v"] if not singleProjection else [singleProjection]) %}
100
  {% if not gemvWalk %}
 
101
  {{ matmul_nbits_packed_code(fn=stream ~ "_code", buffer=stream ~ "_b", kBlocks="KBLOCKS", blobSize="BLOB_SIZE", bits=bits) }}
102
  // Decode two consecutive reduction-axis codes from one stored byte. An odd
103
  // offset would straddle bytes, so callers advance by two from an even start.
@@ -112,7 +110,6 @@ fn {{ stream }}_code_pair(n: u32, block: u32, offset: u32) -> vec2<u32> {
112
 
113
  {% if gemvWalk %}
114
  var<workgroup> partials: array<vec4<f32>, WG>;
115
-
116
  {% set codesPerWord = 4 * codesPerByte %}
117
  {% macro gemv_project(weights, scalesBuffer) %}
118
  for (var g = tid; g < VEC_GROUPS; g = g + WG) {
@@ -156,8 +153,7 @@ var<workgroup> partials: array<vec4<f32>, WG>;
156
  acc.{{ comp }} = acc.{{ comp }} + (dot - ZERO * asum) * scale;
157
  }
158
  {% endfor %}
159
- }
160
- {%- endmacro %}
161
 
162
  @compute @workgroup_size(WG, 1, 1)
163
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
@@ -226,17 +222,25 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
226
  {% else %}
227
  partials[tid] = acc;
228
  workgroupBarrier();
229
- var stride = WG / 2u;
 
 
 
 
 
230
  loop {
231
- if (stride == 0u) {
232
  break;
233
  }
234
- if (tid < stride) {
235
- partials[tid] = partials[tid] + partials[tid + stride];
 
 
236
  }
237
- stride = stride / 2u;
238
  workgroupBarrier();
239
- }
 
240
 
241
  if (tid == 0u && wg.y < ROWS) {
242
  let total = partials[0];
@@ -266,7 +270,6 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
266
  }
267
  {% else %}
268
  var<workgroup> reduction: array<f32, WG * ROW_TILE>;
269
-
270
  {% macro walk_block(codeFn, guarded) %}
271
  for (var offset = lane * 2u; offset + 1u < BLOCK_SIZE; offset = offset + LANES * 2u) {
272
  let k = k_base + offset;
@@ -288,9 +291,7 @@ var<workgroup> reduction: array<f32, WG * ROW_TILE>;
288
  {% endfor %}
289
  }
290
  {% endif %}
291
- }
292
- {%- endmacro %}
293
-
294
  {% macro project(codeFn, scalesBuffer) %}
295
  for (var block = 0u; block < KBLOCKS; block = block + 1u) {
296
  let scale = f32({{ scalesBuffer }}[n * KBLOCKS + block]);
@@ -309,8 +310,7 @@ var<workgroup> reduction: array<f32, WG * ROW_TILE>;
309
  {% for r in range(rowTile) %}
310
  acc_{{ r }} = acc_{{ r }} + block_acc_{{ r }} * scale;
311
  {% endfor %}
312
- }
313
- {%- endmacro %}
314
 
315
  @compute @workgroup_size(WG, 1, 1)
316
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
 
12
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
13
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
14
  {% endif %}
15
+ }{% endmacro %}
 
 
16
  {% if gemvWalk and useSubgroups %}
17
  enable subgroups;
18
  {% endif %}
 
93
  const VEC_GROUPS: u32 = KBLOCKS * VEC_PER_BLOCK;
94
  const CODES_PER_VEC: u32 = 16u * CODES_PER_BYTE;
95
  {% endif %}
 
96
  {% for stream in (["q", "k", "v"] if not singleProjection else [singleProjection]) %}
97
  {% if not gemvWalk %}
98
+
99
  {{ matmul_nbits_packed_code(fn=stream ~ "_code", buffer=stream ~ "_b", kBlocks="KBLOCKS", blobSize="BLOB_SIZE", bits=bits) }}
100
  // Decode two consecutive reduction-axis codes from one stored byte. An odd
101
  // offset would straddle bytes, so callers advance by two from an even start.
 
110
 
111
  {% if gemvWalk %}
112
  var<workgroup> partials: array<vec4<f32>, WG>;
 
113
  {% set codesPerWord = 4 * codesPerByte %}
114
  {% macro gemv_project(weights, scalesBuffer) %}
115
  for (var g = tid; g < VEC_GROUPS; g = g + WG) {
 
153
  acc.{{ comp }} = acc.{{ comp }} + (dot - ZERO * asum) * scale;
154
  }
155
  {% endfor %}
156
+ }{% endmacro %}
 
157
 
158
  @compute @workgroup_size(WG, 1, 1)
159
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
 
222
  {% else %}
223
  partials[tid] = acc;
224
  workgroupBarrier();
225
+ {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
226
+ {% if op == "max" or op == "min" %}
227
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
228
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
229
+ {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false, reuse=false) %}
230
+ var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
231
  loop {
232
+ if ({{ svar }} == 0u) {
233
  break;
234
  }
235
+ if ({{ idx }} < {{ svar }}) {
236
+ {% for a in arrays %}
237
+ {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
238
+ {% endfor %}
239
  }
240
+ {{ svar }} = {{ svar }} / 2u;
241
  workgroupBarrier();
242
+ }{% endmacro %}
243
+ {{ wgsl_tree_fold(["partials"], idx="tid", wg="WG", form="head") }}
244
 
245
  if (tid == 0u && wg.y < ROWS) {
246
  let total = partials[0];
 
270
  }
271
  {% else %}
272
  var<workgroup> reduction: array<f32, WG * ROW_TILE>;
 
273
  {% macro walk_block(codeFn, guarded) %}
274
  for (var offset = lane * 2u; offset + 1u < BLOCK_SIZE; offset = offset + LANES * 2u) {
275
  let k = k_base + offset;
 
291
  {% endfor %}
292
  }
293
  {% endif %}
294
+ }{% endmacro %}
 
 
295
  {% macro project(codeFn, scalesBuffer) %}
296
  for (var block = 0u; block < KBLOCKS; block = block + 1u) {
297
  let scale = f32({{ scalesBuffer }}[n * KBLOCKS + block]);
 
310
  {% for r in range(rowTile) %}
311
  acc_{{ r }} = acc_{{ r }} + block_acc_{{ r }} * scale;
312
  {% endfor %}
313
+ }{% endmacro %}
 
314
 
315
  @compute @workgroup_size(WG, 1, 1)
316
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {