Xenova HF Staff commited on
Commit
187f23d
·
verified ·
1 Parent(s): 3f9b2cb

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -62,14 +62,14 @@ Attributes and default values (overridable per request):
62
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
63
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
64
  - [`test.json`](build/webgpu/test.json) — correctness cases
65
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
66
  - [`matmul-nbits-fused-rms-norm.wgsl.jinja`](build/webgpu/matmul-nbits-fused-rms-norm.wgsl.jinja)
67
  - [`mlp-gate-up.wgsl.jinja`](build/webgpu/mlp-gate-up.wgsl.jinja)
68
 
69
  ## Use with `@huggingface/kernels`
70
 
71
  ```sh
72
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
73
  ```
74
 
75
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
62
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
63
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
64
  - [`test.json`](build/webgpu/test.json) — correctness cases
65
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
66
  - [`matmul-nbits-fused-rms-norm.wgsl.jinja`](build/webgpu/matmul-nbits-fused-rms-norm.wgsl.jinja)
67
  - [`mlp-gate-up.wgsl.jinja`](build/webgpu/mlp-gate-up.wgsl.jinja)
68
 
69
  ## Use with `@huggingface/kernels`
70
 
71
  ```sh
72
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
73
  ```
74
 
75
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/bench.json CHANGED
@@ -18,6 +18,24 @@
18
  "outputs": { "yT": { "shape": [1, 5632], "dtype": "float32" } },
19
  "bench": { "metrics": [{ "type": "bandwidth", "value": "2 * 5632 * 64 * 16" }] }
20
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
21
  {
22
  "name": "mlp-q4-prefill-m64-k2048-n5632",
23
  "preset": "model",
 
18
  "outputs": { "yT": { "shape": [1, 5632], "dtype": "float32" } },
19
  "bench": { "metrics": [{ "type": "bandwidth", "value": "2 * 5632 * 64 * 16" }] }
20
  },
21
+ {
22
+ "name": "mlp-q4-b16-skip-decode-k2048-n5632",
23
+ "tunableSpace": {},
24
+ "preset": "smoke",
25
+ "vars": { "dtype": "float32" },
26
+ "attrs": { "K": 2048, "N": 5632, "bits": 4, "block_size": 16, "activation": "silu" },
27
+ "inputs": {
28
+ "aT": { "shape": [1, 2048], "dtype": "float32", "dist": "normal", "seed": 8121, "scale": 1 },
29
+ "skipT": { "shape": [1, 2048], "dtype": "float32", "dist": "normal", "seed": 8122, "scale": 1 },
30
+ "normScaleT": { "shape": [2048], "dtype": "float32", "dist": "normal", "seed": 8123, "scale": 1 },
31
+ "gateBT": { "shape": [5632, 128, 8], "dtype": "uint8", "dist": "uniform", "seed": 8124, "min": 0, "max": 256 },
32
+ "gateScalesT": { "shape": [5632, 128], "dtype": "float32", "dist": "normal", "seed": 8125, "scale": 0.05 },
33
+ "upBT": { "shape": [5632, 128, 8], "dtype": "uint8", "dist": "uniform", "seed": 8126, "min": 0, "max": 256 },
34
+ "upScalesT": { "shape": [5632, 128], "dtype": "float32", "dist": "normal", "seed": 8127, "scale": 0.05 }
35
+ },
36
+ "outputs": { "yT": { "shape": [1, 5632], "dtype": "float32" } },
37
+ "bench": { "metrics": [{ "type": "bandwidth", "value": "2 * 5632 * 128 * 8" }] }
38
+ },
39
  {
40
  "name": "mlp-q4-prefill-m64-k2048-n5632",
41
  "preset": "model",
build/webgpu/manifest.json CHANGED
@@ -50,92 +50,93 @@
50
  },
51
  "derive": {
52
  "aRows": "numel(shapes.aT) / max(1, attrs.K)",
53
- "rowTilePlan": "1 if aRows <= 1 else min(aRows, tunables.ROW_TILE)",
54
- "rowGroups": "ceilDiv(aRows, rowTilePlan)",
55
- "kBlocks": "dim(shapes.gateBT, 1)",
56
- "blobSize": "dim(shapes.gateBT, 2)",
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
  "bitsSupported": "attrs.bits == 2 or attrs.bits == 4 or attrs.bits == 8",
61
  "weightShapeOk": "ranks.gateBT == 3 and ranks.upBT == 3 and dim(shapes.gateBT, 0) == attrs.N and dim(shapes.upBT, 0) == attrs.N and dim(shapes.upBT, 1) == kBlocks and dim(shapes.upBT, 2) == blobSize and kBlocks == ceilDiv(attrs.K, attrs.block_size) and blobSize * 8 == attrs.block_size * attrs.bits",
62
  "scaleShapeOk": "ranks.gateScalesT == 2 and ranks.upScalesT == 2 and dim(shapes.gateScalesT, 0) == attrs.N and dim(shapes.gateScalesT, 1) == kBlocks and dim(shapes.upScalesT, 0) == attrs.N and dim(shapes.upScalesT, 1) == kBlocks",
63
  "ioShapeOk": "(ranks.aT == 2 or ranks.aT == 3) and dim(shapes.aT, ranks.aT - 1) == attrs.K and ranks.yT == ranks.aT and dim(shapes.yT, ranks.yT - 1) == attrs.N and sameShape(prefix(shapes.yT, ranks.yT - 1), prefix(shapes.aT, ranks.aT - 1))",
64
  "biasShapeOk": "(ranks.gateBiasT == 1 and dim(shapes.gateBiasT, 0) == attrs.N if present.gateBiasT else true) and (ranks.upBiasT == 1 and dim(shapes.upBiasT, 0) == attrs.N if present.upBiasT else true)",
65
  "dtypeOk": "tensorDtypes.gateScalesT == tensorDtypes.aT and tensorDtypes.upScalesT == tensorDtypes.aT and tensorDtypes.yT == tensorDtypes.aT and f16Ok(tensorDtypes.aT)",
66
- "lanesPow2": "tunables.LANES == pow2ceil(tunables.LANES)",
67
  "mlpShapeOk": "bitsSupported and weightShapeOk and scaleShapeOk and ioShapeOk and biasShapeOk and dtypeOk and lanesPow2 and attrs.K > 0 and attrs.N > 0 and attrs.block_size > 0",
68
  "normContractOk": "present.normScaleT and ranks.normScaleT == 1 and dim(shapes.normScaleT, 0) == attrs.K and tensorDtypes.normScaleT == tensorDtypes.aT 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
- "decodeWalk": "aRows <= 1",
70
  "decodeVec": "decodeWalk and blobSize % 16 == 0",
71
  "decodeActVec4": "decodeVec and attrs.K % attrs.block_size == 0",
72
  "decodeLaneSplit": "decodeVec and kBlocks * blobSize <= tunables.DECODE_WORKGROUP_SIZE * 16",
73
  "decodeCols": "8 if decodeLaneSplit else 4",
74
  "decodeLanes": "tunables.DECODE_WORKGROUP_SIZE * (2 if decodeLaneSplit else 1)",
75
- "tileCols": "decodeCols if decodeWalk else tunables.TILE_N",
76
  "decodeWorkgroupOk": "tunables.DECODE_WORKGROUP_SIZE >= 4 and pow2ceil(tunables.DECODE_WORKGROUP_SIZE) == tunables.DECODE_WORKGROUP_SIZE and decodeLanes <= device.limits.maxComputeInvocationsPerWorkgroup and decodeLanes <= device.limits.maxComputeWorkgroupSizeX",
77
- "gateUpDispatchFits": "decodeWorkgroupOk and ceilDiv(attrs.N, tileCols) <= 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",
78
  "normDispatchFits": "tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX",
79
  "biasPresence_nogb_noub": "not present.gateBiasT and not present.upBiasT",
80
  "biasPresence_nogb_ub": "not present.gateBiasT and present.upBiasT",
81
  "biasPresence_gb_noub": "present.gateBiasT and not present.upBiasT",
82
  "biasPresence_gb_ub": "present.gateBiasT and present.upBiasT",
83
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
84
- "scalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
85
- "K": "attrs.K",
86
  "N": "attrs.N",
87
- "blockSize": "attrs.block_size",
88
- "bits": "attrs.bits",
89
  "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
90
- "tileN": "tunables.TILE_N",
91
- "lanes": "tunables.LANES",
92
- "rowTile": "rowTilePlan",
93
- "rowCount": "aRows",
94
  "decodeNCols": "decodeCols",
95
  "decodeWorkgroupSize": "decodeLanes",
96
  "laneGroups": "2 if decodeLaneSplit else 1",
97
- "useSubgroups": "device.features.has(\"subgroups\")",
98
  "weightElement": "\"vec4<u32>\" if decodeVec else \"u32\"",
99
  "actVec4": "decodeActVec4",
100
  "normedElement": "\"vec4<f32>\" if decodeActVec4 else \"f32\"",
101
- "hidden": "attrs.K",
102
- "workgroupSize": "tunables.NORM_WORKGROUP_SIZE",
103
- "epsilon": "epsilonValue",
104
  "hasGateBias": "present.gateBiasT",
105
  "hasUpBias": "present.upBiasT",
106
- "hasSkip": "present.skipT",
107
- "writeResidual": "present.residualT",
108
- "K_LEN": "attrs.K",
109
  "N_LEN": "attrs.N"
110
  },
111
  "when": ["mlpShapeOk", "gateUpDispatchFits"],
112
  "bindings": {
113
- "a": { "arg": "aT", "buffer": "read-only-storage", "elementType": "$aScalar" },
114
- "gate_b": { "arg": "gateBT", "buffer": "read-only-storage", "elementType": "$weightElement" },
115
- "gate_scales": { "arg": "gateScalesT", "buffer": "read-only-storage", "elementType": "$aScalar" },
116
- "up_b": { "arg": "upBT", "buffer": "read-only-storage", "elementType": "$weightElement" },
117
- "up_scales": { "arg": "upScalesT", "buffer": "read-only-storage", "elementType": "$aScalar" },
118
- "y": { "arg": "yT", "buffer": "storage", "elementType": "$aScalar" },
119
- "up_bias": { "arg": "upBiasT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$N_LEN" },
120
- "gate_bias": { "arg": "gateBiasT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$N_LEN" },
121
- "norm_scale": { "arg": "normScaleT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$K_LEN" },
122
- "normed": { "scratch": "normedA", "buffer": "storage", "elementType": "f32" },
123
- "params": { "buffer": "uniform", "struct": [{ "name": "rows", "type": "u32", "value": "aRows" }] },
124
- "normed_2": {
125
  "scratch": "normedA",
126
  "name": "normed",
127
  "buffer": "read-only-storage",
128
  "elementType": "$normedElement"
129
  },
130
- "skip": { "arg": "skipT", "buffer": "read-only-storage", "elementType": "$aScalar" },
131
- "residual": { "arg": "residualT", "buffer": "storage", "elementType": "$aScalar" }
132
  },
133
  "variants": [
134
  {
135
  "id": "plain_nogb_noub",
136
  "priority": 10,
137
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_nogb_noub"],
138
- "derive": { "inlineNorm": "0", "fromNormed": "0", "rowTile": "rowTilePlan" },
139
  "passes": [
140
  {
141
  "id": "main",
@@ -150,7 +151,7 @@
150
  "id": "plain_nogb_ub",
151
  "priority": 10,
152
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_nogb_ub"],
153
- "derive": { "inlineNorm": "0", "fromNormed": "0", "rowTile": "rowTilePlan" },
154
  "passes": [
155
  {
156
  "id": "main",
@@ -165,7 +166,7 @@
165
  "id": "plain_gb_noub",
166
  "priority": 10,
167
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_gb_noub"],
168
- "derive": { "inlineNorm": "0", "fromNormed": "0", "rowTile": "rowTilePlan" },
169
  "passes": [
170
  {
171
  "id": "main",
@@ -180,23 +181,22 @@
180
  "id": "plain_gb_ub",
181
  "priority": 10,
182
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_gb_ub"],
183
- "derive": { "inlineNorm": "0", "fromNormed": "0", "rowTile": "rowTilePlan" },
184
  "passes": [
185
  {
186
  "id": "main",
187
  "name": "MatMulNBitsMlp.GateUp",
188
  "shader": "mlp-gate-up.wgsl.jinja",
189
  "bindings": ["a", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
190
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
191
- "subgroupCollectivesWidth": "portable"
192
  }
193
  ]
194
  },
195
  {
196
  "id": "staged_norm_nogb_noub",
197
  "priority": 10,
198
- "when": ["normContractOk", "biasPresence_nogb_noub", "normDispatchFits", "not present.skipT", "not present.residualT"],
199
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
200
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
201
  "passes": [
202
  {
@@ -210,17 +210,16 @@
210
  "id": "main",
211
  "name": "MatMulNBitsMlp.GateUp",
212
  "shader": "mlp-gate-up.wgsl.jinja",
213
- "bindings": ["normed_2", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
214
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
215
- "subgroupCollectivesWidth": "portable"
216
  }
217
  ]
218
  },
219
  {
220
  "id": "staged_skip_nogb_noub",
221
  "priority": 10,
222
- "when": ["normContractOk", "biasPresence_nogb_noub", "normDispatchFits", "present.skipT", "not present.residualT"],
223
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
224
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
225
  "passes": [
226
  {
@@ -234,17 +233,16 @@
234
  "id": "main",
235
  "name": "MatMulNBitsMlp.GateUp",
236
  "shader": "mlp-gate-up.wgsl.jinja",
237
- "bindings": ["normed_2", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
238
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
239
- "subgroupCollectivesWidth": "portable"
240
  }
241
  ]
242
  },
243
  {
244
  "id": "staged_skipsum_nogb_noub",
245
  "priority": 10,
246
- "when": ["normContractOk", "biasPresence_nogb_noub", "normDispatchFits", "present.skipT", "present.residualT"],
247
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
248
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
249
  "passes": [
250
  {
@@ -258,17 +256,16 @@
258
  "id": "main",
259
  "name": "MatMulNBitsMlp.GateUp",
260
  "shader": "mlp-gate-up.wgsl.jinja",
261
- "bindings": ["normed_2", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
262
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
263
- "subgroupCollectivesWidth": "portable"
264
  }
265
  ]
266
  },
267
  {
268
  "id": "staged_norm_nogb_ub",
269
  "priority": 10,
270
- "when": ["normContractOk", "biasPresence_nogb_ub", "normDispatchFits", "not present.skipT", "not present.residualT"],
271
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
272
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
273
  "passes": [
274
  {
@@ -282,17 +279,16 @@
282
  "id": "main",
283
  "name": "MatMulNBitsMlp.GateUp",
284
  "shader": "mlp-gate-up.wgsl.jinja",
285
- "bindings": ["normed_2", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
286
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
287
- "subgroupCollectivesWidth": "portable"
288
  }
289
  ]
290
  },
291
  {
292
  "id": "staged_skip_nogb_ub",
293
  "priority": 10,
294
- "when": ["normContractOk", "biasPresence_nogb_ub", "normDispatchFits", "present.skipT", "not present.residualT"],
295
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
296
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
297
  "passes": [
298
  {
@@ -306,17 +302,16 @@
306
  "id": "main",
307
  "name": "MatMulNBitsMlp.GateUp",
308
  "shader": "mlp-gate-up.wgsl.jinja",
309
- "bindings": ["normed_2", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
310
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
311
- "subgroupCollectivesWidth": "portable"
312
  }
313
  ]
314
  },
315
  {
316
  "id": "staged_skipsum_nogb_ub",
317
  "priority": 10,
318
- "when": ["normContractOk", "biasPresence_nogb_ub", "normDispatchFits", "present.skipT", "present.residualT"],
319
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
320
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
321
  "passes": [
322
  {
@@ -330,17 +325,16 @@
330
  "id": "main",
331
  "name": "MatMulNBitsMlp.GateUp",
332
  "shader": "mlp-gate-up.wgsl.jinja",
333
- "bindings": ["normed_2", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
334
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
335
- "subgroupCollectivesWidth": "portable"
336
  }
337
  ]
338
  },
339
  {
340
  "id": "staged_norm_gb_noub",
341
  "priority": 10,
342
- "when": ["normContractOk", "biasPresence_gb_noub", "normDispatchFits", "not present.skipT", "not present.residualT"],
343
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
344
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
345
  "passes": [
346
  {
@@ -354,17 +348,16 @@
354
  "id": "main",
355
  "name": "MatMulNBitsMlp.GateUp",
356
  "shader": "mlp-gate-up.wgsl.jinja",
357
- "bindings": ["normed_2", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
358
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
359
- "subgroupCollectivesWidth": "portable"
360
  }
361
  ]
362
  },
363
  {
364
  "id": "staged_skip_gb_noub",
365
  "priority": 10,
366
- "when": ["normContractOk", "biasPresence_gb_noub", "normDispatchFits", "present.skipT", "not present.residualT"],
367
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
368
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
369
  "passes": [
370
  {
@@ -378,17 +371,16 @@
378
  "id": "main",
379
  "name": "MatMulNBitsMlp.GateUp",
380
  "shader": "mlp-gate-up.wgsl.jinja",
381
- "bindings": ["normed_2", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
382
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
383
- "subgroupCollectivesWidth": "portable"
384
  }
385
  ]
386
  },
387
  {
388
  "id": "staged_skipsum_gb_noub",
389
  "priority": 10,
390
- "when": ["normContractOk", "biasPresence_gb_noub", "normDispatchFits", "present.skipT", "present.residualT"],
391
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
392
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
393
  "passes": [
394
  {
@@ -402,17 +394,16 @@
402
  "id": "main",
403
  "name": "MatMulNBitsMlp.GateUp",
404
  "shader": "mlp-gate-up.wgsl.jinja",
405
- "bindings": ["normed_2", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
406
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
407
- "subgroupCollectivesWidth": "portable"
408
  }
409
  ]
410
  },
411
  {
412
  "id": "staged_norm_gb_ub",
413
  "priority": 10,
414
- "when": ["normContractOk", "biasPresence_gb_ub", "normDispatchFits", "not present.skipT", "not present.residualT"],
415
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
416
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
417
  "passes": [
418
  {
@@ -426,17 +417,16 @@
426
  "id": "main",
427
  "name": "MatMulNBitsMlp.GateUp",
428
  "shader": "mlp-gate-up.wgsl.jinja",
429
- "bindings": ["normed_2", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
430
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
431
- "subgroupCollectivesWidth": "portable"
432
  }
433
  ]
434
  },
435
  {
436
  "id": "staged_skip_gb_ub",
437
  "priority": 10,
438
- "when": ["normContractOk", "biasPresence_gb_ub", "normDispatchFits", "present.skipT", "not present.residualT"],
439
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
440
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
441
  "passes": [
442
  {
@@ -450,17 +440,16 @@
450
  "id": "main",
451
  "name": "MatMulNBitsMlp.GateUp",
452
  "shader": "mlp-gate-up.wgsl.jinja",
453
- "bindings": ["normed_2", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
454
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
455
- "subgroupCollectivesWidth": "portable"
456
  }
457
  ]
458
  },
459
  {
460
  "id": "staged_skipsum_gb_ub",
461
  "priority": 10,
462
- "when": ["normContractOk", "biasPresence_gb_ub", "normDispatchFits", "present.skipT", "present.residualT"],
463
- "derive": { "inlineNorm": "0", "fromNormed": "1", "rowTile": "rowTilePlan" },
464
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
465
  "passes": [
466
  {
@@ -474,17 +463,15 @@
474
  "id": "main",
475
  "name": "MatMulNBitsMlp.GateUp",
476
  "shader": "mlp-gate-up.wgsl.jinja",
477
- "bindings": ["normed_2", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
478
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" },
479
- "subgroupCollectivesWidth": "portable"
480
  }
481
  ]
482
  },
483
  {
484
  "id": "fused_norm_nogb_noub",
485
  "priority": 30,
486
- "when": ["normContractOk", "biasPresence_nogb_noub", "aRows == 1", "not decodeVec", "not present.skipT", "not present.residualT"],
487
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 7 } },
488
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
489
  "passes": [
490
  {
@@ -492,16 +479,14 @@
492
  "name": "MatMulNBitsMlp.FusedDecode",
493
  "shader": "mlp-gate-up.wgsl.jinja",
494
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
495
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
496
- "subgroupCollectivesWidth": "portable"
497
  }
498
  ]
499
  },
500
  {
501
  "id": "fused_skip_nogb_noub",
502
  "priority": 30,
503
- "when": ["normContractOk", "biasPresence_nogb_noub", "aRows == 1", "not decodeVec", "present.skipT", "not present.residualT"],
504
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 8 } },
505
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
506
  "passes": [
507
  {
@@ -509,16 +494,14 @@
509
  "name": "MatMulNBitsMlp.FusedDecode",
510
  "shader": "mlp-gate-up.wgsl.jinja",
511
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
512
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
513
- "subgroupCollectivesWidth": "portable"
514
  }
515
  ]
516
  },
517
  {
518
  "id": "fused_skipsum_nogb_noub",
519
  "priority": 30,
520
- "when": ["normContractOk", "biasPresence_nogb_noub", "aRows == 1", "not decodeVec", "present.skipT", "present.residualT"],
521
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 9 } },
522
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
523
  "passes": [
524
  {
@@ -526,16 +509,14 @@
526
  "name": "MatMulNBitsMlp.FusedDecode",
527
  "shader": "mlp-gate-up.wgsl.jinja",
528
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "y", "residual"],
529
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
530
- "subgroupCollectivesWidth": "portable"
531
  }
532
  ]
533
  },
534
  {
535
  "id": "fused_norm_nogb_ub",
536
  "priority": 30,
537
- "when": ["normContractOk", "biasPresence_nogb_ub", "aRows == 1", "not decodeVec", "not present.skipT", "not present.residualT"],
538
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 8 } },
539
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
540
  "passes": [
541
  {
@@ -543,16 +524,14 @@
543
  "name": "MatMulNBitsMlp.FusedDecode",
544
  "shader": "mlp-gate-up.wgsl.jinja",
545
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
546
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
547
- "subgroupCollectivesWidth": "portable"
548
  }
549
  ]
550
  },
551
  {
552
  "id": "fused_skip_nogb_ub",
553
  "priority": 30,
554
- "when": ["normContractOk", "biasPresence_nogb_ub", "aRows == 1", "not decodeVec", "present.skipT", "not present.residualT"],
555
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 9 } },
556
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
557
  "passes": [
558
  {
@@ -560,16 +539,14 @@
560
  "name": "MatMulNBitsMlp.FusedDecode",
561
  "shader": "mlp-gate-up.wgsl.jinja",
562
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
563
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
564
- "subgroupCollectivesWidth": "portable"
565
  }
566
  ]
567
  },
568
  {
569
  "id": "fused_skipsum_nogb_ub",
570
  "priority": 30,
571
- "when": ["normContractOk", "biasPresence_nogb_ub", "aRows == 1", "not decodeVec", "present.skipT", "present.residualT"],
572
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 10 } },
573
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
574
  "passes": [
575
  {
@@ -577,16 +554,14 @@
577
  "name": "MatMulNBitsMlp.FusedDecode",
578
  "shader": "mlp-gate-up.wgsl.jinja",
579
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y", "residual"],
580
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
581
- "subgroupCollectivesWidth": "portable"
582
  }
583
  ]
584
  },
585
  {
586
  "id": "fused_norm_gb_noub",
587
  "priority": 30,
588
- "when": ["normContractOk", "biasPresence_gb_noub", "aRows == 1", "not decodeVec", "not present.skipT", "not present.residualT"],
589
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 8 } },
590
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
591
  "passes": [
592
  {
@@ -594,16 +569,14 @@
594
  "name": "MatMulNBitsMlp.FusedDecode",
595
  "shader": "mlp-gate-up.wgsl.jinja",
596
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
597
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
598
- "subgroupCollectivesWidth": "portable"
599
  }
600
  ]
601
  },
602
  {
603
  "id": "fused_skip_gb_noub",
604
  "priority": 30,
605
- "when": ["normContractOk", "biasPresence_gb_noub", "aRows == 1", "not decodeVec", "present.skipT", "not present.residualT"],
606
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 9 } },
607
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
608
  "passes": [
609
  {
@@ -611,16 +584,14 @@
611
  "name": "MatMulNBitsMlp.FusedDecode",
612
  "shader": "mlp-gate-up.wgsl.jinja",
613
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
614
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
615
- "subgroupCollectivesWidth": "portable"
616
  }
617
  ]
618
  },
619
  {
620
  "id": "fused_skipsum_gb_noub",
621
  "priority": 30,
622
- "when": ["normContractOk", "biasPresence_gb_noub", "aRows == 1", "not decodeVec", "present.skipT", "present.residualT"],
623
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 10 } },
624
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
625
  "passes": [
626
  {
@@ -628,16 +599,14 @@
628
  "name": "MatMulNBitsMlp.FusedDecode",
629
  "shader": "mlp-gate-up.wgsl.jinja",
630
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y", "residual"],
631
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
632
- "subgroupCollectivesWidth": "portable"
633
  }
634
  ]
635
  },
636
  {
637
  "id": "fused_norm_gb_ub",
638
  "priority": 30,
639
- "when": ["normContractOk", "biasPresence_gb_ub", "aRows == 1", "not decodeVec", "not present.skipT", "not present.residualT"],
640
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 9 } },
641
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
642
  "passes": [
643
  {
@@ -645,16 +614,14 @@
645
  "name": "MatMulNBitsMlp.FusedDecode",
646
  "shader": "mlp-gate-up.wgsl.jinja",
647
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
648
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
649
- "subgroupCollectivesWidth": "portable"
650
  }
651
  ]
652
  },
653
  {
654
  "id": "fused_skip_gb_ub",
655
  "priority": 30,
656
- "when": ["normContractOk", "biasPresence_gb_ub", "aRows == 1", "not decodeVec", "present.skipT", "not present.residualT"],
657
- "requires": { "limits": { "maxStorageBuffersPerShaderStage": 10 } },
658
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
659
  "passes": [
660
  {
@@ -662,8 +629,7 @@
662
  "name": "MatMulNBitsMlp.FusedDecode",
663
  "shader": "mlp-gate-up.wgsl.jinja",
664
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
665
- "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" },
666
- "subgroupCollectivesWidth": "portable"
667
  }
668
  ]
669
  }
 
50
  },
51
  "derive": {
52
  "aRows": "numel(shapes.aT) / max(1, attrs.K)",
 
 
 
 
53
  "codesPerByte": "8 / attrs.bits",
54
  "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
55
  "epsilonValue": "attrs.epsilon",
56
+ "lanesPow2": "tunables.LANES == pow2ceil(tunables.LANES)",
57
+ "decodeWalk": "aRows <= 1",
58
+ "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
59
+ "scalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
60
+ "K": "attrs.K",
61
+ "blockSize": "attrs.block_size",
62
+ "bits": "attrs.bits",
63
+ "lanes": "tunables.LANES",
64
+ "rowCount": "aRows",
65
+ "useSubgroups": "device.features.has(\"subgroups\")",
66
+ "hidden": "attrs.K",
67
+ "workgroupSize": "tunables.NORM_WORKGROUP_SIZE",
68
+ "epsilon": "epsilonValue",
69
+ "hasSkip": "present.skipT",
70
+ "writeResidual": "present.residualT",
71
+ "K_LEN": "attrs.K",
72
+ "variableSubgroup16To32": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 16 and device.adapterInfo.subgroupMaxSize == 32",
73
+ "portableSmallTile": "variableSubgroup16To32 and tensorDtypes.aT == \"float32\" and attrs.bits == 4 and aRows >= 2 * tunables.ROW_TILE and tunables.TILE_N >= 4 and tunables.LANES >= 1 and floor(tunables.TILE_N / 2) * tunables.LANES <= device.limits.maxComputeInvocationsPerWorkgroup and floor(tunables.TILE_N / 2) * tunables.LANES <= device.limits.maxComputeWorkgroupSizeX and ceilDiv(attrs.N, floor(tunables.TILE_N / 2)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and device.limits.maxComputeWorkgroupStorageSize >= 8 * floor(tunables.TILE_N / 2) * tunables.LANES * 2 * tunables.ROW_TILE",
74
+ "tileN": "floor(tunables.TILE_N / 2) if portableSmallTile else tunables.TILE_N",
75
+ "rowTile": "min(aRows, 2 * tunables.ROW_TILE, floor(device.limits.maxComputeWorkgroupStorageSize / (8 * tileN * tunables.LANES))) if portableSmallTile else (1 if aRows <= 1 else min(aRows, tunables.ROW_TILE))",
76
+ "rowGroups": "ceilDiv(aRows, rowTile)",
77
+ "kBlocks": "dim(shapes.gateBT, 1)",
78
+ "blobSize": "dim(shapes.gateBT, 2)",
79
  "bitsSupported": "attrs.bits == 2 or attrs.bits == 4 or attrs.bits == 8",
80
  "weightShapeOk": "ranks.gateBT == 3 and ranks.upBT == 3 and dim(shapes.gateBT, 0) == attrs.N and dim(shapes.upBT, 0) == attrs.N and dim(shapes.upBT, 1) == kBlocks and dim(shapes.upBT, 2) == blobSize and kBlocks == ceilDiv(attrs.K, attrs.block_size) and blobSize * 8 == attrs.block_size * attrs.bits",
81
  "scaleShapeOk": "ranks.gateScalesT == 2 and ranks.upScalesT == 2 and dim(shapes.gateScalesT, 0) == attrs.N and dim(shapes.gateScalesT, 1) == kBlocks and dim(shapes.upScalesT, 0) == attrs.N and dim(shapes.upScalesT, 1) == kBlocks",
82
  "ioShapeOk": "(ranks.aT == 2 or ranks.aT == 3) and dim(shapes.aT, ranks.aT - 1) == attrs.K and ranks.yT == ranks.aT and dim(shapes.yT, ranks.yT - 1) == attrs.N and sameShape(prefix(shapes.yT, ranks.yT - 1), prefix(shapes.aT, ranks.aT - 1))",
83
  "biasShapeOk": "(ranks.gateBiasT == 1 and dim(shapes.gateBiasT, 0) == attrs.N if present.gateBiasT else true) and (ranks.upBiasT == 1 and dim(shapes.upBiasT, 0) == attrs.N if present.upBiasT else true)",
84
  "dtypeOk": "tensorDtypes.gateScalesT == tensorDtypes.aT and tensorDtypes.upScalesT == tensorDtypes.aT and tensorDtypes.yT == tensorDtypes.aT and f16Ok(tensorDtypes.aT)",
 
85
  "mlpShapeOk": "bitsSupported and weightShapeOk and scaleShapeOk and ioShapeOk and biasShapeOk and dtypeOk and lanesPow2 and attrs.K > 0 and attrs.N > 0 and attrs.block_size > 0",
86
  "normContractOk": "present.normScaleT and ranks.normScaleT == 1 and dim(shapes.normScaleT, 0) == attrs.K and tensorDtypes.normScaleT == tensorDtypes.aT 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)",
 
87
  "decodeVec": "decodeWalk and blobSize % 16 == 0",
88
  "decodeActVec4": "decodeVec and attrs.K % attrs.block_size == 0",
89
  "decodeLaneSplit": "decodeVec and kBlocks * blobSize <= tunables.DECODE_WORKGROUP_SIZE * 16",
90
  "decodeCols": "8 if decodeLaneSplit else 4",
91
  "decodeLanes": "tunables.DECODE_WORKGROUP_SIZE * (2 if decodeLaneSplit else 1)",
92
+ "tileCols": "decodeCols if decodeWalk else tileN",
93
  "decodeWorkgroupOk": "tunables.DECODE_WORKGROUP_SIZE >= 4 and pow2ceil(tunables.DECODE_WORKGROUP_SIZE) == tunables.DECODE_WORKGROUP_SIZE and decodeLanes <= device.limits.maxComputeInvocationsPerWorkgroup and decodeLanes <= device.limits.maxComputeWorkgroupSizeX",
94
+ "gateUpDispatchFits": "decodeWorkgroupOk and ceilDiv(attrs.N, tileCols) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and aRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and tileN * tunables.LANES <= device.limits.maxComputeInvocationsPerWorkgroup and tileN * tunables.LANES <= device.limits.maxComputeWorkgroupSizeX",
95
  "normDispatchFits": "tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.NORM_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX",
96
  "biasPresence_nogb_noub": "not present.gateBiasT and not present.upBiasT",
97
  "biasPresence_nogb_ub": "not present.gateBiasT and present.upBiasT",
98
  "biasPresence_gb_noub": "present.gateBiasT and not present.upBiasT",
99
  "biasPresence_gb_ub": "present.gateBiasT and present.upBiasT",
 
 
 
100
  "N": "attrs.N",
 
 
101
  "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
 
 
 
 
102
  "decodeNCols": "decodeCols",
103
  "decodeWorkgroupSize": "decodeLanes",
104
  "laneGroups": "2 if decodeLaneSplit else 1",
 
105
  "weightElement": "\"vec4<u32>\" if decodeVec else \"u32\"",
106
  "actVec4": "decodeActVec4",
107
  "normedElement": "\"vec4<f32>\" if decodeActVec4 else \"f32\"",
 
 
 
108
  "hasGateBias": "present.gateBiasT",
109
  "hasUpBias": "present.upBiasT",
 
 
 
110
  "N_LEN": "attrs.N"
111
  },
112
  "when": ["mlpShapeOk", "gateUpDispatchFits"],
113
  "bindings": {
114
+ "a": { "arg": "aT", "elementType": "$aScalar" },
115
+ "gate_b": { "arg": "gateBT", "elementType": "$weightElement" },
116
+ "gate_scales": { "arg": "gateScalesT", "elementType": "$aScalar" },
117
+ "up_b": { "arg": "upBT", "elementType": "$weightElement" },
118
+ "up_scales": { "arg": "upScalesT", "elementType": "$aScalar" },
119
+ "y": { "arg": "yT", "elementType": "$aScalar" },
120
+ "up_bias": { "arg": "upBiasT", "elementType": "$aScalar", "length": "$N_LEN" },
121
+ "gate_bias": { "arg": "gateBiasT", "elementType": "$aScalar", "length": "$N_LEN" },
122
+ "norm_scale": { "arg": "normScaleT", "elementType": "$aScalar", "length": "$K_LEN" },
123
+ "normed": { "scratch": "normedA", "elementType": "f32" },
124
+ "params": { "struct": [{ "name": "rows", "type": "u32", "value": "aRows" }] },
125
+ "normed_main": {
126
  "scratch": "normedA",
127
  "name": "normed",
128
  "buffer": "read-only-storage",
129
  "elementType": "$normedElement"
130
  },
131
+ "skip": { "arg": "skipT", "elementType": "$aScalar" },
132
+ "residual": { "arg": "residualT", "elementType": "$aScalar" }
133
  },
134
  "variants": [
135
  {
136
  "id": "plain_nogb_noub",
137
  "priority": 10,
138
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_nogb_noub"],
139
+ "derive": { "inlineNorm": "0", "fromNormed": "0" },
140
  "passes": [
141
  {
142
  "id": "main",
 
151
  "id": "plain_nogb_ub",
152
  "priority": 10,
153
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_nogb_ub"],
154
+ "derive": { "inlineNorm": "0", "fromNormed": "0" },
155
  "passes": [
156
  {
157
  "id": "main",
 
166
  "id": "plain_gb_noub",
167
  "priority": 10,
168
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_gb_noub"],
169
+ "derive": { "inlineNorm": "0", "fromNormed": "0" },
170
  "passes": [
171
  {
172
  "id": "main",
 
181
  "id": "plain_gb_ub",
182
  "priority": 10,
183
  "when": ["not present.normScaleT", "not present.skipT", "not present.residualT", "biasPresence_gb_ub"],
184
+ "derive": { "inlineNorm": "0", "fromNormed": "0" },
185
  "passes": [
186
  {
187
  "id": "main",
188
  "name": "MatMulNBitsMlp.GateUp",
189
  "shader": "mlp-gate-up.wgsl.jinja",
190
  "bindings": ["a", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
191
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
192
  }
193
  ]
194
  },
195
  {
196
  "id": "staged_norm_nogb_noub",
197
  "priority": 10,
198
+ "when": ["normContractOk", "biasPresence_nogb_noub", "normDispatchFits", "not hasSkip and not writeResidual"],
199
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
200
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
201
  "passes": [
202
  {
 
210
  "id": "main",
211
  "name": "MatMulNBitsMlp.GateUp",
212
  "shader": "mlp-gate-up.wgsl.jinja",
213
+ "bindings": ["normed_main", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
214
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
215
  }
216
  ]
217
  },
218
  {
219
  "id": "staged_skip_nogb_noub",
220
  "priority": 10,
221
+ "when": ["normContractOk", "biasPresence_nogb_noub", "normDispatchFits", "hasSkip and not writeResidual"],
222
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
223
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
224
  "passes": [
225
  {
 
233
  "id": "main",
234
  "name": "MatMulNBitsMlp.GateUp",
235
  "shader": "mlp-gate-up.wgsl.jinja",
236
+ "bindings": ["normed_main", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
237
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
238
  }
239
  ]
240
  },
241
  {
242
  "id": "staged_skipsum_nogb_noub",
243
  "priority": 10,
244
+ "when": ["normContractOk", "biasPresence_nogb_noub", "normDispatchFits", "hasSkip and writeResidual"],
245
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
246
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
247
  "passes": [
248
  {
 
256
  "id": "main",
257
  "name": "MatMulNBitsMlp.GateUp",
258
  "shader": "mlp-gate-up.wgsl.jinja",
259
+ "bindings": ["normed_main", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
260
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
261
  }
262
  ]
263
  },
264
  {
265
  "id": "staged_norm_nogb_ub",
266
  "priority": 10,
267
+ "when": ["normContractOk", "biasPresence_nogb_ub", "normDispatchFits", "not hasSkip and not writeResidual"],
268
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
269
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
270
  "passes": [
271
  {
 
279
  "id": "main",
280
  "name": "MatMulNBitsMlp.GateUp",
281
  "shader": "mlp-gate-up.wgsl.jinja",
282
+ "bindings": ["normed_main", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
283
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
284
  }
285
  ]
286
  },
287
  {
288
  "id": "staged_skip_nogb_ub",
289
  "priority": 10,
290
+ "when": ["normContractOk", "biasPresence_nogb_ub", "normDispatchFits", "hasSkip and not writeResidual"],
291
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
292
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
293
  "passes": [
294
  {
 
302
  "id": "main",
303
  "name": "MatMulNBitsMlp.GateUp",
304
  "shader": "mlp-gate-up.wgsl.jinja",
305
+ "bindings": ["normed_main", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
306
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
307
  }
308
  ]
309
  },
310
  {
311
  "id": "staged_skipsum_nogb_ub",
312
  "priority": 10,
313
+ "when": ["normContractOk", "biasPresence_nogb_ub", "normDispatchFits", "hasSkip and writeResidual"],
314
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
315
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
316
  "passes": [
317
  {
 
325
  "id": "main",
326
  "name": "MatMulNBitsMlp.GateUp",
327
  "shader": "mlp-gate-up.wgsl.jinja",
328
+ "bindings": ["normed_main", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
329
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
330
  }
331
  ]
332
  },
333
  {
334
  "id": "staged_norm_gb_noub",
335
  "priority": 10,
336
+ "when": ["normContractOk", "biasPresence_gb_noub", "normDispatchFits", "not hasSkip and not writeResidual"],
337
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
338
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
339
  "passes": [
340
  {
 
348
  "id": "main",
349
  "name": "MatMulNBitsMlp.GateUp",
350
  "shader": "mlp-gate-up.wgsl.jinja",
351
+ "bindings": ["normed_main", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
352
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
353
  }
354
  ]
355
  },
356
  {
357
  "id": "staged_skip_gb_noub",
358
  "priority": 10,
359
+ "when": ["normContractOk", "biasPresence_gb_noub", "normDispatchFits", "hasSkip and not writeResidual"],
360
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
361
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
362
  "passes": [
363
  {
 
371
  "id": "main",
372
  "name": "MatMulNBitsMlp.GateUp",
373
  "shader": "mlp-gate-up.wgsl.jinja",
374
+ "bindings": ["normed_main", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
375
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
376
  }
377
  ]
378
  },
379
  {
380
  "id": "staged_skipsum_gb_noub",
381
  "priority": 10,
382
+ "when": ["normContractOk", "biasPresence_gb_noub", "normDispatchFits", "hasSkip and writeResidual"],
383
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
384
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
385
  "passes": [
386
  {
 
394
  "id": "main",
395
  "name": "MatMulNBitsMlp.GateUp",
396
  "shader": "mlp-gate-up.wgsl.jinja",
397
+ "bindings": ["normed_main", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
398
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
399
  }
400
  ]
401
  },
402
  {
403
  "id": "staged_norm_gb_ub",
404
  "priority": 10,
405
+ "when": ["normContractOk", "biasPresence_gb_ub", "normDispatchFits", "not hasSkip and not writeResidual"],
406
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
407
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
408
  "passes": [
409
  {
 
417
  "id": "main",
418
  "name": "MatMulNBitsMlp.GateUp",
419
  "shader": "mlp-gate-up.wgsl.jinja",
420
+ "bindings": ["normed_main", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
421
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
422
  }
423
  ]
424
  },
425
  {
426
  "id": "staged_skip_gb_ub",
427
  "priority": 10,
428
+ "when": ["normContractOk", "biasPresence_gb_ub", "normDispatchFits", "hasSkip and not writeResidual"],
429
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
430
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
431
  "passes": [
432
  {
 
440
  "id": "main",
441
  "name": "MatMulNBitsMlp.GateUp",
442
  "shader": "mlp-gate-up.wgsl.jinja",
443
+ "bindings": ["normed_main", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
444
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
445
  }
446
  ]
447
  },
448
  {
449
  "id": "staged_skipsum_gb_ub",
450
  "priority": 10,
451
+ "when": ["normContractOk", "biasPresence_gb_ub", "normDispatchFits", "hasSkip and writeResidual"],
452
+ "derive": { "inlineNorm": "0", "fromNormed": "1" },
453
  "intermediates": [{ "id": "normedA", "dtype": "float32", "shape": "[numel(shapes.aT)]" }],
454
  "passes": [
455
  {
 
463
  "id": "main",
464
  "name": "MatMulNBitsMlp.GateUp",
465
  "shader": "mlp-gate-up.wgsl.jinja",
466
+ "bindings": ["normed_main", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
467
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "rowGroups" }
 
468
  }
469
  ]
470
  },
471
  {
472
  "id": "fused_norm_nogb_noub",
473
  "priority": 30,
474
+ "when": ["normContractOk", "biasPresence_nogb_noub", "aRows == 1", "not decodeVec", "not hasSkip and not writeResidual"],
 
475
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
476
  "passes": [
477
  {
 
479
  "name": "MatMulNBitsMlp.FusedDecode",
480
  "shader": "mlp-gate-up.wgsl.jinja",
481
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
482
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
483
  }
484
  ]
485
  },
486
  {
487
  "id": "fused_skip_nogb_noub",
488
  "priority": 30,
489
+ "when": ["normContractOk", "biasPresence_nogb_noub", "aRows == 1", "not decodeVec", "hasSkip and not writeResidual"],
 
490
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
491
  "passes": [
492
  {
 
494
  "name": "MatMulNBitsMlp.FusedDecode",
495
  "shader": "mlp-gate-up.wgsl.jinja",
496
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "y"],
497
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
498
  }
499
  ]
500
  },
501
  {
502
  "id": "fused_skipsum_nogb_noub",
503
  "priority": 30,
504
+ "when": ["normContractOk", "biasPresence_nogb_noub", "aRows == 1", "not decodeVec", "hasSkip and writeResidual"],
 
505
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
506
  "passes": [
507
  {
 
509
  "name": "MatMulNBitsMlp.FusedDecode",
510
  "shader": "mlp-gate-up.wgsl.jinja",
511
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "y", "residual"],
512
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
513
  }
514
  ]
515
  },
516
  {
517
  "id": "fused_norm_nogb_ub",
518
  "priority": 30,
519
+ "when": ["normContractOk", "biasPresence_nogb_ub", "aRows == 1", "not decodeVec", "not hasSkip and not writeResidual"],
 
520
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
521
  "passes": [
522
  {
 
524
  "name": "MatMulNBitsMlp.FusedDecode",
525
  "shader": "mlp-gate-up.wgsl.jinja",
526
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
527
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
528
  }
529
  ]
530
  },
531
  {
532
  "id": "fused_skip_nogb_ub",
533
  "priority": 30,
534
+ "when": ["normContractOk", "biasPresence_nogb_ub", "aRows == 1", "not decodeVec", "hasSkip and not writeResidual"],
 
535
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
536
  "passes": [
537
  {
 
539
  "name": "MatMulNBitsMlp.FusedDecode",
540
  "shader": "mlp-gate-up.wgsl.jinja",
541
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y"],
542
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
543
  }
544
  ]
545
  },
546
  {
547
  "id": "fused_skipsum_nogb_ub",
548
  "priority": 30,
549
+ "when": ["normContractOk", "biasPresence_nogb_ub", "aRows == 1", "not decodeVec", "hasSkip and writeResidual"],
 
550
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
551
  "passes": [
552
  {
 
554
  "name": "MatMulNBitsMlp.FusedDecode",
555
  "shader": "mlp-gate-up.wgsl.jinja",
556
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "up_b", "up_scales", "up_bias", "y", "residual"],
557
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
558
  }
559
  ]
560
  },
561
  {
562
  "id": "fused_norm_gb_noub",
563
  "priority": 30,
564
+ "when": ["normContractOk", "biasPresence_gb_noub", "aRows == 1", "not decodeVec", "not hasSkip and not writeResidual"],
 
565
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
566
  "passes": [
567
  {
 
569
  "name": "MatMulNBitsMlp.FusedDecode",
570
  "shader": "mlp-gate-up.wgsl.jinja",
571
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
572
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
573
  }
574
  ]
575
  },
576
  {
577
  "id": "fused_skip_gb_noub",
578
  "priority": 30,
579
+ "when": ["normContractOk", "biasPresence_gb_noub", "aRows == 1", "not decodeVec", "hasSkip and not writeResidual"],
 
580
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
581
  "passes": [
582
  {
 
584
  "name": "MatMulNBitsMlp.FusedDecode",
585
  "shader": "mlp-gate-up.wgsl.jinja",
586
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y"],
587
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
588
  }
589
  ]
590
  },
591
  {
592
  "id": "fused_skipsum_gb_noub",
593
  "priority": 30,
594
+ "when": ["normContractOk", "biasPresence_gb_noub", "aRows == 1", "not decodeVec", "hasSkip and writeResidual"],
 
595
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
596
  "passes": [
597
  {
 
599
  "name": "MatMulNBitsMlp.FusedDecode",
600
  "shader": "mlp-gate-up.wgsl.jinja",
601
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "y", "residual"],
602
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
603
  }
604
  ]
605
  },
606
  {
607
  "id": "fused_norm_gb_ub",
608
  "priority": 30,
609
+ "when": ["normContractOk", "biasPresence_gb_ub", "aRows == 1", "not decodeVec", "not hasSkip and not writeResidual"],
 
610
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
611
  "passes": [
612
  {
 
614
  "name": "MatMulNBitsMlp.FusedDecode",
615
  "shader": "mlp-gate-up.wgsl.jinja",
616
  "bindings": ["a", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
617
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
618
  }
619
  ]
620
  },
621
  {
622
  "id": "fused_skip_gb_ub",
623
  "priority": 30,
624
+ "when": ["normContractOk", "biasPresence_gb_ub", "aRows == 1", "not decodeVec", "hasSkip and not writeResidual"],
 
625
  "derive": { "inlineNorm": "1", "fromNormed": "0", "rowTile": "1" },
626
  "passes": [
627
  {
 
629
  "name": "MatMulNBitsMlp.FusedDecode",
630
  "shader": "mlp-gate-up.wgsl.jinja",
631
  "bindings": ["a", "skip", "norm_scale", "gate_b", "gate_scales", "gate_bias", "up_b", "up_scales", "up_bias", "y"],
632
+ "dispatch": { "x": "ceilDiv(attrs.N, tileCols)", "y": "aRows" }
 
633
  }
634
  ]
635
  }
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,22 +1,22 @@
1
  {
2
  "name": "com.microsoft.MatMulNBitsMlp",
3
- "id": "_com_microsoft_matmulnbitsmlp_webgpu_d711f87",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "teEo4C2oaUqxuIKr6wLZW6vFlmI9gTl/akkWJhjWxqg=",
11
- "manifest.json": "ceiKDb7Itix3E2Ab3NY1WqE4vb5P3kKtrPxTrImqQQk=",
12
- "matmul-nbits-fused-rms-norm.wgsl.jinja": "4lOdB+RprQh3iv29i6RV8UWkn8S5y1aK6Te5nJpxuEk=",
13
- "mlp-gate-up.wgsl.jinja": "VuwljvMtS5vV09dkhewYJLe6s6Syr6zEtOTiGQMMukI=",
14
- "test.json": "x1T4v2UMpwHHuiJcAq6Y8p4TkcyvZwz5pbpbV38m5D4="
15
  }
16
  },
17
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
18
  "webgpu": {
19
- "manifestSpec": "2.0",
20
  "variants": {
21
  "plain_nogb_noub": ["mlp-gate-up.wgsl.jinja"],
22
  "plain_nogb_ub": ["mlp-gate-up.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.MatMulNBitsMlp",
3
+ "id": "_com_microsoft_matmulnbitsmlp_webgpu_9ceab4a",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "ehJ4b7fUHWFgrEfT9QMCQJ+Lmm7CNvhClWNKMbmSgo0=",
11
+ "manifest.json": "wcISD9DIQ/KU8a1i5qDiwKVx1Tnk79JURhMQLb9w62s=",
12
+ "matmul-nbits-fused-rms-norm.wgsl.jinja": "lbzl6cUXSdORYxwAlSR40vtv2Qq+oNUsoJ7DAzgk0zg=",
13
+ "mlp-gate-up.wgsl.jinja": "AEO81iEkyXsDSnjuJdFYKcSmCnQXKFLWbtrOn1pWVlc=",
14
+ "test.json": "SKeEd86oERPUrGq2L1eafvM53Ft4NkoMTAwNuBb3BgQ="
15
  }
16
  },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
  "webgpu": {
19
+ "manifestSpec": "2.1",
20
  "variants": {
21
  "plain_nogb_noub": ["mlp-gate-up.wgsl.jinja"],
22
  "plain_nogb_ub": ["mlp-gate-up.wgsl.jinja"],
build/webgpu/mlp-gate-up.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 rowTile == 1 and useSubgroups %}
19
  enable subgroups;
20
  {% endif %}
@@ -76,12 +74,13 @@ const CODES_PER_VEC: u32 = 16u * CODES_PER_BYTE;
76
  {% if inlineNorm %}
77
  const EPSILON: f32 = {{ epsilon }};
78
  {% endif %}
79
-
80
  {% for stream in ["gate", "up"] %}
81
  {% if not gemvWalk %}
 
82
  {{ matmul_nbits_packed_code(fn=stream ~ "_code", buffer=stream ~ "_b", kBlocks="KBLOCKS", blobSize="BLOB_SIZE", bits=bits) }}
83
  {% endif %}
84
  {% if not decodeVec %}
 
85
  // Decode two consecutive reduction-axis codes. Below 8 bits an even offset and
86
  // its successor share one stored byte; at 8 bits they occupy adjacent bytes of
87
  // one word, because an even byte index in a 4-byte-aligned blob never ends a
@@ -113,71 +112,7 @@ var<workgroup> red_gate: array<f32, WG * ROW_TILE>;
113
  var<workgroup> red_up: array<f32, WG * ROW_TILE>;
114
  {% endif %}
115
  {% if inlineNorm %}
116
- var<workgroup> partial: array<f32, WG>;
117
- var<workgroup> row_inv: f32;
118
-
119
- {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
120
- {% if op == "max" %}
121
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
122
- {%- else %}
123
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
124
- {%- endif %}
125
- {% endmacro %}
126
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
127
- var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
128
- loop {
129
- {% if form == "head" %}
130
- {% if breakInline %}
131
- if ({{ svar }} == 0u) { break; }
132
- {% else %}
133
- if ({{ svar }} == 0u) {
134
- break;
135
- }
136
- {% endif %}
137
- {% endif %}
138
- {% if bodyInline %}
139
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
140
- {% else %}
141
- if ({{ idx }} < {{ svar }}) {
142
- {% for a in arrays %}
143
- {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
144
- {% endfor %}
145
- }
146
- {% endif %}
147
- {% if form == "head" %}
148
- {% if barrierFirst %}
149
- workgroupBarrier();
150
- {{ svar }} = {{ svar }} / 2u;
151
- {% else %}
152
- {{ svar }} = {{ svar }} / 2u;
153
- workgroupBarrier();
154
- {% endif %}
155
- {% else %}
156
- workgroupBarrier();
157
- if ({{ svar }} == 1u) {
158
- break;
159
- }
160
- {{ svar }} = {{ svar }} / 2u;
161
- {% endif %}
162
- }
163
- {%- endmacro %}
164
-
165
- // Reusing partial after this reduction requires a barrier between the read of
166
- // partial[0] and the next write, or the next round can race the prior readers.
167
- {% set trailingBarrier = trailingBarrier is defined and trailingBarrier %}
168
- fn reduce_sum(value: f32, tid: u32) -> f32 {
169
- partial[tid] = value;
170
- workgroupBarrier();
171
- {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
172
- {% if trailingBarrier %}
173
- let total = partial[0];
174
- workgroupBarrier();
175
- return total;
176
- {% else %}
177
- return partial[0];
178
- {% endif %}
179
- }
180
-
181
 
182
  fn row_value(index: u32) -> f32 {
183
  {% if hasSkip %}
@@ -188,8 +123,7 @@ fn row_value(index: u32) -> f32 {
188
  }
189
  {% endif %}
190
 
191
- {% macro act(b, k) %}{% if inlineNorm %}row_value({{ b }} + {{ k }}) * row_inv * f32(norm_scale[{{ k }}]){% elif fromNormed %}normed[{{ b }} + {{ k }}]{% else %}f32(a[{{ b }} + {{ k }}]){% endif %}{%- endmacro %}
192
-
193
  {% if gemvWalk %}
194
  @compute @workgroup_size(WG, 1, 1)
195
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
@@ -200,20 +134,7 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
200
  let row = min(wg.y, ROWS - 1u);
201
  let base_0 = row * K;
202
  let col_base = wg.x * N_COLS;
203
- {% if inlineNorm %}
204
- var local_sq = 0.0;
205
- for (var d = tid; d < K; d = d + WG) {
206
- let value = row_value(base_0 + d);
207
- local_sq = local_sq + value * value;
208
- }
209
- let inv = inverseSqrt(reduce_sum(local_sq, tid) / f32(K) + EPSILON);
210
- if (tid == 0u) {
211
- row_inv = inv;
212
- }
213
- // Separates the reduction's readers of partial[0] from the projection's
214
- // reuse of the same workgroup array below.
215
- workgroupBarrier();
216
- {% if writeResidual %}
217
  // Every N tile computes the same residual row; only the first one stores it,
218
  // so the tiles never write the same location.
219
  if (wg.x == 0u) {
@@ -222,9 +143,8 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
222
  }
223
  }
224
  {% endif %}
225
- {% endif %}
226
- // Whole workgroups past the last column return here, after the normalization
227
- // barriers above and before the reduction's.
228
  if (col_base >= N) {
229
  return;
230
  }
@@ -233,6 +153,9 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
233
  var acc_gate{{ sfx }} = vec4<f32>(0.0);
234
  var acc_up{{ sfx }} = vec4<f32>(0.0);
235
  {% endfor %}
 
 
 
236
  {% if decodeVec %}
237
  {% set codesPerWord = 4 * codesPerByte %}
238
  {% if laneGroups == 2 %}
@@ -325,9 +248,18 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
325
  for (var k = tid * 2u; k < K; k = k + WG * 2u) {
326
  let block = k / BLOCK_SIZE;
327
  let offset = k - block * BLOCK_SIZE;
 
 
 
 
 
 
 
 
328
  let v0 = {{ act("base_0", "k") }};
329
  // K need not be even; a code past the end contributes zero.
330
  let v1 = select(0.0, {{ act("base_0", "min(k + 1u, K - 1u)") }}, k + 1u < K);
 
331
  {% for c in range(4) %}
332
  {% set comp = ["x", "y", "z", "w"][c] %}
333
  {% if c == 0 %}
@@ -360,7 +292,13 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
360
  let sgGate{{ sfx }} = subgroupAdd(acc_gate{{ sfx }});
361
  let sgUp{{ sfx }} = subgroupAdd(acc_up{{ sfx }});
362
  {% endfor %}
 
 
 
363
  if (sgLane == 0u) {
 
 
 
364
  {% for gi in range(colGroups) %}
365
  {% set sfx = "" if gi == 0 else gi %}
366
  red_gate{{ sfx }}[tid / sgSize] = sgGate{{ sfx }};
@@ -376,7 +314,13 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
376
  var gate_total{{ sfx }} = red_gate{{ sfx }}[0];
377
  var up_total{{ sfx }} = red_up{{ sfx }}[0];
378
  {% endfor %}
 
 
 
379
  for (var i = 1u; i < subgroupCount; i = i + 1u) {
 
 
 
380
  {% for gi in range(colGroups) %}
381
  {% set sfx = "" if gi == 0 else gi %}
382
  gate_total{{ sfx }} = gate_total{{ sfx }} + red_gate{{ sfx }}[i];
@@ -389,6 +333,9 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
389
  red_gate{{ sfx }}[tid] = acc_gate{{ sfx }};
390
  red_up{{ sfx }}[tid] = acc_up{{ sfx }};
391
  {% endfor %}
 
 
 
392
  workgroupBarrier();
393
  var stride = WG / 2u;
394
  loop {
@@ -396,6 +343,9 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
396
  break;
397
  }
398
  if (tid < stride) {
 
 
 
399
  {% for gi in range(colGroups) %}
400
  {% set sfx = "" if gi == 0 else gi %}
401
  red_gate{{ sfx }}[tid] = red_gate{{ sfx }}[tid] + red_gate{{ sfx }}[tid + stride];
@@ -412,6 +362,12 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
412
  let gate_total{{ sfx }} = red_gate{{ sfx }}[0];
413
  let up_total{{ sfx }} = red_up{{ sfx }}[0];
414
  {% endfor %}
 
 
 
 
 
 
415
  {% endif %}
416
  {% for gi in range(colGroups) %}
417
  {% set sfx = "" if gi == 0 else gi %}
@@ -424,8 +380,8 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
424
  if (col_base + {{ i }}u < N) {
425
  {% endif %}
426
  let n = col_base + {{ i }}u;
427
- var gate_value = gate_total{{ sfx }}.{{ comp }};
428
- var up_value = up_total{{ sfx }}.{{ comp }};
429
  {% if hasGateBias %}
430
  gate_value = gate_value + f32(gate_bias[n]);
431
  {% endif %}
@@ -472,9 +428,7 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
472
  {% endfor %}
473
  }
474
  {% endif %}
475
- }
476
- {%- endmacro %}
477
-
478
  @compute @workgroup_size(WG, 1, 1)
479
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
480
  let row0 = wg.y * ROW_TILE;
@@ -488,7 +442,6 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
488
  let base_{{ r }} = min(row0 + {{ r }}u, ROWS - 1u) * K;
489
  {% endfor %}
490
 
491
-
492
  {% for r in range(rowTile) %}
493
  var acc_gate_{{ r }} = 0.0;
494
  var acc_up_{{ r }} = 0.0;
 
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 rowTile == 1 and useSubgroups %}
17
  enable subgroups;
18
  {% endif %}
 
74
  {% if inlineNorm %}
75
  const EPSILON: f32 = {{ epsilon }};
76
  {% endif %}
 
77
  {% for stream in ["gate", "up"] %}
78
  {% if not gemvWalk %}
79
+
80
  {{ matmul_nbits_packed_code(fn=stream ~ "_code", buffer=stream ~ "_b", kBlocks="KBLOCKS", blobSize="BLOB_SIZE", bits=bits) }}
81
  {% endif %}
82
  {% if not decodeVec %}
83
+
84
  // Decode two consecutive reduction-axis codes. Below 8 bits an even offset and
85
  // its successor share one stored byte; at 8 bits they occupy adjacent bytes of
86
  // one word, because an even byte index in a 4-byte-aligned blob never ends a
 
112
  var<workgroup> red_up: array<f32, WG * ROW_TILE>;
113
  {% endif %}
114
  {% if inlineNorm %}
115
+ var<workgroup> red_sq: array<f32, WG>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
116
 
117
  fn row_value(index: u32) -> f32 {
118
  {% if hasSkip %}
 
123
  }
124
  {% endif %}
125
 
126
+ {% macro act(b, k) %}{% if fromNormed %}normed[{{ b }} + {{ k }}]{% else %}f32(a[{{ b }} + {{ k }}]){% endif %}{% endmacro %}
 
127
  {% if gemvWalk %}
128
  @compute @workgroup_size(WG, 1, 1)
129
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %},
 
134
  let row = min(wg.y, ROWS - 1u);
135
  let base_0 = row * K;
136
  let col_base = wg.x * N_COLS;
137
+ {% if inlineNorm and writeResidual %}
 
 
 
 
 
 
 
 
 
 
 
 
 
138
  // Every N tile computes the same residual row; only the first one stores it,
139
  // so the tiles never write the same location.
140
  if (wg.x == 0u) {
 
143
  }
144
  }
145
  {% endif %}
146
+ // Whole workgroups past the last column return here, before the reduction's
147
+ // barriers.
 
148
  if (col_base >= N) {
149
  return;
150
  }
 
153
  var acc_gate{{ sfx }} = vec4<f32>(0.0);
154
  var acc_up{{ sfx }} = vec4<f32>(0.0);
155
  {% endfor %}
156
+ {% if inlineNorm %}
157
+ var acc_sq = 0.0;
158
+ {% endif %}
159
  {% if decodeVec %}
160
  {% set codesPerWord = 4 * codesPerByte %}
161
  {% if laneGroups == 2 %}
 
248
  for (var k = tid * 2u; k < K; k = k + WG * 2u) {
249
  let block = k / BLOCK_SIZE;
250
  let offset = k - block * BLOCK_SIZE;
251
+ {% if inlineNorm %}
252
+ let r0 = row_value(base_0 + k);
253
+ // K need not be even; a code past the end contributes zero.
254
+ let r1 = select(0.0, row_value(base_0 + min(k + 1u, K - 1u)), k + 1u < K);
255
+ acc_sq = acc_sq + r0 * r0 + r1 * r1;
256
+ let v0 = r0 * f32(norm_scale[k]);
257
+ let v1 = r1 * f32(norm_scale[min(k + 1u, K - 1u)]);
258
+ {% else %}
259
  let v0 = {{ act("base_0", "k") }};
260
  // K need not be even; a code past the end contributes zero.
261
  let v1 = select(0.0, {{ act("base_0", "min(k + 1u, K - 1u)") }}, k + 1u < K);
262
+ {% endif %}
263
  {% for c in range(4) %}
264
  {% set comp = ["x", "y", "z", "w"][c] %}
265
  {% if c == 0 %}
 
292
  let sgGate{{ sfx }} = subgroupAdd(acc_gate{{ sfx }});
293
  let sgUp{{ sfx }} = subgroupAdd(acc_up{{ sfx }});
294
  {% endfor %}
295
+ {% if inlineNorm %}
296
+ let sgSq = subgroupAdd(acc_sq);
297
+ {% endif %}
298
  if (sgLane == 0u) {
299
+ {% if inlineNorm %}
300
+ red_sq[tid / sgSize] = sgSq;
301
+ {% endif %}
302
  {% for gi in range(colGroups) %}
303
  {% set sfx = "" if gi == 0 else gi %}
304
  red_gate{{ sfx }}[tid / sgSize] = sgGate{{ sfx }};
 
314
  var gate_total{{ sfx }} = red_gate{{ sfx }}[0];
315
  var up_total{{ sfx }} = red_up{{ sfx }}[0];
316
  {% endfor %}
317
+ {% if inlineNorm %}
318
+ var sq_total = red_sq[0];
319
+ {% endif %}
320
  for (var i = 1u; i < subgroupCount; i = i + 1u) {
321
+ {% if inlineNorm %}
322
+ sq_total = sq_total + red_sq[i];
323
+ {% endif %}
324
  {% for gi in range(colGroups) %}
325
  {% set sfx = "" if gi == 0 else gi %}
326
  gate_total{{ sfx }} = gate_total{{ sfx }} + red_gate{{ sfx }}[i];
 
333
  red_gate{{ sfx }}[tid] = acc_gate{{ sfx }};
334
  red_up{{ sfx }}[tid] = acc_up{{ sfx }};
335
  {% endfor %}
336
+ {% if inlineNorm %}
337
+ red_sq[tid] = acc_sq;
338
+ {% endif %}
339
  workgroupBarrier();
340
  var stride = WG / 2u;
341
  loop {
 
343
  break;
344
  }
345
  if (tid < stride) {
346
+ {% if inlineNorm %}
347
+ red_sq[tid] = red_sq[tid] + red_sq[tid + stride];
348
+ {% endif %}
349
  {% for gi in range(colGroups) %}
350
  {% set sfx = "" if gi == 0 else gi %}
351
  red_gate{{ sfx }}[tid] = red_gate{{ sfx }}[tid] + red_gate{{ sfx }}[tid + stride];
 
362
  let gate_total{{ sfx }} = red_gate{{ sfx }}[0];
363
  let up_total{{ sfx }} = red_up{{ sfx }}[0];
364
  {% endfor %}
365
+ {% if inlineNorm %}
366
+ let sq_total = red_sq[0];
367
+ {% endif %}
368
+ {% endif %}
369
+ {% if inlineNorm %}
370
+ let inv = inverseSqrt(sq_total / f32(K) + EPSILON);
371
  {% endif %}
372
  {% for gi in range(colGroups) %}
373
  {% set sfx = "" if gi == 0 else gi %}
 
380
  if (col_base + {{ i }}u < N) {
381
  {% endif %}
382
  let n = col_base + {{ i }}u;
383
+ var gate_value = gate_total{{ sfx }}.{{ comp }}{% if inlineNorm %} * inv{% endif %};
384
+ var up_value = up_total{{ sfx }}.{{ comp }}{% if inlineNorm %} * inv{% endif %};
385
  {% if hasGateBias %}
386
  gate_value = gate_value + f32(gate_bias[n]);
387
  {% endif %}
 
428
  {% endfor %}
429
  }
430
  {% endif %}
431
+ }{% endmacro %}
 
 
432
  @compute @workgroup_size(WG, 1, 1)
433
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
434
  let row0 = wg.y * ROW_TILE;
 
442
  let base_{{ r }} = min(row0 + {{ r }}u, ROWS - 1u) * K;
443
  {% endfor %}
444
 
 
445
  {% for r in range(rowTile) %}
446
  var acc_gate_{{ r }} = 0.0;
447
  var acc_up_{{ r }} = 0.0;
build/webgpu/test.json CHANGED
@@ -1618,7 +1618,7 @@
1618
  {
1619
  "name": "prefill_skip_nogb_noub",
1620
  "provenance": {
1621
- "notes": "Five activation rows force the two-pass prefill schedule. Supplying skip and norm_scale while omitting both projection biases exercises staged SkipSimplifiedLayerNormalization without the residual-sum output."
1622
  },
1623
  "attrs": { "K": 32, "N": 8, "bits": 4, "block_size": 16, "activation": "silu" },
1624
  "inputs": {
@@ -1798,6 +1798,54 @@
1798
  "notes": "Single row with an odd K, so the last reduction pair has one live code; skip input and residual output."
1799
  }
1800
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1801
  {
1802
  "name": "bits8_decode_two_trips",
1803
  "attrs": { "K": 160, "N": 8, "bits": 8, "block_size": 32, "activation": "silu" },
 
1618
  {
1619
  "name": "prefill_skip_nogb_noub",
1620
  "provenance": {
1621
+ "notes": "Five activation rows (M=5) with skip and norm_scale supplied but both projection biases omitted checks staged SkipSimplifiedLayerNormalization without the residual-sum output, at 4-bit block_size=16 quantization."
1622
  },
1623
  "attrs": { "K": 32, "N": 8, "bits": 4, "block_size": 16, "activation": "silu" },
1624
  "inputs": {
 
1798
  "notes": "Single row with an odd K, so the last reduction pair has one live code; skip input and residual output."
1799
  }
1800
  },
1801
+ {
1802
+ "name": "fused_skipsum_multi_trip_odd_k",
1803
+ "attrs": { "K": 301, "N": 5, "bits": 4, "block_size": 16, "activation": "silu", "epsilon": 0.001 },
1804
+ "inputs": {
1805
+ "aT": {
1806
+ "dtype": "float32",
1807
+ "shape": [1, 301],
1808
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.8 }
1809
+ },
1810
+ "skipT": {
1811
+ "dtype": "float32",
1812
+ "shape": [1, 301],
1813
+ "data": { "kind": "fillFloat32", "sinStep": 0.37, "cosStep": 0.11, "scale": 0.5 }
1814
+ },
1815
+ "normScaleT": {
1816
+ "dtype": "float32",
1817
+ "shape": [301],
1818
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.41, "scale": 0.4, "offset": 1.0 }
1819
+ },
1820
+ "gateBT": {
1821
+ "dtype": "uint8",
1822
+ "shape": [5, 19, 8],
1823
+ "data": { "kind": "cycle", "values": [27, 180, 75, 226, 33, 150, 201, 108, 57, 246, 129, 66, 195] }
1824
+ },
1825
+ "gateScalesT": {
1826
+ "dtype": "float32",
1827
+ "shape": [5, 19],
1828
+ "data": { "kind": "cycle", "values": [0.04, 0.055, 0.05, 0.065, 0.06, 0.075, 0.07] }
1829
+ },
1830
+ "upBT": {
1831
+ "dtype": "uint8",
1832
+ "shape": [5, 19, 8],
1833
+ "data": { "kind": "cycle", "values": [211, 44, 137, 98, 165, 20, 233, 121, 78, 190, 15, 252, 87, 143, 61] }
1834
+ },
1835
+ "upScalesT": {
1836
+ "dtype": "float32",
1837
+ "shape": [5, 19],
1838
+ "data": { "kind": "cycle", "values": [0.045, 0.03, 0.07, 0.05, 0.08, 0.035] }
1839
+ }
1840
+ },
1841
+ "outputs": {
1842
+ "yT": { "dtype": "float32", "shape": [1, 5], "tolerance": 0.0001, "relTolerance": 0.0001 },
1843
+ "residualT": { "dtype": "float32", "shape": [1, 301], "tolerance": 0.000001, "relTolerance": 0.000001 }
1844
+ },
1845
+ "provenance": {
1846
+ "notes": "Inline-norm fused route (blob 8, so not the vector walk): the sum of squares accumulates across three 128-code trips of the pair walk, the odd K leaves one live code in the last pair, and N is not a whole column group."
1847
+ }
1848
+ },
1849
  {
1850
  "name": "bits8_decode_two_trips",
1851
  "attrs": { "K": 160, "N": 8, "bits": 8, "block_size": 32, "activation": "silu" },