Xenova HF Staff commited on
Commit
21ba9b4
·
verified ·
1 Parent(s): 9008499

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -37,8 +37,8 @@ Attributes and default values (overridable per request):
37
 
38
  | Attribute | Default | Description |
39
  | --- | --- | --- |
40
- | `activation` | — | Optional fused activation name: `Relu`, `LeakyRelu`, `Sigmoid`, `Tanh`, `HardSigmoid`, `HardSwish`, or `Clip`. Omission applies no activation. |
41
- | `activation_params` | — | Positional parameters for the fused activation: exactly `[alpha]` is required for `LeakyRelu`, and exactly `[alpha, beta]` or `[min, max]` is required for `HardSigmoid` or `Clip`, respectively. Parameter-free activations ignore this attribute. |
42
  | `auto_pad` | `"NOTSET"` | Automatic padding mode. `NOTSET` uses `pads`; `SAME_UPPER` and `SAME_LOWER` choose padding so each output spatial size is `ceil(input / stride)`; `VALID` uses no padding. |
43
  | `dilations` | — | Optional dilation factors, one positive integer per spatial axis. Omission means all ones. |
44
  | `group` | `1` | Number of groups that input and output channels are split into; defaults to 1. |
@@ -78,6 +78,16 @@ One implementation is selected per call from the device capabilities, the reques
78
  - `im2col_half_direct_subgroup_matrix_z` — Materialize aligned f16 convolution columns without widening their storage, then load weights and columns directly into subgroup matrices. Preserve f32 accumulation and the existing bias, residual, activation and f16 output rounding while removing operand staging and K-loop barriers.
79
  - `im2col_half_direct_subgroup_matrix_bias` — Materialize aligned f16 convolution columns without widening their storage, then load weights and columns directly into subgroup matrices. Preserve f32 accumulation and the existing bias, residual, activation and f16 output rounding while removing operand staging and K-loop barriers.
80
  - `im2col_half_direct_subgroup_matrix_bias_z` — Materialize aligned f16 convolution columns without widening their storage, then load weights and columns directly into subgroup matrices. Preserve f32 accumulation and the existing bias, residual, activation and f16 output rounding while removing operand staging and K-loop barriers.
 
 
 
 
 
 
 
 
 
 
81
  - `implicit_im2col_subgroup_matrix` — Gathers convolution input tiles directly into subgroup-matrix operands and applies the existing bias, residual, and activation epilogue, avoiding a materialized column matrix when output-channel reuse permits or both column layouts exceed device allocation limits.
82
  - `implicit_im2col_subgroup_matrix_z` — Gathers convolution input tiles directly into subgroup-matrix operands and applies the existing bias, residual, and activation epilogue, avoiding a materialized column matrix when output-channel reuse permits or both column layouts exceed device allocation limits.
83
  - `implicit_im2col_subgroup_matrix_bias` — Gathers convolution input tiles directly into subgroup-matrix operands and applies the existing bias, residual, and activation epilogue, avoiding a materialized column matrix when output-channel reuse permits or both column layouts exceed device allocation limits.
@@ -96,7 +106,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
96
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
97
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
98
  - [`test.json`](build/webgpu/test.json) — correctness cases
99
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
100
  - [`conv-1x1-gemm-tiled-reg.wgsl.jinja`](build/webgpu/conv-1x1-gemm-tiled-reg.wgsl.jinja)
101
  - [`conv-1x1-gemm-tiled.wgsl.jinja`](build/webgpu/conv-1x1-gemm-tiled.wgsl.jinja)
102
  - [`conv-1x1-subgroup-matrix.wgsl.jinja`](build/webgpu/conv-1x1-subgroup-matrix.wgsl.jinja)
@@ -110,7 +120,7 @@ Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgr
110
  ## Use with `@huggingface/kernels`
111
 
112
  ```sh
113
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
114
  ```
115
 
116
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
37
 
38
  | Attribute | Default | Description |
39
  | --- | --- | --- |
40
+ | `activation` | — | Optional fused activation name: `Relu`, `LeakyRelu`, `Sigmoid`, `Tanh`, `HardSigmoid`, `HardSwish`, `Clip`, `QuickGelu` (x * sigmoid(alpha * x); alpha 1 is SiLU), `Elu`, `Gelu` (erf form, or the tanh approximation when `activation_params[0]` is nonzero), `FastGelu` (always the tanh approximation), `Softplus`, `ThresholdedRelu`, or `Erf`. Omission applies no activation. The last seven are the set onnxruntime's WebGPU provider fuses into Conv; its CPU provider rejects them. |
41
+ | `activation_params` | — | Positional parameters for the fused activation: exactly `[alpha]` is required for `LeakyRelu`, and exactly `[alpha, beta]` or `[min, max]` is required for `HardSigmoid` or `Clip`, respectively. `QuickGelu`, `Elu` and `ThresholdedRelu` take an optional `[alpha]` (defaults 1.702, 1.0 and 1.0); `Gelu` takes an optional `[approximate]` flag (0 = erf, nonzero = tanh). Parameter-free activations ignore this attribute. |
42
  | `auto_pad` | `"NOTSET"` | Automatic padding mode. `NOTSET` uses `pads`; `SAME_UPPER` and `SAME_LOWER` choose padding so each output spatial size is `ceil(input / stride)`; `VALID` uses no padding. |
43
  | `dilations` | — | Optional dilation factors, one positive integer per spatial axis. Omission means all ones. |
44
  | `group` | `1` | Number of groups that input and output channels are split into; defaults to 1. |
 
78
  - `im2col_half_direct_subgroup_matrix_z` — Materialize aligned f16 convolution columns without widening their storage, then load weights and columns directly into subgroup matrices. Preserve f32 accumulation and the existing bias, residual, activation and f16 output rounding while removing operand staging and K-loop barriers.
79
  - `im2col_half_direct_subgroup_matrix_bias` — Materialize aligned f16 convolution columns without widening their storage, then load weights and columns directly into subgroup matrices. Preserve f32 accumulation and the existing bias, residual, activation and f16 output rounding while removing operand staging and K-loop barriers.
80
  - `im2col_half_direct_subgroup_matrix_bias_z` — Materialize aligned f16 convolution columns without widening their storage, then load weights and columns directly into subgroup matrices. Preserve f32 accumulation and the existing bias, residual, activation and f16 output rounding while removing operand staging and K-loop barriers.
81
+ - `im2col_gemm_tiled_reg` — Materialize f32 columns and run a register-tiled GEMM with f32 accumulation.
82
+ - `im2col_gemm_tiled_bias_reg` — Materialize f32 columns and run a register-tiled GEMM with f32 accumulation. Add bias before activation.
83
+ - `im2col_gemm_tiled_reg_f16_columns` — Materialize f16 inputs directly into f16 columns, then run the register-tiled GEMM with f32 accumulation. This preserves the input values exactly and halves column-buffer traffic and storage; eligibility uses the actual half-precision allocation and WebGPU buffer limits.
84
+ - `im2col_gemm_tiled_bias_reg_f16_columns` — Materialize f16 inputs directly into f16 columns, then run the register-tiled GEMM with f32 accumulation. This preserves the input values exactly and halves column-buffer traffic and storage; eligibility uses the actual half-precision allocation and WebGPU buffer limits. Add bias before activation.
85
+ - `grouped_large_kernel_w4` — Shares input windows across channels and columns with f32 accumulation. Filter area, channel bytes and activation presence choose row looping or unrolling. Bias seeds the accumulators; activation follows the reduction. Small workloads and large unrolled kernels outside a fixed 32-wide subgroup range remain eligible as demoted fallbacks.
86
+ - `grouped_large_kernel_w4_bias` — Shares input windows across channels and columns with f32 accumulation. Filter area, channel bytes and activation presence choose row looping or unrolling. Bias seeds the accumulators; activation follows the reduction. Small workloads and large unrolled kernels outside a fixed 32-wide subgroup range remain eligible as demoted fallbacks.
87
+ - `grouped_large_kernel_w4_tail` — Shares input windows across channels and columns with f32 accumulation. Filter area, channel bytes and activation presence choose row looping or unrolling. Bias seeds the accumulators; activation follows the reduction. Small workloads and large unrolled kernels outside a fixed 32-wide subgroup range remain eligible as demoted fallbacks.
88
+ - `grouped_large_kernel_w4_tail_bias` — Shares input windows across channels and columns with f32 accumulation. Filter area, channel bytes and activation presence choose row looping or unrolling. Bias seeds the accumulators; activation follows the reduction. Small workloads and large unrolled kernels outside a fixed 32-wide subgroup range remain eligible as demoted fallbacks.
89
+ - `grouped_large_kernel_w4_dilated_lanes` — Shares input windows across channels and columns with f32 accumulation. Filter area, channel bytes and activation presence choose row looping or unrolling. Bias seeds the accumulators; activation follows the reduction. Small workloads and large unrolled kernels outside a fixed 32-wide subgroup range remain eligible as demoted fallbacks.
90
+ - `grouped_large_kernel_w4_dilated_lanes_bias` — Shares input windows across channels and columns with f32 accumulation. Filter area, channel bytes and activation presence choose row looping or unrolling. Bias seeds the accumulators; activation follows the reduction. Small workloads and large unrolled kernels outside a fixed 32-wide subgroup range remain eligible as demoted fallbacks.
91
  - `implicit_im2col_subgroup_matrix` — Gathers convolution input tiles directly into subgroup-matrix operands and applies the existing bias, residual, and activation epilogue, avoiding a materialized column matrix when output-channel reuse permits or both column layouts exceed device allocation limits.
92
  - `implicit_im2col_subgroup_matrix_z` — Gathers convolution input tiles directly into subgroup-matrix operands and applies the existing bias, residual, and activation epilogue, avoiding a materialized column matrix when output-channel reuse permits or both column layouts exceed device allocation limits.
93
  - `implicit_im2col_subgroup_matrix_bias` — Gathers convolution input tiles directly into subgroup-matrix operands and applies the existing bias, residual, and activation epilogue, avoiding a materialized column matrix when output-channel reuse permits or both column layouts exceed device allocation limits.
 
106
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
107
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
108
  - [`test.json`](build/webgpu/test.json) — correctness cases
109
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
110
  - [`conv-1x1-gemm-tiled-reg.wgsl.jinja`](build/webgpu/conv-1x1-gemm-tiled-reg.wgsl.jinja)
111
  - [`conv-1x1-gemm-tiled.wgsl.jinja`](build/webgpu/conv-1x1-gemm-tiled.wgsl.jinja)
112
  - [`conv-1x1-subgroup-matrix.wgsl.jinja`](build/webgpu/conv-1x1-subgroup-matrix.wgsl.jinja)
 
120
  ## Use with `@huggingface/kernels`
121
 
122
  ```sh
123
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
124
  ```
125
 
126
  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
The diff for this file is too large to render. See raw diff
 
build/webgpu/conv-1x1-gemm-tiled-reg.wgsl.jinja CHANGED
@@ -14,14 +14,11 @@
14
  // instead of 4 * (TM + TN) scalars, so a shared word feeds four times as many
15
  // FMAs and the K loop runs four accumulation steps per iteration.
16
  {{ env.wgsl.resourceDeclarations }}
17
- {% set hasActivation = hasActivation is defined and hasActivation %}
18
- {% set activation = activation | default("") %}
19
- {% set actAlpha = actAlpha | default(0.0) %}
20
- {% set actBeta = actBeta | default(0.0) %}
21
  {% if hasActivation %}
22
  // Apply the fused activation in the f32 accumulator before the single output
23
  // cast, avoiding an intermediate convolution tensor.
24
- {% macro fused_act_return(mode, alpha, beta) -%}
 
25
  {% if mode == "Relu" %}
26
  return max(v, 0.0);
27
  {% elif mode == "Clip" %}
@@ -36,19 +33,49 @@
36
  return tanh(clamp(v, -10.0, 10.0));
37
  {% elif mode == "HardSigmoid" %}
38
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39
  {% else %}
40
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
41
  {% endif %}
42
- {%- endmacro -%}
43
- fn fused_act(v: f32) -> f32 {
44
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
45
- }
46
  {% endif %}
47
-
48
- {% set tileT = "f16" if usesF16 else "f32" -%}
49
- {%- set GEMM_BK = gemmKTile if gemmKTile is defined else 16 %}
50
  {% set GEMM_TM = gemmThreadRows if gemmThreadRows is defined else 8 %}
51
- {% set GEMM_WG_X = gemmWorkgroupX if gemmWorkgroupX is defined else 16 %}
52
  {% set GEMM_BM = gemmMTile if gemmMTile is defined else 64 %}
53
  {% set GEMM_BN = gemmNTile if gemmNTile is defined else GEMM_WG_X * 4 %}
54
  {% set GEMM_TN = gemmThreadColumns if gemmThreadColumns is defined else (GEMM_BN / GEMM_WG_X)|int %}
@@ -70,11 +97,6 @@ const BN_VECS: u32 = BN / 4u;
70
  // Overlapping windows may reread input values, trading address arithmetic and
71
  // repeated input loads for writing and rereading an expanded column matrix.
72
  {% set implicitIm2col = implicitIm2col is defined and implicitIm2col %}
73
- {% set fusedNarrowProjection = fusedNarrowProjection is defined and fusedNarrowProjection %}
74
- {% set projectionChannels = projectionOutChannels | default(0) %}
75
- {% set projectionInputAct = inputActivation | default("none") %}
76
- {% set projectionOutputAct = outputActivation | default("none") %}
77
- {% set narrowProjectionTile = "tileB" if GEMM_BK >= GEMM_BM else "projectionTile" %}
78
  {% set splitKValue = splitK if splitK is defined else 1 %}
79
  {% set splitKPartial = splitKValue > 1 %}
80
  {% set kTiles = kTiles | default(0) %}
@@ -116,37 +138,7 @@ const PARTIAL_COLS: u32 = {{ nPadded }}u;
116
  const PARTIAL_BATCH_STRIDE: u32 = PARTIAL_ROWS * PARTIAL_COLS;
117
  const PARTIAL_SLICE_STRIDE: u32 = {{ batchCount }}u * PARTIAL_BATCH_STRIDE;
118
  {% endif %}
119
-
120
- {% if fusedNarrowProjection %}
121
- {% set inputActivation = inputActivation | default("none") %}
122
- {% set outputActivation = outputActivation | default("none") %}
123
- {% if inputActivation == "relu" or outputActivation == "relu" %}
124
- fn is_nan_f32(value: f32) -> bool {
125
- let bits = bitcast<u32>(value);
126
- return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
127
- }
128
-
129
- fn projection_relu(value: f32) -> f32 {
130
- var out = max(value, 0.0);
131
- if (is_nan_f32(value)) {
132
- out = value;
133
- }
134
- return out;
135
- }
136
- {% endif %}
137
- {% if outputActivation == "sigmoid" %}
138
- fn sigmoid_safe(x: f32) -> f32 {
139
- if (x >= 0.0) {
140
- let z = exp(-x);
141
- return 1.0 / (1.0 + z);
142
- }
143
- let z = exp(x);
144
- return z / (1.0 + z);
145
- }
146
- {% endif %}
147
-
148
- const PROJECTION_OUT_C: u32 = {{ projectionChannels }}u;
149
- {% else %}{% set emitConvStoreOut = not splitKPartial %}{% set hasZ = hasZ is defined and hasZ %}
150
  {% if emitConvStoreOut | default(true) %}fn store_out(m: u32, n: u32, yBase: u32, raw: f32) {
151
  if (m >= params.M || n >= params.N) {
152
  return;
@@ -164,30 +156,9 @@ const PROJECTION_OUT_C: u32 = {{ projectionChannels }}u;
164
  y[yBase + m * params.N + n] = {{ T }}(v);
165
  }
166
  {% endif %}
167
- {% endif %}
168
 
169
  var<workgroup> tileA: array<array<vec4<{{ tileT }}>, AK_VECS>, BM>;
170
  var<workgroup> tileB: array<array<vec4<{{ tileT }}>, BN_VECS>, BK>;
171
- {% if fusedNarrowProjection and GEMM_BK < GEMM_BM %}
172
- // BK16 cannot reuse the 16-row input tile to publish a BM32 intermediate.
173
- // The extra 32x64 f32 tile keeps the total at 14 KiB, below WebGPU's
174
- // guaranteed 16 KiB workgroup-storage floor.
175
- var<workgroup> projectionTile: array<array<vec4<f32>, BN_VECS>, BM>;
176
- {% endif %}
177
- {% macro publish_projection_input(localRow, localVec, component, globalRow, globalColumn, raw) %}
178
- if ({{ globalRow }} < M && {{ globalColumn }} < N) {
179
- var projectionInput = {{ raw }};
180
- {% if hasBias %}
181
- projectionInput = projectionInput + f32(bias[{{ globalRow }}]);
182
- {% endif %}
183
- {% if projectionInputAct == "relu" %}
184
- projectionInput = projection_relu(projectionInput);
185
- {% endif %}
186
- {{ narrowProjectionTile }}[{{ localRow }}][{{ localVec }}].{{ component }} = projectionInput;
187
- } else {
188
- {{ narrowProjectionTile }}[{{ localRow }}][{{ localVec }}].{{ component }} = 0.0;
189
- }
190
- {%- endmacro %}
191
  {% macro load_a_vec4(rowBase, columnBase, rowLimit, columnLimit) %}
192
  for (var linear = li; linear < BM * AK_VECS; linear += WG_SIZE) {
193
  let ar = linear / AK_VECS;
@@ -208,8 +179,7 @@ var<workgroup> projectionTile: array<array<vec4<f32>, BN_VECS>, BM>;
208
  }
209
  }
210
  tileA[ar][ac4] = av;
211
- }
212
- {%- endmacro %}
213
  {% macro load_b_vec4(rowBase, columnBase, rowLimit, columnLimit) %}
214
  for (var linear = li; linear < BK * BN_VECS; linear += WG_SIZE) {
215
  let br = linear / BN_VECS;
@@ -230,8 +200,7 @@ var<workgroup> projectionTile: array<array<vec4<f32>, BN_VECS>, BM>;
230
  }
231
  }
232
  tileB[br][bc4] = bvec;
233
- }
234
- {%- endmacro %}
235
  {% macro load_b_implicit_vec4(rowBase, columnBase, rowLimit, columnLimit) %}
236
  for (var linear = li; linear < BK * BN_VECS; linear += WG_SIZE) {
237
  let br = linear / BN_VECS;
@@ -259,8 +228,7 @@ var<workgroup> projectionTile: array<array<vec4<f32>, BN_VECS>, BM>;
259
  }
260
  }
261
  tileB[br][bc4] = bvec;
262
- }
263
- {%- endmacro %}
264
  {% macro load_b_implicit_carried(rowBase, columnBase, rowLimit, columnLimit) %}
265
  let bColVec = li / {{ implicitGatherKChunks }}u;
266
  let bChunk = li % {{ implicitGatherKChunks }}u;
@@ -344,8 +312,7 @@ var<workgroup> projectionTile: array<array<vec4<f32>, BN_VECS>, BM>;
344
  }
345
  }
346
  }
347
- }
348
- {%- endmacro %}
349
 
350
  @compute @workgroup_size({{ GEMM_WG_X }}, {{ GEMM_WG_Y }}, 1)
351
  fn main(
@@ -423,9 +390,7 @@ fn main(
423
  let n0 = nBase + lid.x * TN;
424
  let partialBase = slice * PARTIAL_SLICE_STRIDE + batch * PARTIAL_BATCH_STRIDE;
425
  {% else %}
426
- {% if not fusedNarrowProjection %}
427
  let yBatchBase = batch * M * N;
428
- {% endif %}
429
  let m0 = mBase + lid.y * TM;
430
  let n0 = nBase + lid.x * TN;
431
  {% endif %}
@@ -433,52 +398,8 @@ fn main(
433
  {% for column in range(GEMM_TN) %}
434
  {% if splitKPartial %}
435
  y[partialBase + (m0 + {{ row }}u) * PARTIAL_COLS + n0 + {{ column }}u] = acc{{ row }}.{{ components[column] }};
436
- {% elif fusedNarrowProjection %}
437
- {{ publish_projection_input("lid.y * TM + " ~ row ~ "u", "lid.x * " ~ ((GEMM_TN / 4)|int) ~ "u + " ~ ((column / 4)|int) ~ "u", components[column % 4], "m0 + " ~ row ~ "u", "n0 + " ~ column ~ "u", "acc" ~ row ~ "." ~ components[column]) }}
438
  {% else %}
439
  store_out(m0 + {{ row }}u, n0 + {{ column }}u, yBatchBase, acc{{ row }}.{{ components[column] }});
440
  {% endif %}
441
  {% endfor %}
442
- {% endfor %}{% if fusedNarrowProjection %}
443
-
444
- // Conv's final K-tile barrier makes tileB dead before it becomes the
445
- // intermediate tile. Every lane publishes its unique micro-tile, then one
446
- // lane per spatial vector word consumes all logical M rows in increasing
447
- // order, four adjacent columns at a time. Invalid M/N-tail cells are
448
- // initialized above, so the barrier is uniform and no lane can observe a
449
- // stale input-tile value.
450
- workgroupBarrier();
451
- if (li < BN_VECS) {
452
- let projectionNBase = nBase + li * 4u;
453
- {% for oc in range(projectionChannels) %}
454
- var projectionAcc{{ oc }} = vec4<f32>(0.0);
455
- {% endfor %}
456
- for (var channel = 0u; channel < M; channel = channel + 1u) {
457
- let projectionValues = vec4<f32>({{ narrowProjectionTile }}[channel][li]);
458
- {% for oc in range(projectionChannels) %}
459
- projectionAcc{{ oc }} = projectionAcc{{ oc }} + projectionValues * f32(projectionW[{{ oc }}u * M + channel]);
460
- {% endfor %}
461
- }
462
-
463
- let projectionYBase = batch * PROJECTION_OUT_C * N;
464
- {% for oc in range(projectionChannels) %}
465
- let projected{{ oc }} = projectionAcc{{ oc }};
466
- {% endfor %}
467
- {% for column in range(4) %}
468
- let projectionN{{ column }} = projectionNBase + {{ column }}u;
469
- if (projectionN{{ column }} < N) {
470
- {% for oc in range(projectionChannels) %}
471
- {% if projectionOutputAct == "relu" %}
472
- let activated{{ oc }}_{{ column }} = projection_relu(projected{{ oc }}.{{ components[column] }});
473
- {% elif projectionOutputAct == "sigmoid" %}
474
- let activated{{ oc }}_{{ column }} = sigmoid_safe(projected{{ oc }}.{{ components[column] }});
475
- {% else %}
476
- let activated{{ oc }}_{{ column }} = projected{{ oc }}.{{ components[column] }};
477
- {% endif %}
478
- y[projectionYBase + {{ oc }}u * N + projectionN{{ column }}] = activated{{ oc }}_{{ column }};
479
- {% endfor %}
480
- }
481
- {% endfor %}
482
- }
483
- {% endif %}
484
- }
 
14
  // instead of 4 * (TM + TN) scalars, so a shared word feeds four times as many
15
  // FMAs and the K loop runs four accumulation steps per iteration.
16
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
17
  {% if hasActivation %}
18
  // Apply the fused activation in the f32 accumulator before the single output
19
  // cast, avoiding an intermediate convolution tensor.
20
+ {% macro fused_act_fn(mode, alpha, beta) %}
21
+ fn fused_act(v: f32) -> f32 {
22
  {% if mode == "Relu" %}
23
  return max(v, 0.0);
24
  {% elif mode == "Clip" %}
 
33
  return tanh(clamp(v, -10.0, 10.0));
34
  {% elif mode == "HardSigmoid" %}
35
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
36
+ {% elif mode == "QuickGelu" %}
37
+ {% if alpha == 1 %}
38
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
39
+ return v * (1.0 / (1.0 + exp(-v)));
40
+ {% else %}
41
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
42
+ {% endif %}
43
+ {% elif mode == "Elu" %}
44
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
45
+ {% elif mode == "ThresholdedRelu" %}
46
+ return select(0.0, v, v > f32({{ alpha }}));
47
+ {% elif mode == "Softplus" %}
48
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
49
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
50
+ {% elif mode == "Erf" %}
51
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
52
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
53
+ let a = abs(v);
54
+ let t = 1.0 / (1.0 + 0.3275911 * a);
55
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
56
+ return sign(v) * (1.0 - poly * exp(-a * a));
57
+ {% elif mode == "Gelu" and alpha == 0 %}
58
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
59
+ let u = v * 0.70710678118654752;
60
+ let a = abs(u);
61
+ let t = 1.0 / (1.0 + 0.3275911 * a);
62
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
63
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
64
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
65
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
66
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
67
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
68
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
69
  {% else %}
70
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
71
  {% endif %}
72
+ }{% endmacro %}
73
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
74
  {% endif %}
75
+ {% set tileT = "f16" if usesF16 else "f32" %}
76
+ {% set GEMM_BK = gemmKTile if gemmKTile is defined else 16 %}
 
77
  {% set GEMM_TM = gemmThreadRows if gemmThreadRows is defined else 8 %}
78
+ {% set GEMM_WG_X = 16 %}
79
  {% set GEMM_BM = gemmMTile if gemmMTile is defined else 64 %}
80
  {% set GEMM_BN = gemmNTile if gemmNTile is defined else GEMM_WG_X * 4 %}
81
  {% set GEMM_TN = gemmThreadColumns if gemmThreadColumns is defined else (GEMM_BN / GEMM_WG_X)|int %}
 
97
  // Overlapping windows may reread input values, trading address arithmetic and
98
  // repeated input loads for writing and rereading an expanded column matrix.
99
  {% set implicitIm2col = implicitIm2col is defined and implicitIm2col %}
 
 
 
 
 
100
  {% set splitKValue = splitK if splitK is defined else 1 %}
101
  {% set splitKPartial = splitKValue > 1 %}
102
  {% set kTiles = kTiles | default(0) %}
 
138
  const PARTIAL_BATCH_STRIDE: u32 = PARTIAL_ROWS * PARTIAL_COLS;
139
  const PARTIAL_SLICE_STRIDE: u32 = {{ batchCount }}u * PARTIAL_BATCH_STRIDE;
140
  {% endif %}
141
+ {% set emitConvStoreOut = not splitKPartial %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
142
  {% if emitConvStoreOut | default(true) %}fn store_out(m: u32, n: u32, yBase: u32, raw: f32) {
143
  if (m >= params.M || n >= params.N) {
144
  return;
 
156
  y[yBase + m * params.N + n] = {{ T }}(v);
157
  }
158
  {% endif %}
 
159
 
160
  var<workgroup> tileA: array<array<vec4<{{ tileT }}>, AK_VECS>, BM>;
161
  var<workgroup> tileB: array<array<vec4<{{ tileT }}>, BN_VECS>, BK>;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
162
  {% macro load_a_vec4(rowBase, columnBase, rowLimit, columnLimit) %}
163
  for (var linear = li; linear < BM * AK_VECS; linear += WG_SIZE) {
164
  let ar = linear / AK_VECS;
 
179
  }
180
  }
181
  tileA[ar][ac4] = av;
182
+ }{% endmacro %}
 
183
  {% macro load_b_vec4(rowBase, columnBase, rowLimit, columnLimit) %}
184
  for (var linear = li; linear < BK * BN_VECS; linear += WG_SIZE) {
185
  let br = linear / BN_VECS;
 
200
  }
201
  }
202
  tileB[br][bc4] = bvec;
203
+ }{% endmacro %}
 
204
  {% macro load_b_implicit_vec4(rowBase, columnBase, rowLimit, columnLimit) %}
205
  for (var linear = li; linear < BK * BN_VECS; linear += WG_SIZE) {
206
  let br = linear / BN_VECS;
 
228
  }
229
  }
230
  tileB[br][bc4] = bvec;
231
+ }{% endmacro %}
 
232
  {% macro load_b_implicit_carried(rowBase, columnBase, rowLimit, columnLimit) %}
233
  let bColVec = li / {{ implicitGatherKChunks }}u;
234
  let bChunk = li % {{ implicitGatherKChunks }}u;
 
312
  }
313
  }
314
  }
315
+ }{% endmacro %}
 
316
 
317
  @compute @workgroup_size({{ GEMM_WG_X }}, {{ GEMM_WG_Y }}, 1)
318
  fn main(
 
390
  let n0 = nBase + lid.x * TN;
391
  let partialBase = slice * PARTIAL_SLICE_STRIDE + batch * PARTIAL_BATCH_STRIDE;
392
  {% else %}
 
393
  let yBatchBase = batch * M * N;
 
394
  let m0 = mBase + lid.y * TM;
395
  let n0 = nBase + lid.x * TN;
396
  {% endif %}
 
398
  {% for column in range(GEMM_TN) %}
399
  {% if splitKPartial %}
400
  y[partialBase + (m0 + {{ row }}u) * PARTIAL_COLS + n0 + {{ column }}u] = acc{{ row }}.{{ components[column] }};
 
 
401
  {% else %}
402
  store_out(m0 + {{ row }}u, n0 + {{ column }}u, yBatchBase, acc{{ row }}.{{ components[column] }});
403
  {% endif %}
404
  {% endfor %}
405
+ {% endfor %}}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
build/webgpu/conv-1x1-gemm-tiled.wgsl.jinja CHANGED
@@ -6,14 +6,11 @@
6
  // f16 operands remain packed in workgroup memory and widen when consumed. All
7
  // M, N, and K accesses are bounds-checked.
8
  {{ env.wgsl.resourceDeclarations }}
9
- {% set hasActivation = hasActivation is defined and hasActivation %}
10
- {% set activation = activation | default("") %}
11
- {% set actAlpha = actAlpha | default(0.0) %}
12
- {% set actBeta = actBeta | default(0.0) %}
13
  {% if hasActivation %}
14
  // Apply the fused activation in the f32 accumulator before the single output
15
  // cast, avoiding an intermediate convolution tensor.
16
- {% macro fused_act_return(mode, alpha, beta) -%}
 
17
  {% if mode == "Relu" %}
18
  return max(v, 0.0);
19
  {% elif mode == "Clip" %}
@@ -28,23 +25,55 @@
28
  return tanh(clamp(v, -10.0, 10.0));
29
  {% elif mode == "HardSigmoid" %}
30
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31
  {% else %}
32
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
33
  {% endif %}
34
- {%- endmacro -%}
35
- fn fused_act(v: f32) -> f32 {
36
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
37
- }
38
  {% endif %}
39
-
40
- {% set tileT = "f16" if usesF16 else "f32" -%}
 
 
41
  const BK: u32 = 16u;
42
  const BM: u32 = 32u;
43
  const BN: u32 = 32u;
44
 
45
  // Store one output element with the configured bias/residual/activation
46
  // epilogue in the f32 accumulator domain.
47
- {% set hasZ = hasZ is defined and hasZ %}
48
  fn store_out(m: u32, n: u32, yBase: u32, raw: f32) {
49
  if (m >= params.M || n >= params.N) {
50
  return;
@@ -62,11 +91,10 @@ fn store_out(m: u32, n: u32, yBase: u32, raw: f32) {
62
  y[yBase + m * params.N + n] = {{ T }}(v);
63
  }
64
 
65
-
66
  var<workgroup> tileA: array<array<{{ tileT }}, BK>, BM>;
67
  var<workgroup> tileB: array<array<{{ tileT }}, BN>, BK>;
68
 
69
- @compute @workgroup_size(16, 16, 1)
70
  fn main(
71
  @builtin(workgroup_id) wg: vec3<u32>,
72
  @builtin(local_invocation_id) lid: vec3<u32>
@@ -78,7 +106,7 @@ fn main(
78
  let nBase = wg.x * BN;
79
  let batch = wg.z;
80
  let xBatchBase = batch * K * N;
81
- let li = lid.y * 16u + lid.x;
82
 
83
  var acc00: f32 = 0.0;
84
  var acc01: f32 = 0.0;
@@ -88,8 +116,8 @@ fn main(
88
  for (var kt: u32 = 0u; kt < numTiles; kt = kt + 1u) {
89
  let kBase = kt * BK;
90
  // Cooperative load: 32x16 A(=W) tile + 16x32 B(=X[batch]) tile, 256 threads x 2 each.
91
- for (var e: u32 = 0u; e < (BM * BK) / 256u; e = e + 1u) {
92
- let idx = li + e * 256u;
93
  let ar = idx / BK;
94
  let ac = idx % BK;
95
  let am = mBase + ar;
@@ -110,16 +138,26 @@ fn main(
110
  }
111
  }
112
  workgroupBarrier();
 
 
 
 
 
 
113
  for (var kk: u32 = 0u; kk < BK; kk = kk + 1u) {
114
  let a0 = f32(tileA[lid.y * 2u][kk]);
115
  let a1 = f32(tileA[lid.y * 2u + 1u][kk]);
116
  let b0 = f32(tileB[kk][lid.x * 2u]);
117
  let b1 = f32(tileB[kk][lid.x * 2u + 1u]);
118
- acc00 = acc00 + a0 * b0;
119
- acc01 = acc01 + a0 * b1;
120
- acc10 = acc10 + a1 * b0;
121
- acc11 = acc11 + a1 * b1;
122
  }
 
 
 
 
 
123
  workgroupBarrier();
124
  }
125
 
 
6
  // f16 operands remain packed in workgroup memory and widen when consumed. All
7
  // M, N, and K accesses are bounds-checked.
8
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
9
  {% if hasActivation %}
10
  // Apply the fused activation in the f32 accumulator before the single output
11
  // cast, avoiding an intermediate convolution tensor.
12
+ {% macro fused_act_fn(mode, alpha, beta) %}
13
+ fn fused_act(v: f32) -> f32 {
14
  {% if mode == "Relu" %}
15
  return max(v, 0.0);
16
  {% elif mode == "Clip" %}
 
25
  return tanh(clamp(v, -10.0, 10.0));
26
  {% elif mode == "HardSigmoid" %}
27
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
28
+ {% elif mode == "QuickGelu" %}
29
+ {% if alpha == 1 %}
30
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
31
+ return v * (1.0 / (1.0 + exp(-v)));
32
+ {% else %}
33
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
34
+ {% endif %}
35
+ {% elif mode == "Elu" %}
36
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
37
+ {% elif mode == "ThresholdedRelu" %}
38
+ return select(0.0, v, v > f32({{ alpha }}));
39
+ {% elif mode == "Softplus" %}
40
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
41
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
42
+ {% elif mode == "Erf" %}
43
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
44
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
45
+ let a = abs(v);
46
+ let t = 1.0 / (1.0 + 0.3275911 * a);
47
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
48
+ return sign(v) * (1.0 - poly * exp(-a * a));
49
+ {% elif mode == "Gelu" and alpha == 0 %}
50
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
51
+ let u = v * 0.70710678118654752;
52
+ let a = abs(u);
53
+ let t = 1.0 / (1.0 + 0.3275911 * a);
54
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
55
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
56
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
57
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
58
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
59
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
60
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
61
  {% else %}
62
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
63
  {% endif %}
64
+ }{% endmacro %}
65
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
66
  {% endif %}
67
+ {% set accumulatorPrefix = "sum" if tilePartialAccumulation else "acc" %}
68
+ {% set tileT = "f16" if usesF16 else "f32" %}
69
+ {% set wgX = 16 %}
70
+ {% set wgY = 16 %}
71
  const BK: u32 = 16u;
72
  const BM: u32 = 32u;
73
  const BN: u32 = 32u;
74
 
75
  // Store one output element with the configured bias/residual/activation
76
  // epilogue in the f32 accumulator domain.
 
77
  fn store_out(m: u32, n: u32, yBase: u32, raw: f32) {
78
  if (m >= params.M || n >= params.N) {
79
  return;
 
91
  y[yBase + m * params.N + n] = {{ T }}(v);
92
  }
93
 
 
94
  var<workgroup> tileA: array<array<{{ tileT }}, BK>, BM>;
95
  var<workgroup> tileB: array<array<{{ tileT }}, BN>, BK>;
96
 
97
+ @compute @workgroup_size({{ wgX }}, {{ wgY }}, 1)
98
  fn main(
99
  @builtin(workgroup_id) wg: vec3<u32>,
100
  @builtin(local_invocation_id) lid: vec3<u32>
 
106
  let nBase = wg.x * BN;
107
  let batch = wg.z;
108
  let xBatchBase = batch * K * N;
109
+ let li = lid.y * {{ wgX }}u + lid.x;
110
 
111
  var acc00: f32 = 0.0;
112
  var acc01: f32 = 0.0;
 
116
  for (var kt: u32 = 0u; kt < numTiles; kt = kt + 1u) {
117
  let kBase = kt * BK;
118
  // Cooperative load: 32x16 A(=W) tile + 16x32 B(=X[batch]) tile, 256 threads x 2 each.
119
+ for (var e: u32 = 0u; e < (BM * BK) / {{ wgX * wgY }}u; e = e + 1u) {
120
+ let idx = li + e * {{ wgX * wgY }}u;
121
  let ar = idx / BK;
122
  let ac = idx % BK;
123
  let am = mBase + ar;
 
138
  }
139
  }
140
  workgroupBarrier();
141
+ {% if tilePartialAccumulation %}
142
+ // Bound each serial dot-product chain by the K tile before merging it.
143
+ {% for i in range(2) %}{% for j in range(2) %}
144
+ var sum{{ i }}{{ j }}: f32 = 0.0;
145
+ {% endfor %}{% endfor %}
146
+ {% endif %}
147
  for (var kk: u32 = 0u; kk < BK; kk = kk + 1u) {
148
  let a0 = f32(tileA[lid.y * 2u][kk]);
149
  let a1 = f32(tileA[lid.y * 2u + 1u][kk]);
150
  let b0 = f32(tileB[kk][lid.x * 2u]);
151
  let b1 = f32(tileB[kk][lid.x * 2u + 1u]);
152
+ {% for i in range(2) %}{% for j in range(2) %}
153
+ {{ accumulatorPrefix }}{{ i }}{{ j }} = {{ accumulatorPrefix }}{{ i }}{{ j }} + a{{ i }} * b{{ j }};
154
+ {% endfor %}{% endfor %}
 
155
  }
156
+ {% if tilePartialAccumulation %}
157
+ {% for i in range(2) %}{% for j in range(2) %}
158
+ acc{{ i }}{{ j }} = acc{{ i }}{{ j }} + sum{{ i }}{{ j }};
159
+ {% endfor %}{% endfor %}
160
+ {% endif %}
161
  workgroupBarrier();
162
  }
163
 
build/webgpu/conv-1x1-subgroup-matrix.wgsl.jinja CHANGED
@@ -12,18 +12,13 @@
12
  // the slices and applies bias or an epilogue. Padded tail cells remain zero and
13
  // are never copied to the logical output.
14
  {% set directMatrixInputs = directMatrixInputs is defined and directMatrixInputs %}
15
- {% set fusedNarrowProjection = fusedNarrowProjection if fusedNarrowProjection is defined else false %}
16
- {% set polyphase = polyphaseConvTranspose is defined and polyphaseConvTranspose %}
17
  {% set implicit = implicitIm2col is defined and implicitIm2col %}
18
- {% set implicit3d = implicit and (convSpatialDims is defined and convSpatialDims == 3) %}
19
  {% set splitKValue = splitK if splitK is defined else 1 %}
20
  {% set splitKPartial = splitKValue > 1 %}
21
  {% set kLoopVar = "K_LOOP" if padded else "K" %}
22
  {% set nColsVar = "N_COLS" if padded else "N" %}
23
- {% set hasZ = hasZ is defined and hasZ %}
24
- {% set projectionOutChannels = projectionOutChannels | default(0) %}
25
- {% set hasProjectionBias = hasProjectionBias is defined and hasProjectionBias %}
26
- {% set hasOutputScale = hasOutputScale is defined and hasOutputScale %}
27
  {% set convKernelH = convKernelH | default(0) %}
28
  {% set convKernelW = convKernelW | default(0) %}
29
  {% set convStrideH = convStrideH | default(0) %}
@@ -36,15 +31,9 @@
36
  {% set convInW = convInW | default(0) %}
37
  {% set convOutW = convOutW | default(0) %}
38
  {% set convInChannels = convInChannels | default(0) %}
39
- {% set convKernelD = convKernelD | default(0) %}
40
- {% set convStrideD = convStrideD | default(0) %}
41
- {% set convDilationD = convDilationD | default(0) %}
42
- {% set convPadFront = convPadFront | default(0) %}
43
- {% set convInD = convInD | default(0) %}
44
- {% set convOutH = convOutH | default(0) %}
45
  {% set nPadded = nPadded | default(0) %}
46
  {% set mPadded = mPadded | default(0) %}
47
- {% set kChunk = kChunk | default(0) %}
48
  {% set batchCount = batchCount | default(0) %}
49
  enable subgroups;
50
  {% if pinSubgroupSize32 %}
@@ -53,46 +42,16 @@ enable subgroup_size_control;
53
  enable chromium_experimental_subgroup_matrix;
54
  diagnostic(off, chromium.subgroup_matrix_uniformity);
55
 
56
-
57
  {{ env.wgsl.resourceDeclarations }}
58
- {% if fusedNarrowProjection %}{% set inputActivation = inputActivation | default("none") %}
59
- {% set outputActivation = outputActivation | default("none") %}
60
- {% if inputActivation == "relu" or outputActivation == "relu" %}
61
- fn is_nan_f32(value: f32) -> bool {
62
- let bits = bitcast<u32>(value);
63
- return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
64
- }
65
-
66
- fn projection_relu(value: f32) -> f32 {
67
- var out = max(value, 0.0);
68
- if (is_nan_f32(value)) {
69
- out = value;
70
- }
71
- return out;
72
- }
73
- {% endif %}
74
- {% if outputActivation == "sigmoid" %}
75
- fn sigmoid_safe(x: f32) -> f32 {
76
- if (x >= 0.0) {
77
- let z = exp(-x);
78
- return 1.0 / (1.0 + z);
79
- }
80
- let z = exp(x);
81
- return z / (1.0 + z);
82
- }
83
- {% endif %}
84
- {% endif %}
85
-
86
  {% set operandScalar = fScalar %}
87
  {% set accScalar = "f32" %}
88
  {% set useDirectMatrixStore = directMatrixStore is defined and directMatrixStore %}
89
  {% set tileRowsValue = tileRows if tileRows is defined else 32 %}
90
- {% set tileColsValue = tileCols if tileCols is defined else 64 %}
91
  {% set workgroupThreadsValue = workgroupThreads if workgroupThreads is defined else 128 %}
92
  {% set subgroupRowsValue = subgroupRows if subgroupRows is defined else 2 %}
93
- {% set subgroupColsValue = subgroupCols if subgroupCols is defined else 2 %}
94
  {% set subRowsValue = (tileRowsValue / subgroupRowsValue)|int %}
95
- {% set subColsValue = (tileColsValue / subgroupColsValue)|int %}
96
  {% set bLoadWidth = (tileColsValue * 32 / workgroupThreadsValue)|int %}
97
  {% set bKChunks = (32 / bLoadWidth)|int %}
98
 
@@ -107,18 +66,7 @@ const K_LOOP: u32 = {{ kPadded }}u;
107
  {% if implicit %}
108
  const CONV_KERNEL_H: u32 = {{ convKernelH }}u;
109
  const CONV_KERNEL_W: u32 = {{ convKernelW }}u;
110
- {% if implicit3d %}
111
- const CONV_KERNEL_D: u32 = {{ convKernelD }}u;
112
- const CONV_KSIZE_HW: u32 = CONV_KERNEL_H * CONV_KERNEL_W;
113
- const CONV_KSIZE: u32 = CONV_KERNEL_D * CONV_KSIZE_HW;
114
- const CONV_STRIDE_D: u32 = {{ convStrideD }}u;
115
- const CONV_DILATION_D: u32 = {{ convDilationD }}u;
116
- const CONV_PAD_FRONT: i32 = {{ convPadFront }};
117
- const CONV_IN_D: i32 = {{ convInD }};
118
- const CONV_OUT_H: u32 = {{ convOutH }}u;
119
- {% else %}
120
  const CONV_KSIZE: u32 = CONV_KERNEL_H * CONV_KERNEL_W;
121
- {% endif %}
122
  const CONV_STRIDE_H: u32 = {{ convStrideH }}u;
123
  const CONV_STRIDE_W: u32 = {{ convStrideW }}u;
124
  const CONV_DILATION_H: u32 = {{ convDilationH }}u;
@@ -128,13 +76,8 @@ const CONV_PAD_LEFT: i32 = {{ convPadLeft }};
128
  const CONV_IN_H: i32 = {{ convInH }};
129
  const CONV_IN_W: i32 = {{ convInW }};
130
  const CONV_OUT_W: u32 = {{ convOutW }}u;
131
- {% if implicit3d %}
132
- // X is [batch, inChannels, inD, inH, inW]; the tile gather indexes it directly.
133
- const B_BATCH_STRIDE: u32 = {{ convInChannels }}u * u32(CONV_IN_D) * u32(CONV_IN_H) * u32(CONV_IN_W);
134
- {% else %}
135
  // X is [batch, inChannels, inH, inW]; the tile gather indexes it directly.
136
  const B_BATCH_STRIDE: u32 = {{ convInChannels }}u * u32(CONV_IN_H) * u32(CONV_IN_W);
137
- {% endif %}
138
  {% else %}
139
  {% if padded %}
140
  const N_COLS: u32 = {{ nPadded }}u;
@@ -176,7 +119,7 @@ fn loadSHMA(tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
176
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
177
  let k = k_idx + col + col_offset;
178
  if (a_global < M{% if padded %} && k < K{% endif %}) {
179
- {% set A_PHASE = "phase_base + " if polyphase else "" %}
180
  {% if operandScalar == "f16" %}
181
  tile_A[row * TILE_K + col + col_offset] = f16(w[{{ A_PHASE }}a_global * K + k]);
182
  {% else %}
@@ -197,12 +140,7 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
197
  let col = c_idx * {{ bLoadWidth }}u;
198
  {% if implicit %}
199
  // One output position per tile column, so its window origin is loop-invariant.
200
- {% if implicit3d %}
201
- let id0 = i32((b_col / (CONV_OUT_W * CONV_OUT_H)) * CONV_STRIDE_D) - CONV_PAD_FRONT;
202
- let ih0 = i32(((b_col / CONV_OUT_W) % CONV_OUT_H) * CONV_STRIDE_H) - CONV_PAD_TOP;
203
- {% else %}
204
  let ih0 = i32((b_col / CONV_OUT_W) * CONV_STRIDE_H) - CONV_PAD_TOP;
205
- {% endif %}
206
  let iw0 = i32((b_col % CONV_OUT_W) * CONV_STRIDE_W) - CONV_PAD_LEFT;
207
  let in_column = b_col < N;
208
  // k advances by exactly one per element, so the (in-channel, tap-row, tap-col)
@@ -213,16 +151,8 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
213
  var k = k_idx + col;
214
  var ic = k / CONV_KSIZE;
215
  var kq = k % CONV_KSIZE;
216
- {% if implicit3d %}
217
- var kd = kq / CONV_KSIZE_HW;
218
- var khw = kq % CONV_KSIZE_HW;
219
- var kh = khw / CONV_KERNEL_W;
220
- var kw = khw % CONV_KERNEL_W;
221
- var id = id0 + i32(kd * CONV_DILATION_D);
222
- {% else %}
223
  var kh = kq / CONV_KERNEL_W;
224
  var kw = kq % CONV_KERNEL_W;
225
- {% endif %}
226
  var ih = ih0 + i32(kh * CONV_DILATION_H);
227
  var iw = iw0 + i32(kw * CONV_DILATION_W);
228
  // A tile column whose whole kernel window and K span are in bounds cannot
@@ -230,30 +160,16 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
230
  // the interior loop carries a plain address instead of recomputing coordinates.
231
  let interior = in_column
232
  && k_idx + col + {{ bLoadWidth }}u <= K
233
- {% if implicit3d %}
234
- && id0 >= 0 && id0 + i32((CONV_KERNEL_D - 1u) * CONV_DILATION_D) < CONV_IN_D
235
- {% endif %}
236
  && ih0 >= 0 && ih0 + i32((CONV_KERNEL_H - 1u) * CONV_DILATION_H) < CONV_IN_H
237
  && iw0 >= 0 && iw0 + i32((CONV_KERNEL_W - 1u) * CONV_DILATION_W) < CONV_IN_W;
238
  if (interior) {
239
  let colStep = i32(CONV_DILATION_W);
240
  let rowStep = i32(CONV_DILATION_H) * CONV_IN_W;
241
- {% if implicit3d %}
242
- let depthStep = i32(CONV_DILATION_D) * CONV_IN_H * CONV_IN_W;
243
- let planeStep = CONV_IN_D * CONV_IN_H * CONV_IN_W;
244
- {% else %}
245
  let planeStep = CONV_IN_H * CONV_IN_W;
246
- {% endif %}
247
  // Advancing k already added one colStep, so each wrap rewinds that plus the taps it walked.
248
  let kwWrap = colStep * i32(CONV_KERNEL_W) - rowStep;
249
- {% if implicit3d %}
250
- let khWrap = rowStep * i32(CONV_KERNEL_H) - depthStep;
251
- let kdWrap = depthStep * i32(CONV_KERNEL_D) - planeStep;
252
- var addr = i32(b_base) + i32(ic) * planeStep + (id * CONV_IN_H + ih) * CONV_IN_W + iw;
253
- {% else %}
254
  let khWrap = rowStep * i32(CONV_KERNEL_H) - planeStep;
255
  var addr = i32(b_base) + i32(ic) * planeStep + ih * CONV_IN_W + iw;
256
- {% endif %}
257
  for (var i = 0u; i < {{ bLoadWidth }}u; i = i + 1u) {
258
  tile_B[row * TILE_K + col + i] = {{ operandScalar }}(xm[u32(addr)]);
259
  addr = addr + colStep;
@@ -265,13 +181,6 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
265
  if (kh == CONV_KERNEL_H) {
266
  kh = 0u;
267
  addr = addr - khWrap;
268
- {% if implicit3d %}
269
- kd = kd + 1u;
270
- if (kd == CONV_KERNEL_D) {
271
- kd = 0u;
272
- addr = addr - kdWrap;
273
- }
274
- {% endif %}
275
  }
276
  }
277
  }
@@ -281,12 +190,8 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
281
  var value = {{ operandScalar }}(0.0);
282
  // The K and N tails read as exact zeros, exactly as the materialized padded
283
  // cols buffer does, so they contribute nothing to the dot product.
284
- if (in_column && k < K{% if implicit3d %} && id >= 0 && id < CONV_IN_D{% endif %} && ih >= 0 && ih < CONV_IN_H && iw >= 0 && iw < CONV_IN_W) {
285
- {% if implicit3d %}
286
- value = {{ operandScalar }}(xm[b_base + ((ic * u32(CONV_IN_D) + u32(id)) * u32(CONV_IN_H) + u32(ih)) * u32(CONV_IN_W) + u32(iw)]);
287
- {% else %}
288
  value = {{ operandScalar }}(xm[b_base + (ic * u32(CONV_IN_H) + u32(ih)) * u32(CONV_IN_W) + u32(iw)]);
289
- {% endif %}
290
  }
291
  tile_B[row * TILE_K + col + i] = value;
292
  k = k + 1u;
@@ -300,17 +205,7 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
300
  if (kh == CONV_KERNEL_H) {
301
  kh = 0u;
302
  ih = ih0;
303
- {% if implicit3d %}
304
- kd = kd + 1u;
305
- id = id + i32(CONV_DILATION_D);
306
- if (kd == CONV_KERNEL_D) {
307
- kd = 0u;
308
- id = id0;
309
- ic = ic + 1u;
310
- }
311
- {% else %}
312
  ic = ic + 1u;
313
- {% endif %}
314
  }
315
  }
316
  }
@@ -325,16 +220,15 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
325
  }
326
  {% endif %}
327
  }
 
328
 
329
  {% endif %}
330
- {% set hasActivation = hasActivation is defined and hasActivation %}
331
- {% set activation = activation | default("") %}
332
- {% set actAlpha = actAlpha | default(0.0) %}
333
- {% set actBeta = actBeta | default(0.0) %}
334
  {% if hasActivation %}
335
  // Apply the fused activation in the f32 accumulator before the single output
336
  // cast, avoiding an intermediate convolution tensor.
337
- {% macro fused_act_return(mode, alpha, beta) -%}
 
338
  {% if mode == "Relu" %}
339
  return max(v, 0.0);
340
  {% elif mode == "Clip" %}
@@ -349,15 +243,45 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
349
  return tanh(clamp(v, -10.0, 10.0));
350
  {% elif mode == "HardSigmoid" %}
351
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
352
  {% else %}
353
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
354
  {% endif %}
355
- {%- endmacro -%}
356
- fn fused_act(v: f32) -> f32 {
357
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
358
- }
359
  {% endif %}
360
-
361
  {% set biasAdd = " + bv" if hasBias else "" %}
362
  {% if hasActivation or hasZ %}
363
  // Fused epilogue: Y = activation(conv + bias + Z), applied at the output store.
@@ -374,10 +298,9 @@ fn epi(raw: f32{% if hasZ %}, yIndex: u32{% endif %}) -> {{ T }} {
374
  }
375
  {% endif %}
376
  {% macro store_val(valExpr, idxExpr) %}
377
- {% if hasActivation or hasZ %}epi({{ valExpr }}{% if hasZ %}, {{ idxExpr }}{% endif %})
378
- {%- else %}{{ T }}({{ valExpr }})
379
- {%- endif %}
380
- {% endmacro %}
381
 
382
  {% if (not useDirectMatrixStore or padded) and not splitKPartial and not fusedNarrowProjection %}
383
  fn storeOutput(offset: u32, {% if hasBias or polyphase %}row_base: u32, {% endif %}row: u32, col: u32, src_slot: u32, row_limit: i32{% if padded %}, col_anchor: u32{% endif %}) {
@@ -431,9 +354,7 @@ fn main(
431
  {% else %}
432
  let batch = workgroup_id.z;
433
  let b_base = batch * B_BATCH_STRIDE;
434
- {% if not fusedNarrowProjection %}
435
  let c_base = batch * C_BATCH_STRIDE;
436
- {% endif %}
437
  {% endif %}
438
  let a_global_base = workgroup_id.y * TILE_ROWS;
439
  let b_global_base = workgroup_id.x * TILE_COLS;
@@ -491,57 +412,7 @@ fn main(
491
  {% endif %}
492
  }
493
 
494
- {% if fusedNarrowProjection %}
495
- // The Conv producer owns exactly one 32-row tile. Its input tile is dead after
496
- // the final K iteration, so reuse those 32x64 f32 cells as the only
497
- // cross-subgroup handoff into the narrow projection. All collective stores
498
- // remain subgroup-uniform, and every fragment has a disjoint destination.
499
- let fused_tile_offset = base_A * TILE_COLS + base_B;
500
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 0u * TILE_COLS + 0u, matC00, TILE_COLS);
501
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 0u * TILE_COLS + 8u, matC01, TILE_COLS);
502
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 0u * TILE_COLS + 16u, matC02, TILE_COLS);
503
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 0u * TILE_COLS + 24u, matC03, TILE_COLS);
504
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 8u * TILE_COLS + 0u, matC10, TILE_COLS);
505
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 8u * TILE_COLS + 8u, matC11, TILE_COLS);
506
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 8u * TILE_COLS + 16u, matC12, TILE_COLS);
507
- subgroupMatrixStore<row_major>(&tile_B, fused_tile_offset + 8u * TILE_COLS + 24u, matC13, TILE_COLS);
508
- workgroupBarrier();
509
-
510
- // One invocation owns one spatial column and walks the producer channels in
511
- // increasing order to preserve the scalar projection's accumulation order.
512
- // The inactive half of the workgroup has no remaining collective operation
513
- // to reach.
514
- let fused_global_col = b_global_base + local_idx;
515
- if (local_idx < TILE_COLS && fused_global_col < N) {
516
- {% for oc in range(projectionOutChannels) %}
517
- var projected_acc{{ oc }} = f32(0.0);
518
- {% endfor %}
519
- for (var channel = 0u; channel < M; channel = channel + 1u) {
520
- let producer_value = tile_B[channel * TILE_COLS + local_idx]{% if hasBias %} + f32(bias[channel]){% endif %};
521
- {% if inputActivation == "relu" %}
522
- let activated_value = projection_relu(producer_value);
523
- {% else %}
524
- let activated_value = producer_value;
525
- {% endif %}
526
- {% for oc in range(projectionOutChannels) %}
527
- projected_acc{{ oc }} = projected_acc{{ oc }} + activated_value * f32(projectionW[{{ oc }}u * M + channel]);
528
- {% endfor %}
529
- }
530
-
531
- let fused_y_base = batch * {{ projectionOutChannels }}u * N + fused_global_col;
532
- {% for oc in range(projectionOutChannels) %}
533
- let projected{{ oc }} = projected_acc{{ oc }}{% if hasProjectionBias %} + f32(projectionBias[{{ oc }}u]){% endif %};
534
- {% if outputActivation == "relu" %}
535
- let activated{{ oc }} = projection_relu(projected{{ oc }});
536
- {% elif outputActivation == "sigmoid" %}
537
- let activated{{ oc }} = sigmoid_safe(projected{{ oc }});
538
- {% else %}
539
- let activated{{ oc }} = projected{{ oc }};
540
- {% endif %}
541
- y[fused_y_base + {{ oc }}u * N] = activated{{ oc }}{% if hasOutputScale %} * params.outputScale{% endif %};
542
- {% endfor %}
543
- }
544
- {% elif splitKPartial %}
545
  // Raw partials, no bias and no epilogue — the reduce pass owns both. The
546
  // scratch is whole tiles in both axes, so this needs none of the column or
547
  // row guards the real output store carries.
 
12
  // the slices and applies bias or an epilogue. Padded tail cells remain zero and
13
  // are never copied to the logical output.
14
  {% set directMatrixInputs = directMatrixInputs is defined and directMatrixInputs %}
15
+ {% set fusedNarrowProjection = false %}
16
+ {% set polyphase = false %}
17
  {% set implicit = implicitIm2col is defined and implicitIm2col %}
 
18
  {% set splitKValue = splitK if splitK is defined else 1 %}
19
  {% set splitKPartial = splitKValue > 1 %}
20
  {% set kLoopVar = "K_LOOP" if padded else "K" %}
21
  {% set nColsVar = "N_COLS" if padded else "N" %}
 
 
 
 
22
  {% set convKernelH = convKernelH | default(0) %}
23
  {% set convKernelW = convKernelW | default(0) %}
24
  {% set convStrideH = convStrideH | default(0) %}
 
31
  {% set convInW = convInW | default(0) %}
32
  {% set convOutW = convOutW | default(0) %}
33
  {% set convInChannels = convInChannels | default(0) %}
 
 
 
 
 
 
34
  {% set nPadded = nPadded | default(0) %}
35
  {% set mPadded = mPadded | default(0) %}
36
+ {% set kChunk = 0 %}
37
  {% set batchCount = batchCount | default(0) %}
38
  enable subgroups;
39
  {% if pinSubgroupSize32 %}
 
42
  enable chromium_experimental_subgroup_matrix;
43
  diagnostic(off, chromium.subgroup_matrix_uniformity);
44
 
 
45
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
  {% set operandScalar = fScalar %}
47
  {% set accScalar = "f32" %}
48
  {% set useDirectMatrixStore = directMatrixStore is defined and directMatrixStore %}
49
  {% set tileRowsValue = tileRows if tileRows is defined else 32 %}
50
+ {% set tileColsValue = 64 %}
51
  {% set workgroupThreadsValue = workgroupThreads if workgroupThreads is defined else 128 %}
52
  {% set subgroupRowsValue = subgroupRows if subgroupRows is defined else 2 %}
 
53
  {% set subRowsValue = (tileRowsValue / subgroupRowsValue)|int %}
54
+ {% set subColsValue = 32 %}
55
  {% set bLoadWidth = (tileColsValue * 32 / workgroupThreadsValue)|int %}
56
  {% set bKChunks = (32 / bLoadWidth)|int %}
57
 
 
66
  {% if implicit %}
67
  const CONV_KERNEL_H: u32 = {{ convKernelH }}u;
68
  const CONV_KERNEL_W: u32 = {{ convKernelW }}u;
 
 
 
 
 
 
 
 
 
 
69
  const CONV_KSIZE: u32 = CONV_KERNEL_H * CONV_KERNEL_W;
 
70
  const CONV_STRIDE_H: u32 = {{ convStrideH }}u;
71
  const CONV_STRIDE_W: u32 = {{ convStrideW }}u;
72
  const CONV_DILATION_H: u32 = {{ convDilationH }}u;
 
76
  const CONV_IN_H: i32 = {{ convInH }};
77
  const CONV_IN_W: i32 = {{ convInW }};
78
  const CONV_OUT_W: u32 = {{ convOutW }}u;
 
 
 
 
79
  // X is [batch, inChannels, inH, inW]; the tile gather indexes it directly.
80
  const B_BATCH_STRIDE: u32 = {{ convInChannels }}u * u32(CONV_IN_H) * u32(CONV_IN_W);
 
81
  {% else %}
82
  {% if padded %}
83
  const N_COLS: u32 = {{ nPadded }}u;
 
119
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
120
  let k = k_idx + col + col_offset;
121
  if (a_global < M{% if padded %} && k < K{% endif %}) {
122
+ {% set A_PHASE = "" %}
123
  {% if operandScalar == "f16" %}
124
  tile_A[row * TILE_K + col + col_offset] = f16(w[{{ A_PHASE }}a_global * K + k]);
125
  {% else %}
 
140
  let col = c_idx * {{ bLoadWidth }}u;
141
  {% if implicit %}
142
  // One output position per tile column, so its window origin is loop-invariant.
 
 
 
 
143
  let ih0 = i32((b_col / CONV_OUT_W) * CONV_STRIDE_H) - CONV_PAD_TOP;
 
144
  let iw0 = i32((b_col % CONV_OUT_W) * CONV_STRIDE_W) - CONV_PAD_LEFT;
145
  let in_column = b_col < N;
146
  // k advances by exactly one per element, so the (in-channel, tap-row, tap-col)
 
151
  var k = k_idx + col;
152
  var ic = k / CONV_KSIZE;
153
  var kq = k % CONV_KSIZE;
 
 
 
 
 
 
 
154
  var kh = kq / CONV_KERNEL_W;
155
  var kw = kq % CONV_KERNEL_W;
 
156
  var ih = ih0 + i32(kh * CONV_DILATION_H);
157
  var iw = iw0 + i32(kw * CONV_DILATION_W);
158
  // A tile column whose whole kernel window and K span are in bounds cannot
 
160
  // the interior loop carries a plain address instead of recomputing coordinates.
161
  let interior = in_column
162
  && k_idx + col + {{ bLoadWidth }}u <= K
 
 
 
163
  && ih0 >= 0 && ih0 + i32((CONV_KERNEL_H - 1u) * CONV_DILATION_H) < CONV_IN_H
164
  && iw0 >= 0 && iw0 + i32((CONV_KERNEL_W - 1u) * CONV_DILATION_W) < CONV_IN_W;
165
  if (interior) {
166
  let colStep = i32(CONV_DILATION_W);
167
  let rowStep = i32(CONV_DILATION_H) * CONV_IN_W;
 
 
 
 
168
  let planeStep = CONV_IN_H * CONV_IN_W;
 
169
  // Advancing k already added one colStep, so each wrap rewinds that plus the taps it walked.
170
  let kwWrap = colStep * i32(CONV_KERNEL_W) - rowStep;
 
 
 
 
 
171
  let khWrap = rowStep * i32(CONV_KERNEL_H) - planeStep;
172
  var addr = i32(b_base) + i32(ic) * planeStep + ih * CONV_IN_W + iw;
 
173
  for (var i = 0u; i < {{ bLoadWidth }}u; i = i + 1u) {
174
  tile_B[row * TILE_K + col + i] = {{ operandScalar }}(xm[u32(addr)]);
175
  addr = addr + colStep;
 
181
  if (kh == CONV_KERNEL_H) {
182
  kh = 0u;
183
  addr = addr - khWrap;
 
 
 
 
 
 
 
184
  }
185
  }
186
  }
 
190
  var value = {{ operandScalar }}(0.0);
191
  // The K and N tails read as exact zeros, exactly as the materialized padded
192
  // cols buffer does, so they contribute nothing to the dot product.
193
+ if (in_column && k < K && ih >= 0 && ih < CONV_IN_H && iw >= 0 && iw < CONV_IN_W) {
 
 
 
194
  value = {{ operandScalar }}(xm[b_base + (ic * u32(CONV_IN_H) + u32(ih)) * u32(CONV_IN_W) + u32(iw)]);
 
195
  }
196
  tile_B[row * TILE_K + col + i] = value;
197
  k = k + 1u;
 
205
  if (kh == CONV_KERNEL_H) {
206
  kh = 0u;
207
  ih = ih0;
 
 
 
 
 
 
 
 
 
208
  ic = ic + 1u;
 
209
  }
210
  }
211
  }
 
220
  }
221
  {% endif %}
222
  }
223
+ {% if hasActivation or hasZ %}
224
 
225
  {% endif %}
226
+ {% endif %}
 
 
 
227
  {% if hasActivation %}
228
  // Apply the fused activation in the f32 accumulator before the single output
229
  // cast, avoiding an intermediate convolution tensor.
230
+ {% macro fused_act_fn(mode, alpha, beta) %}
231
+ fn fused_act(v: f32) -> f32 {
232
  {% if mode == "Relu" %}
233
  return max(v, 0.0);
234
  {% elif mode == "Clip" %}
 
243
  return tanh(clamp(v, -10.0, 10.0));
244
  {% elif mode == "HardSigmoid" %}
245
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
246
+ {% elif mode == "QuickGelu" %}
247
+ {% if alpha == 1 %}
248
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
249
+ return v * (1.0 / (1.0 + exp(-v)));
250
+ {% else %}
251
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
252
+ {% endif %}
253
+ {% elif mode == "Elu" %}
254
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
255
+ {% elif mode == "ThresholdedRelu" %}
256
+ return select(0.0, v, v > f32({{ alpha }}));
257
+ {% elif mode == "Softplus" %}
258
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
259
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
260
+ {% elif mode == "Erf" %}
261
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
262
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
263
+ let a = abs(v);
264
+ let t = 1.0 / (1.0 + 0.3275911 * a);
265
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
266
+ return sign(v) * (1.0 - poly * exp(-a * a));
267
+ {% elif mode == "Gelu" and alpha == 0 %}
268
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
269
+ let u = v * 0.70710678118654752;
270
+ let a = abs(u);
271
+ let t = 1.0 / (1.0 + 0.3275911 * a);
272
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
273
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
274
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
275
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
276
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
277
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
278
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
279
  {% else %}
280
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
281
  {% endif %}
282
+ }{% endmacro %}
283
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
284
  {% endif %}
 
285
  {% set biasAdd = " + bv" if hasBias else "" %}
286
  {% if hasActivation or hasZ %}
287
  // Fused epilogue: Y = activation(conv + bias + Z), applied at the output store.
 
298
  }
299
  {% endif %}
300
  {% macro store_val(valExpr, idxExpr) %}
301
+ {% if hasActivation or hasZ %}
302
+ epi({{ valExpr }}{% if hasZ %}, {{ idxExpr }}{% endif %}){% else %}
303
+ {{ T }}({{ valExpr }}){% endif %}{% endmacro %}
 
304
 
305
  {% if (not useDirectMatrixStore or padded) and not splitKPartial and not fusedNarrowProjection %}
306
  fn storeOutput(offset: u32, {% if hasBias or polyphase %}row_base: u32, {% endif %}row: u32, col: u32, src_slot: u32, row_limit: i32{% if padded %}, col_anchor: u32{% endif %}) {
 
354
  {% else %}
355
  let batch = workgroup_id.z;
356
  let b_base = batch * B_BATCH_STRIDE;
 
357
  let c_base = batch * C_BATCH_STRIDE;
 
358
  {% endif %}
359
  let a_global_base = workgroup_id.y * TILE_ROWS;
360
  let b_global_base = workgroup_id.x * TILE_COLS;
 
412
  {% endif %}
413
  }
414
 
415
+ {% if splitKPartial %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
416
  // Raw partials, no bias and no epilogue — the reduce pass owns both. The
417
  // scratch is whole tiles in both axes, so this needs none of the column or
418
  // row guards the real output store carries.
build/webgpu/conv-direct-nd.wgsl.jinja CHANGED
@@ -1,18 +1,20 @@
1
  // Direct N-dimensional convolution for channels-first tensors. Inputs are
2
  // accumulated in f32 and narrowed once at the output.
3
  // A fused epilogue may apply a residual input and activation before the store.
4
- {% if usesF16Spec %}
5
- enable f16;
6
- {% endif %}
 
 
 
 
 
7
  {{ env.wgsl.resourceDeclarations }}
8
- {% set hasActivation = hasActivation is defined and hasActivation %}
9
- {% set activation = activation | default("") %}
10
- {% set actAlpha = actAlpha | default(0.0) %}
11
- {% set actBeta = actBeta | default(0.0) %}
12
  {% if hasActivation %}
13
  // Apply the fused activation in the f32 accumulator before the single output
14
  // cast, avoiding an intermediate convolution tensor.
15
- {% macro fused_act_return(mode, alpha, beta) -%}
 
16
  {% if mode == "Relu" %}
17
  return max(v, 0.0);
18
  {% elif mode == "Clip" %}
@@ -27,24 +29,50 @@ enable f16;
27
  return tanh(clamp(v, -10.0, 10.0));
28
  {% elif mode == "HardSigmoid" %}
29
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30
  {% else %}
31
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
32
  {% endif %}
33
- {%- endmacro -%}
34
- fn fused_act(v: f32) -> f32 {
35
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
36
- }
37
  {% endif %}
38
-
39
- {% set hasZ = hasZ is defined and hasZ %}
40
  const WG: u32 = {{ convWorkgroupSize }}u;
41
 
42
  @compute @workgroup_size(WG)
43
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
44
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
45
- if (index >= params.count) {
46
- return;
47
- }
48
 
49
  let ow = index % params.outW;
50
  var q = index / params.outW;
 
1
  // Direct N-dimensional convolution for channels-first tensors. Inputs are
2
  // accumulated in f32 and narrowed once at the output.
3
  // A fused epilogue may apply a residual input and activation before the store.
4
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
5
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
6
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
7
+ // per-axis workgroup fold width.
8
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
9
+ if ({{ name }} >= {{ bound }}) {
10
+ return;
11
+ }{% endmacro %}
12
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
13
  {% if hasActivation %}
14
  // Apply the fused activation in the f32 accumulator before the single output
15
  // cast, avoiding an intermediate convolution tensor.
16
+ {% macro fused_act_fn(mode, alpha, beta) %}
17
+ fn fused_act(v: f32) -> f32 {
18
  {% if mode == "Relu" %}
19
  return max(v, 0.0);
20
  {% elif mode == "Clip" %}
 
29
  return tanh(clamp(v, -10.0, 10.0));
30
  {% elif mode == "HardSigmoid" %}
31
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
32
+ {% elif mode == "QuickGelu" %}
33
+ {% if alpha == 1 %}
34
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
35
+ return v * (1.0 / (1.0 + exp(-v)));
36
+ {% else %}
37
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
38
+ {% endif %}
39
+ {% elif mode == "Elu" %}
40
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
41
+ {% elif mode == "ThresholdedRelu" %}
42
+ return select(0.0, v, v > f32({{ alpha }}));
43
+ {% elif mode == "Softplus" %}
44
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
45
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
46
+ {% elif mode == "Erf" %}
47
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
48
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
49
+ let a = abs(v);
50
+ let t = 1.0 / (1.0 + 0.3275911 * a);
51
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
52
+ return sign(v) * (1.0 - poly * exp(-a * a));
53
+ {% elif mode == "Gelu" and alpha == 0 %}
54
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
55
+ let u = v * 0.70710678118654752;
56
+ let a = abs(u);
57
+ let t = 1.0 / (1.0 + 0.3275911 * a);
58
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
59
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
60
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
61
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
62
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
63
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
64
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
65
  {% else %}
66
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
67
  {% endif %}
68
+ }{% endmacro %}
69
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
70
  {% endif %}
 
 
71
  const WG: u32 = {{ convWorkgroupSize }}u;
72
 
73
  @compute @workgroup_size(WG)
74
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
75
+ {{ flat_index_2d("WG", "index") }}
 
 
 
76
 
77
  let ow = index % params.outW;
78
  var q = index / params.outW;
build/webgpu/conv-direct-unrolled.wgsl.jinja CHANGED
@@ -6,15 +6,20 @@
6
  // strides fold to literals; the per-row bounds check is hoisted per kh.
7
  // Group offsets are applied to both input channels and weights. f16 inputs are
8
  // widened to an f32 accumulator and narrowed once at store.
 
 
 
 
 
 
 
 
9
  {{ env.wgsl.resourceDeclarations }}
10
- {% set hasActivation = hasActivation is defined and hasActivation %}
11
- {% set activation = activation | default("") %}
12
- {% set actAlpha = actAlpha | default(0.0) %}
13
- {% set actBeta = actBeta | default(0.0) %}
14
  {% if hasActivation %}
15
  // Apply the fused activation in the f32 accumulator before the single output
16
  // cast, avoiding an intermediate convolution tensor.
17
- {% macro fused_act_return(mode, alpha, beta) -%}
 
18
  {% if mode == "Relu" %}
19
  return max(v, 0.0);
20
  {% elif mode == "Clip" %}
@@ -29,16 +34,45 @@
29
  return tanh(clamp(v, -10.0, 10.0));
30
  {% elif mode == "HardSigmoid" %}
31
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
32
  {% else %}
33
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
34
  {% endif %}
35
- {%- endmacro -%}
36
- fn fused_act(v: f32) -> f32 {
37
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
38
- }
39
  {% endif %}
40
-
41
- {% set hasZ = hasZ is defined and hasZ %}
42
  const KERNEL_AREA: u32 = {{ kernelHSpec * kernelWSpec }}u;
43
  const STRIDE_H: u32 = {{ strideHSpec }}u;
44
  const STRIDE_W: u32 = {{ strideWSpec }}u;
@@ -47,12 +81,7 @@ const PAD_LEFT: i32 = {{ padLeftSpec }};
47
 
48
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
49
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
50
- // 2D-folded flat index: gid.y carries the high bits past the
51
- // per-axis dispatch fold width (outputs > 16.7M elements).
52
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
53
- if (index >= params.count) {
54
- return;
55
- }
56
 
57
  let ow = index % params.outW;
58
  var t = index / params.outW;
 
6
  // strides fold to literals; the per-row bounds check is hoisted per kh.
7
  // Group offsets are applied to both input channels and weights. f16 inputs are
8
  // widened to an f32 accumulator and narrowed once at store.
9
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
10
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
11
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
12
+ // per-axis workgroup fold width.
13
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
14
+ if ({{ name }} >= {{ bound }}) {
15
+ return;
16
+ }{% endmacro %}
17
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
18
  {% if hasActivation %}
19
  // Apply the fused activation in the f32 accumulator before the single output
20
  // cast, avoiding an intermediate convolution tensor.
21
+ {% macro fused_act_fn(mode, alpha, beta) %}
22
+ fn fused_act(v: f32) -> f32 {
23
  {% if mode == "Relu" %}
24
  return max(v, 0.0);
25
  {% elif mode == "Clip" %}
 
34
  return tanh(clamp(v, -10.0, 10.0));
35
  {% elif mode == "HardSigmoid" %}
36
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
37
+ {% elif mode == "QuickGelu" %}
38
+ {% if alpha == 1 %}
39
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
40
+ return v * (1.0 / (1.0 + exp(-v)));
41
+ {% else %}
42
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
43
+ {% endif %}
44
+ {% elif mode == "Elu" %}
45
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
46
+ {% elif mode == "ThresholdedRelu" %}
47
+ return select(0.0, v, v > f32({{ alpha }}));
48
+ {% elif mode == "Softplus" %}
49
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
50
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
51
+ {% elif mode == "Erf" %}
52
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
53
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
54
+ let a = abs(v);
55
+ let t = 1.0 / (1.0 + 0.3275911 * a);
56
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
57
+ return sign(v) * (1.0 - poly * exp(-a * a));
58
+ {% elif mode == "Gelu" and alpha == 0 %}
59
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
60
+ let u = v * 0.70710678118654752;
61
+ let a = abs(u);
62
+ let t = 1.0 / (1.0 + 0.3275911 * a);
63
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
64
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
65
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
66
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
67
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
68
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
69
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
70
  {% else %}
71
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
72
  {% endif %}
73
+ }{% endmacro %}
74
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
75
  {% endif %}
 
 
76
  const KERNEL_AREA: u32 = {{ kernelHSpec * kernelWSpec }}u;
77
  const STRIDE_H: u32 = {{ strideHSpec }}u;
78
  const STRIDE_W: u32 = {{ strideWSpec }}u;
 
81
 
82
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
83
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
84
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "index") }}
 
 
 
 
 
85
 
86
  let ow = index % params.outW;
87
  var t = index / params.outW;
build/webgpu/conv-splitk-reduce.wgsl.jinja CHANGED
@@ -10,17 +10,22 @@
10
  //
11
  // Summing slices in index order changes the f32 association relative to one
12
  // uninterrupted K loop, so the two orders need not be bit-identical.
 
 
 
 
 
 
 
 
13
  {{ env.wgsl.resourceDeclarations }}
14
  {% set applyActivation = hasActivation is defined and hasActivation %}
15
  {% if applyActivation %}
16
- {% set hasActivation = hasActivation is defined and hasActivation %}
17
- {% set activation = activation | default("") %}
18
- {% set actAlpha = actAlpha | default(0.0) %}
19
- {% set actBeta = actBeta | default(0.0) %}
20
  {% if hasActivation %}
21
  // Apply the fused activation in the f32 accumulator before the single output
22
  // cast, avoiding an intermediate convolution tensor.
23
- {% macro fused_act_return(mode, alpha, beta) -%}
 
24
  {% if mode == "Relu" %}
25
  return max(v, 0.0);
26
  {% elif mode == "Clip" %}
@@ -35,15 +40,45 @@
35
  return tanh(clamp(v, -10.0, 10.0));
36
  {% elif mode == "HardSigmoid" %}
37
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
38
  {% else %}
39
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
40
  {% endif %}
41
- {%- endmacro -%}
42
- fn fused_act(v: f32) -> f32 {
43
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
44
- }
45
  {% endif %}
46
-
47
  {% endif %}
48
 
49
  const M: u32 = {{ M }}u;
@@ -57,12 +92,7 @@ const WORKGROUP_SIZE: u32 = {{ reduceWorkgroupSize }}u;
57
 
58
  @compute @workgroup_size({{ reduceWorkgroupSize }}, 1, 1)
59
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
60
- // 2D-folded flat index: gid.y carries the high bits past
61
- // the per-axis dispatch fold width and reduces to the 1D form at y=0.
62
- let idx = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WORKGROUP_SIZE;
63
- if (idx >= COUNT) {
64
- return;
65
- }
66
  let col = idx % N;
67
  let row = (idx / N) % M;
68
  let batch = idx / (M * N);
 
10
  //
11
  // Summing slices in index order changes the f32 association relative to one
12
  // uninterrupted K loop, so the two orders need not be bit-identical.
13
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
14
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
15
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
16
+ // per-axis workgroup fold width.
17
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
18
+ if ({{ name }} >= {{ bound }}) {
19
+ return;
20
+ }{% endmacro %}
21
  {{ env.wgsl.resourceDeclarations }}
22
  {% set applyActivation = hasActivation is defined and hasActivation %}
23
  {% if applyActivation %}
 
 
 
 
24
  {% if hasActivation %}
25
  // Apply the fused activation in the f32 accumulator before the single output
26
  // cast, avoiding an intermediate convolution tensor.
27
+ {% macro fused_act_fn(mode, alpha, beta) %}
28
+ fn fused_act(v: f32) -> f32 {
29
  {% if mode == "Relu" %}
30
  return max(v, 0.0);
31
  {% elif mode == "Clip" %}
 
40
  return tanh(clamp(v, -10.0, 10.0));
41
  {% elif mode == "HardSigmoid" %}
42
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
43
+ {% elif mode == "QuickGelu" %}
44
+ {% if alpha == 1 %}
45
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
46
+ return v * (1.0 / (1.0 + exp(-v)));
47
+ {% else %}
48
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
49
+ {% endif %}
50
+ {% elif mode == "Elu" %}
51
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
52
+ {% elif mode == "ThresholdedRelu" %}
53
+ return select(0.0, v, v > f32({{ alpha }}));
54
+ {% elif mode == "Softplus" %}
55
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
56
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
57
+ {% elif mode == "Erf" %}
58
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
59
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
60
+ let a = abs(v);
61
+ let t = 1.0 / (1.0 + 0.3275911 * a);
62
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
63
+ return sign(v) * (1.0 - poly * exp(-a * a));
64
+ {% elif mode == "Gelu" and alpha == 0 %}
65
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
66
+ let u = v * 0.70710678118654752;
67
+ let a = abs(u);
68
+ let t = 1.0 / (1.0 + 0.3275911 * a);
69
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
70
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
71
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
72
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
73
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
74
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
75
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
76
  {% else %}
77
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
78
  {% endif %}
79
+ }{% endmacro %}
80
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
81
  {% endif %}
 
82
  {% endif %}
83
 
84
  const M: u32 = {{ M }}u;
 
92
 
93
  @compute @workgroup_size({{ reduceWorkgroupSize }}, 1, 1)
94
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
95
+ {{ flat_index_2d("WORKGROUP_SIZE", "idx", "COUNT") }}
 
 
 
 
 
96
  let col = idx % N;
97
  let row = (idx / N) % M;
98
  let batch = idx / (M * N);
build/webgpu/conv1d-tiled-reg.wgsl.jinja CHANGED
@@ -10,14 +10,11 @@
10
  // through dot(). Bounds checks cover every partial tile. Accumulation, bias,
11
  // and the optional activation epilogue use f32; f16 operands widen at the FMA.
12
  {{ env.wgsl.resourceDeclarations }}
13
- {% set hasActivation = hasActivation is defined and hasActivation %}
14
- {% set activation = activation | default("") %}
15
- {% set actAlpha = actAlpha | default(0.0) %}
16
- {% set actBeta = actBeta | default(0.0) %}
17
  {% if hasActivation %}
18
  // Apply the fused activation in the f32 accumulator before the single output
19
  // cast, avoiding an intermediate convolution tensor.
20
- {% macro fused_act_return(mode, alpha, beta) -%}
 
21
  {% if mode == "Relu" %}
22
  return max(v, 0.0);
23
  {% elif mode == "Clip" %}
@@ -32,15 +29,45 @@
32
  return tanh(clamp(v, -10.0, 10.0));
33
  {% elif mode == "HardSigmoid" %}
34
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
35
  {% else %}
36
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
37
  {% endif %}
38
- {%- endmacro -%}
39
- fn fused_act(v: f32) -> f32 {
40
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
41
- }
42
  {% endif %}
43
-
44
  {% set tileT = "f16" if usesF16 else "f32" %}
45
 
46
  const WG_X: u32 = {{ convWgX }}u;
 
10
  // through dot(). Bounds checks cover every partial tile. Accumulation, bias,
11
  // and the optional activation epilogue use f32; f16 operands widen at the FMA.
12
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
13
  {% if hasActivation %}
14
  // Apply the fused activation in the f32 accumulator before the single output
15
  // cast, avoiding an intermediate convolution tensor.
16
+ {% macro fused_act_fn(mode, alpha, beta) %}
17
+ fn fused_act(v: f32) -> f32 {
18
  {% if mode == "Relu" %}
19
  return max(v, 0.0);
20
  {% elif mode == "Clip" %}
 
29
  return tanh(clamp(v, -10.0, 10.0));
30
  {% elif mode == "HardSigmoid" %}
31
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
32
+ {% elif mode == "QuickGelu" %}
33
+ {% if alpha == 1 %}
34
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
35
+ return v * (1.0 / (1.0 + exp(-v)));
36
+ {% else %}
37
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
38
+ {% endif %}
39
+ {% elif mode == "Elu" %}
40
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
41
+ {% elif mode == "ThresholdedRelu" %}
42
+ return select(0.0, v, v > f32({{ alpha }}));
43
+ {% elif mode == "Softplus" %}
44
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
45
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
46
+ {% elif mode == "Erf" %}
47
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
48
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
49
+ let a = abs(v);
50
+ let t = 1.0 / (1.0 + 0.3275911 * a);
51
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
52
+ return sign(v) * (1.0 - poly * exp(-a * a));
53
+ {% elif mode == "Gelu" and alpha == 0 %}
54
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
55
+ let u = v * 0.70710678118654752;
56
+ let a = abs(u);
57
+ let t = 1.0 / (1.0 + 0.3275911 * a);
58
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
59
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
60
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
61
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
62
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
63
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
64
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
65
  {% else %}
66
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
67
  {% endif %}
68
+ }{% endmacro %}
69
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
70
  {% endif %}
 
71
  {% set tileT = "f16" if usesF16 else "f32" %}
72
 
73
  const WG_X: u32 = {{ convWgX }}u;
build/webgpu/conv2d-grouped-large-w4.wgsl.jinja CHANGED
@@ -1,12 +1,15 @@
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
- {% set hasActivation = hasActivation is defined and hasActivation %}
3
- {% set activation = activation | default("") %}
4
- {% set actAlpha = actAlpha | default(0.0) %}
5
- {% set actBeta = actBeta | default(0.0) %}
6
  {% if hasActivation %}
7
  // Apply the fused activation in the f32 accumulator before the single output
8
  // cast, avoiding an intermediate convolution tensor.
9
- {% macro fused_act_return(mode, alpha, beta) -%}
 
10
  {% if mode == "Relu" %}
11
  return max(v, 0.0);
12
  {% elif mode == "Clip" %}
@@ -21,23 +24,50 @@
21
  return tanh(clamp(v, -10.0, 10.0));
22
  {% elif mode == "HardSigmoid" %}
23
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  {% else %}
25
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
26
  {% endif %}
27
- {%- endmacro -%}
28
- fn fused_act(v: f32) -> f32 {
29
- {{ fused_act_return(activation, actAlpha, actBeta) -}}
30
- }
31
  {% endif %}
32
-
33
  // Large grouped kernels are dominated by repeatedly fetching the same filter
34
  // value for neighboring output columns and input window for neighboring output
35
  // channels. One invocation computes four adjacent columns for OC_TILE channels,
36
  // sharing each input load across those channels. Shape and window geometry are
37
  // static, so indexing divisors and kernel offsets are shader constants.
38
- {% set dilatedLanes = dilatedLanes | default(false) %}
39
- {% set lanes = lanes | default(4) %}
40
- {% set quads = quads | default(1) %}
41
  {% set laneStep = dilationWSpec if dilatedLanes else strideWSpec %}
42
  {% set comps = ["x", "y", "z", "w"] %}
43
  {% macro acc_ref(oct, lane) %}{% if dilatedLanes %}acc{{ oct }}_{{ lane }}{% else %}acc{{ oct }}.{{ comps[lane] }}{% endif %}{% endmacro %}
@@ -86,8 +116,7 @@ const WG: u32 = {{ workgroupSizeSpec }}u;
86
 
87
  @compute @workgroup_size({{ workgroupSizeSpec }})
88
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
89
- let q = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
90
- if (q >= COUNT_TILES) { return; }
91
 
92
  {% if dilatedLanes %}
93
  let tile = q % TILES_PER_ROW;
@@ -113,19 +142,25 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
113
  var xChannelBase = (batch * IN_C + group * IN_CPG) * IN_PLANE;
114
  {% for oct in range(ocTile) %}
115
  var wBase{{ oct }} = (ocBase + {{ oct }}u) * IN_CPG * KAREA;
 
 
 
 
 
116
  {% if dilatedLanes %}
117
  {% for lane in range(lanes) %}
118
- var acc{{ oct }}_{{ lane }} = 0.0;
119
  {% endfor %}
120
  {% else %}
121
- var acc{{ oct }} = vec4<f32>(0.0);
122
  {% endif %}
123
  {% endfor %}
124
 
125
  for (var ic = 0u; ic < IN_CPG; ic++) {
126
- {% for kh in range(kernelHSpec) %}
 
127
  {
128
- let ih = ihBase + {{ kh * dilationHSpec }};
129
  if (ih >= 0 && ih < IN_H) {
130
  let xRow = xChannelBase + u32(ih) * IN_W_U;
131
  {% if registerForm %}
@@ -165,7 +200,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
165
  {% for kw in range(kernelWSpec) %}
166
  {
167
  {% for oct in range(ocTile) %}
168
- let weight{{ oct }} = f32(w[wBase{{ oct }} + {{ kh * kernelWSpec + kw }}u]);
169
  {% endfor %}
170
  {% for lane in range(lanes) %}
171
  {% for oct in range(ocTile) %}
@@ -178,7 +213,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
178
  {% for kw in range(kernelWSpec) %}
179
  {
180
  {% for oct in range(ocTile) %}
181
- let weight{{ oct }} = f32(w[wBase{{ oct }} + {{ kh * kernelWSpec + kw }}u]);
182
  {% endfor %}
183
  let iwK = iwBase + {{ kw * dilationWSpec }};
184
  {% for lane in range(lanes) %}
@@ -196,7 +231,8 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
196
  }
197
  }
198
  {% endfor %}
199
- xChannelBase += IN_PLANE;
 
200
  {% for oct in range(ocTile) %}
201
  wBase{{ oct }} += KAREA;
202
  {% endfor %}
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
6
+ if ({{ name }} >= {{ bound }}) { return; }{% endmacro %}
7
  {{ env.wgsl.resourceDeclarations }}
 
 
 
 
8
  {% if hasActivation %}
9
  // Apply the fused activation in the f32 accumulator before the single output
10
  // cast, avoiding an intermediate convolution tensor.
11
+ {% macro fused_act_fn(mode, alpha, beta) %}
12
+ fn fused_act(v: f32) -> f32 {
13
  {% if mode == "Relu" %}
14
  return max(v, 0.0);
15
  {% elif mode == "Clip" %}
 
24
  return tanh(clamp(v, -10.0, 10.0));
25
  {% elif mode == "HardSigmoid" %}
26
  return clamp(f32({{ alpha }}) * v + f32({{ beta }}), 0.0, 1.0);
27
+ {% elif mode == "QuickGelu" %}
28
+ {% if alpha == 1 %}
29
+ // alpha == 1 is SiLU: the multiply folds away (the spelling yolov8-class detectors emit).
30
+ return v * (1.0 / (1.0 + exp(-v)));
31
+ {% else %}
32
+ return v * (1.0 / (1.0 + exp(-(f32({{ alpha }}) * v))));
33
+ {% endif %}
34
+ {% elif mode == "Elu" %}
35
+ return select(f32({{ alpha }}) * (exp(v) - 1.0), v, v >= 0.0);
36
+ {% elif mode == "ThresholdedRelu" %}
37
+ return select(0.0, v, v > f32({{ alpha }}));
38
+ {% elif mode == "Softplus" %}
39
+ // log(1 + exp(v)) spelled overflow-safe: max(v, 0) + log(1 + exp(-|v|)).
40
+ return max(v, 0.0) + log(1.0 + exp(-abs(v)));
41
+ {% elif mode == "Erf" %}
42
+ // Abramowitz & Stegun 7.1.26 (|error| < 1.5e-7), the same polynomial the standalone
43
+ // Erf/Gelu kernels use, so a fused and an unfused graph agree to f32 rounding.
44
+ let a = abs(v);
45
+ let t = 1.0 / (1.0 + 0.3275911 * a);
46
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
47
+ return sign(v) * (1.0 - poly * exp(-a * a));
48
+ {% elif mode == "Gelu" and alpha == 0 %}
49
+ // ONNX Gelu, erf form: 0.5 * v * (1 + erf(v / sqrt(2))); activation_params[0] == 0.
50
+ let u = v * 0.70710678118654752;
51
+ let a = abs(u);
52
+ let t = 1.0 / (1.0 + 0.3275911 * a);
53
+ let poly = ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - 0.284496736) * t + 0.254829592) * t;
54
+ return 0.5 * v * (1.0 + sign(u) * (1.0 - poly * exp(-a * a)));
55
+ {% elif mode == "Gelu" or mode == "FastGelu" %}
56
+ // tanh approximation (ONNX Gelu approximate="tanh" arrives as activation_params[0] != 0;
57
+ // com.microsoft.FastGelu always). tanh saturates to +/-1 well inside the clamp.
58
+ let inner = v * (0.035677408136300125 * v * v + 0.7978845608028654);
59
+ return v * (0.5 + 0.5 * tanh(clamp(inner, -10.0, 10.0)));
60
  {% else %}
61
  return v * clamp(v * 0.16666666666666666 + 0.5, 0.0, 1.0);
62
  {% endif %}
63
+ }{% endmacro %}
64
+ {{ fused_act_fn(activation, actAlpha, actBeta) }}
 
 
65
  {% endif %}
 
66
  // Large grouped kernels are dominated by repeatedly fetching the same filter
67
  // value for neighboring output columns and input window for neighboring output
68
  // channels. One invocation computes four adjacent columns for OC_TILE channels,
69
  // sharing each input load across those channels. Shape and window geometry are
70
  // static, so indexing divisors and kernel offsets are shader constants.
 
 
 
71
  {% set laneStep = dilationWSpec if dilatedLanes else strideWSpec %}
72
  {% set comps = ["x", "y", "z", "w"] %}
73
  {% macro acc_ref(oct, lane) %}{% if dilatedLanes %}acc{{ oct }}_{{ lane }}{% else %}acc{{ oct }}.{{ comps[lane] }}{% endif %}{% endmacro %}
 
116
 
117
  @compute @workgroup_size({{ workgroupSizeSpec }})
118
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
119
+ {{ flat_index_2d("WG", "q", "COUNT_TILES", guardInline=true) }}
 
120
 
121
  {% if dilatedLanes %}
122
  let tile = q % TILES_PER_ROW;
 
142
  var xChannelBase = (batch * IN_C + group * IN_CPG) * IN_PLANE;
143
  {% for oct in range(ocTile) %}
144
  var wBase{{ oct }} = (ocBase + {{ oct }}u) * IN_CPG * KAREA;
145
+ {% if hasBias %}
146
+ // Seeding the accumulator with the channel bias replaces the zero seed and
147
+ // incorporates the bias before the reduction.
148
+ let bias{{ oct }} = f32(bias[ocBase + {{ oct }}u]);
149
+ {% endif %}
150
  {% if dilatedLanes %}
151
  {% for lane in range(lanes) %}
152
+ var acc{{ oct }}_{{ lane }} = {% if hasBias %}bias{{ oct }}{% else %}0.0{% endif %};
153
  {% endfor %}
154
  {% else %}
155
+ var acc{{ oct }} = vec4<f32>({% if hasBias %}bias{{ oct }}{% else %}0.0{% endif %});
156
  {% endif %}
157
  {% endfor %}
158
 
159
  for (var ic = 0u; ic < IN_CPG; ic++) {
160
+ {% if rollKernelRows %} for (var kernelRow = 0u; kernelRow < {{ kernelHSpec }}u; kernelRow++) {
161
+ {% endif %}{% for kh in range(1 if rollKernelRows else kernelHSpec) %}
162
  {
163
+ let ih = ihBase + {% if rollKernelRows %}i32(kernelRow) * {{ dilationHSpec }}{% else %}{{ kh * dilationHSpec }}{% endif %};
164
  if (ih >= 0 && ih < IN_H) {
165
  let xRow = xChannelBase + u32(ih) * IN_W_U;
166
  {% if registerForm %}
 
200
  {% for kw in range(kernelWSpec) %}
201
  {
202
  {% for oct in range(ocTile) %}
203
+ let weight{{ oct }} = f32(w[wBase{{ oct }} + {% if rollKernelRows %}kernelRow * {{ kernelWSpec }}u + {{ kw }}u{% else %}{{ kh * kernelWSpec + kw }}u{% endif %}]);
204
  {% endfor %}
205
  {% for lane in range(lanes) %}
206
  {% for oct in range(ocTile) %}
 
213
  {% for kw in range(kernelWSpec) %}
214
  {
215
  {% for oct in range(ocTile) %}
216
+ let weight{{ oct }} = f32(w[wBase{{ oct }} + {% if rollKernelRows %}kernelRow * {{ kernelWSpec }}u + {{ kw }}u{% else %}{{ kh * kernelWSpec + kw }}u{% endif %}]);
217
  {% endfor %}
218
  let iwK = iwBase + {{ kw * dilationWSpec }};
219
  {% for lane in range(lanes) %}
 
231
  }
232
  }
233
  {% endfor %}
234
+ {% if rollKernelRows %} }
235
+ {% endif %} xChannelBase += IN_PLANE;
236
  {% for oct in range(ocTile) %}
237
  wBase{{ oct }} += KAREA;
238
  {% endfor %}
build/webgpu/manifest.json CHANGED
The diff for this file is too large to render. See raw diff
 
build/webgpu/metadata.json CHANGED
@@ -1,29 +1,29 @@
1
  {
2
  "name": "com.microsoft.FusedConv",
3
- "id": "_com_microsoft_fusedconv_webgpu_e4dcf04",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "rDPb1+xln8/qfupw2ajiwaBUEKJvlo+NuJEs81cAgTk=",
11
- "conv-1x1-gemm-tiled-reg.wgsl.jinja": "9yFpwXS1HzmWmm4W72n+aOIryLqc9+e88awZ1ogTe+Q=",
12
- "conv-1x1-gemm-tiled.wgsl.jinja": "uqG5mRVvcHiUo0LDCUcVgqElVTvRNknx19Gm4oIALiU=",
13
- "conv-1x1-subgroup-matrix.wgsl.jinja": "iDJs4RJk4iJarOZCjY7/5MFrHgR7CcqI6x5uwv2dYaA=",
14
- "conv-direct-nd.wgsl.jinja": "5UA6BigLbtPRMre3LRfNShldvupB5yTGn8EzX2tBzP8=",
15
- "conv-direct-unrolled.wgsl.jinja": "vrGUPwYkbYE1slII7nuK5wbERRcG5u0X1cW1kMc2sQs=",
16
  "conv-im2col-nchw.wgsl.jinja": "IRpmtpeVQsLE6ud+Sf1230XaZVS1dxDtTZuSDwU3P8s=",
17
- "conv-splitk-reduce.wgsl.jinja": "LbVU5DRL0Q51fvLjThmxJQZH8aPbZkvNzij4nY/cHmo=",
18
- "conv1d-tiled-reg.wgsl.jinja": "ERqOPl+1RbnJ+CFqZvLSYw45qYMKgB7Je0teew5gyOM=",
19
- "conv2d-grouped-large-w4.wgsl.jinja": "qBG2cfOsQ4FMwtNT8wBOqLJFDOs6w0wJCmfcZuxbho8=",
20
- "manifest.json": "EpZ3lQZ0Ar7ZY4xDkUlcQ3m6lyw6x6+cp9KoBC47TSE=",
21
- "test.json": "g7doPXd+oCU97EVAbYReA8qInOswdv5wrLcMYCX7CcM="
22
  }
23
  },
24
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
25
  "webgpu": {
26
- "manifestSpec": "2.0",
27
  "variants": {
28
  "implicit_im2col_tiled_reg_splitk": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-splitk-reduce.wgsl.jinja"],
29
  "implicit_im2col_tiled_bias_reg_splitk": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-splitk-reduce.wgsl.jinja"],
@@ -57,6 +57,8 @@
57
  "im2col_gemm_tiled_bias_z": ["conv-1x1-gemm-tiled.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
58
  "im2col_gemm_tiled_reg": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
59
  "im2col_gemm_tiled_bias_reg": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
 
 
60
  "gemm_1x1_tiled": ["conv-1x1-gemm-tiled.wgsl.jinja"],
61
  "gemm_1x1_tiled_z": ["conv-1x1-gemm-tiled.wgsl.jinja"],
62
  "gemm_1x1_tiled_bias": ["conv-1x1-gemm-tiled.wgsl.jinja"],
@@ -76,8 +78,11 @@
76
  "ncdhw3d_bias": ["conv-direct-nd.wgsl.jinja"],
77
  "ncdhw3d_bias_z": ["conv-direct-nd.wgsl.jinja"],
78
  "grouped_large_kernel_w4": ["conv2d-grouped-large-w4.wgsl.jinja"],
 
79
  "grouped_large_kernel_w4_tail": ["conv2d-grouped-large-w4.wgsl.jinja"],
 
80
  "grouped_large_kernel_w4_dilated_lanes": ["conv2d-grouped-large-w4.wgsl.jinja"],
 
81
  "implicit_im2col_subgroup_matrix": ["conv-1x1-subgroup-matrix.wgsl.jinja"],
82
  "implicit_im2col_subgroup_matrix_z": ["conv-1x1-subgroup-matrix.wgsl.jinja"],
83
  "implicit_im2col_subgroup_matrix_bias": ["conv-1x1-subgroup-matrix.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.FusedConv",
3
+ "id": "_com_microsoft_fusedconv_webgpu_a0813e7",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "h7j7tWEEOwdOjHUeFKfE46+kvmVgwWNFJTEcVB6dJK8=",
11
+ "conv-1x1-gemm-tiled-reg.wgsl.jinja": "J6MUpDFROmGCbCEln+xq1imr5+In47OPstc5DBcqv8U=",
12
+ "conv-1x1-gemm-tiled.wgsl.jinja": "eFyqUTNykx9plj5o3Se9dEHR4W3H/lFjrTgmQ7N/81A=",
13
+ "conv-1x1-subgroup-matrix.wgsl.jinja": "1f8hTKEnvhrxDuYMBFnXRjDoxKxr9u4ToeanV/Hjk5s=",
14
+ "conv-direct-nd.wgsl.jinja": "7CVIjupPj0kd1lioOiY4PZY2JLcHYiMZ2CnykIoPnpQ=",
15
+ "conv-direct-unrolled.wgsl.jinja": "efRWnTZ8xYUWP+Y8JpGN7w5r7log+H+qTZDTkIyH6gE=",
16
  "conv-im2col-nchw.wgsl.jinja": "IRpmtpeVQsLE6ud+Sf1230XaZVS1dxDtTZuSDwU3P8s=",
17
+ "conv-splitk-reduce.wgsl.jinja": "rxt8ttb0FFw6RpO6hfDbpneMPt5AqD1A7LwA1kiBG+w=",
18
+ "conv1d-tiled-reg.wgsl.jinja": "+SsmtbIw2VkES8J32j1SJ6tGZgHon0Z1ZySA8Ga1iPk=",
19
+ "conv2d-grouped-large-w4.wgsl.jinja": "2KdxhCkpdSFYVRX4LIRnTLHBYW0oCx6MbQmzvauz61Y=",
20
+ "manifest.json": "Kwp+pHg8GOMiDPWE3CeSMu6UAzi+RAOtEA4wninC8Yk=",
21
+ "test.json": "YY8eK1tSLM+NV1kDm5q+BugcjSDAeOS0mpny06NELas="
22
  }
23
  },
24
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
25
  "webgpu": {
26
+ "manifestSpec": "2.1",
27
  "variants": {
28
  "implicit_im2col_tiled_reg_splitk": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-splitk-reduce.wgsl.jinja"],
29
  "implicit_im2col_tiled_bias_reg_splitk": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-splitk-reduce.wgsl.jinja"],
 
57
  "im2col_gemm_tiled_bias_z": ["conv-1x1-gemm-tiled.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
58
  "im2col_gemm_tiled_reg": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
59
  "im2col_gemm_tiled_bias_reg": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
60
+ "im2col_gemm_tiled_reg_f16_columns": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
61
+ "im2col_gemm_tiled_bias_reg_f16_columns": ["conv-1x1-gemm-tiled-reg.wgsl.jinja", "conv-im2col-nchw.wgsl.jinja"],
62
  "gemm_1x1_tiled": ["conv-1x1-gemm-tiled.wgsl.jinja"],
63
  "gemm_1x1_tiled_z": ["conv-1x1-gemm-tiled.wgsl.jinja"],
64
  "gemm_1x1_tiled_bias": ["conv-1x1-gemm-tiled.wgsl.jinja"],
 
78
  "ncdhw3d_bias": ["conv-direct-nd.wgsl.jinja"],
79
  "ncdhw3d_bias_z": ["conv-direct-nd.wgsl.jinja"],
80
  "grouped_large_kernel_w4": ["conv2d-grouped-large-w4.wgsl.jinja"],
81
+ "grouped_large_kernel_w4_bias": ["conv2d-grouped-large-w4.wgsl.jinja"],
82
  "grouped_large_kernel_w4_tail": ["conv2d-grouped-large-w4.wgsl.jinja"],
83
+ "grouped_large_kernel_w4_tail_bias": ["conv2d-grouped-large-w4.wgsl.jinja"],
84
  "grouped_large_kernel_w4_dilated_lanes": ["conv2d-grouped-large-w4.wgsl.jinja"],
85
+ "grouped_large_kernel_w4_dilated_lanes_bias": ["conv2d-grouped-large-w4.wgsl.jinja"],
86
  "implicit_im2col_subgroup_matrix": ["conv-1x1-subgroup-matrix.wgsl.jinja"],
87
  "implicit_im2col_subgroup_matrix_z": ["conv-1x1-subgroup-matrix.wgsl.jinja"],
88
  "implicit_im2col_subgroup_matrix_bias": ["conv-1x1-subgroup-matrix.wgsl.jinja"],
build/webgpu/test.json CHANGED
The diff for this file is too large to render. See raw diff