sync 6fdf6301e2bb
Browse files- README.md +14 -4
- build/webgpu/bench.json +0 -0
- build/webgpu/conv-1x1-gemm-tiled-reg.wgsl.jinja +46 -125
- build/webgpu/conv-1x1-gemm-tiled.wgsl.jinja +59 -21
- build/webgpu/conv-1x1-subgroup-matrix.wgsl.jinja +50 -179
- build/webgpu/conv-direct-nd.wgsl.jinja +46 -18
- build/webgpu/conv-direct-unrolled.wgsl.jinja +46 -17
- build/webgpu/conv-splitk-reduce.wgsl.jinja +46 -16
- build/webgpu/conv1d-tiled-reg.wgsl.jinja +37 -10
- build/webgpu/conv2d-grouped-large-w4.wgsl.jinja +58 -22
- build/webgpu/manifest.json +0 -0
- build/webgpu/metadata.json +19 -14
- build/webgpu/test.json +0 -0
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 `
|
| 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
|
| 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.
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 43 |
-
|
| 44 |
-
{{ fused_act_return(activation, actAlpha, actBeta) -}}
|
| 45 |
-
}
|
| 46 |
{% endif %}
|
| 47 |
-
|
| 48 |
-
{% set
|
| 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 =
|
| 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 %}
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 35 |
-
|
| 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(
|
| 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 *
|
| 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) /
|
| 92 |
-
let idx = li + e *
|
| 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 |
-
|
| 119 |
-
|
| 120 |
-
|
| 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 =
|
| 16 |
-
{% set polyphase =
|
| 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 =
|
| 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 =
|
| 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 =
|
| 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 = "
|
| 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
|
| 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 |
-
{%
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 356 |
-
|
| 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 %}
|
| 378 |
-
{%
|
| 379 |
-
{%
|
| 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
|
| 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 |
-
{%
|
| 5 |
-
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 34 |
-
|
| 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 |
-
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 36 |
-
|
| 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 |
-
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 42 |
-
|
| 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 |
-
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 39 |
-
|
| 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
|
|
|
|
| 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 |
-
{%
|
| 28 |
-
|
| 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 |
-
|
| 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 |
-
{%
|
|
|
|
| 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 |
-
|
|
|
|
| 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": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
-
"bench.json": "
|
| 11 |
-
"conv-1x1-gemm-tiled-reg.wgsl.jinja": "
|
| 12 |
-
"conv-1x1-gemm-tiled.wgsl.jinja": "
|
| 13 |
-
"conv-1x1-subgroup-matrix.wgsl.jinja": "
|
| 14 |
-
"conv-direct-nd.wgsl.jinja": "
|
| 15 |
-
"conv-direct-unrolled.wgsl.jinja": "
|
| 16 |
"conv-im2col-nchw.wgsl.jinja": "IRpmtpeVQsLE6ud+Sf1230XaZVS1dxDtTZuSDwU3P8s=",
|
| 17 |
-
"conv-splitk-reduce.wgsl.jinja": "
|
| 18 |
-
"conv1d-tiled-reg.wgsl.jinja": "
|
| 19 |
-
"conv2d-grouped-large-w4.wgsl.jinja": "
|
| 20 |
-
"manifest.json": "
|
| 21 |
-
"test.json": "
|
| 22 |
}
|
| 23 |
},
|
| 24 |
-
"provenance": { "kernel": { "sha": "
|
| 25 |
"webgpu": {
|
| 26 |
-
"manifestSpec": "2.
|
| 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
|
|
|