sync 6fdf6301e2bb
Browse files- README.md +2 -2
- build/webgpu/batch-normalization-nc-vec4.wgsl.jinja +5 -30
- build/webgpu/batch-normalization-nchw-flat-vec4.wgsl.jinja +5 -30
- build/webgpu/batch-normalization-nchw-vec4.wgsl.jinja +5 -30
- build/webgpu/batch-normalization-nchw.wgsl.jinja +5 -30
- build/webgpu/manifest.json +16 -24
- build/webgpu/metadata.json +8 -8
README.md
CHANGED
|
@@ -53,7 +53,7 @@ Default values (overridable per request):
|
|
| 53 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 54 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 55 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 56 |
-
- [`bench.json`](build/webgpu/bench.json) — benchmark
|
| 57 |
- [`batch-normalization-nc-vec4.wgsl.jinja`](build/webgpu/batch-normalization-nc-vec4.wgsl.jinja)
|
| 58 |
- [`batch-normalization-nchw-flat-vec4.wgsl.jinja`](build/webgpu/batch-normalization-nchw-flat-vec4.wgsl.jinja)
|
| 59 |
- [`batch-normalization-nchw-vec4.wgsl.jinja`](build/webgpu/batch-normalization-nchw-vec4.wgsl.jinja)
|
|
@@ -62,7 +62,7 @@ Default values (overridable per request):
|
|
| 62 |
## Use with `@huggingface/kernels`
|
| 63 |
|
| 64 |
```sh
|
| 65 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 66 |
```
|
| 67 |
|
| 68 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 53 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 54 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 55 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 56 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 57 |
- [`batch-normalization-nc-vec4.wgsl.jinja`](build/webgpu/batch-normalization-nc-vec4.wgsl.jinja)
|
| 58 |
- [`batch-normalization-nchw-flat-vec4.wgsl.jinja`](build/webgpu/batch-normalization-nchw-flat-vec4.wgsl.jinja)
|
| 59 |
- [`batch-normalization-nchw-vec4.wgsl.jinja`](build/webgpu/batch-normalization-nchw-vec4.wgsl.jinja)
|
|
|
|
| 62 |
## Use with `@huggingface/kernels`
|
| 63 |
|
| 64 |
```sh
|
| 65 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 66 |
```
|
| 67 |
|
| 68 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
build/webgpu/batch-normalization-nc-vec4.wgsl.jinja
CHANGED
|
@@ -1,43 +1,18 @@
|
|
| 1 |
-
{% macro flat_index_2d(name="i", bound="params.count", guardInline=false
|
| 2 |
-
{%
|
| 3 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 4 |
-
// per-axis workgroup fold width (outputs > 16.7M elements).
|
| 5 |
-
{% elif note == "limit" %}
|
| 6 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 7 |
-
// per-axis workgroup fold width.
|
| 8 |
-
{% elif note == "device-axis" %}
|
| 9 |
-
// The flat dispatch is folded across x/y at a fixed per-axis workgroup
|
| 10 |
-
// width; gid.y carries the high portion of the output index.
|
| 11 |
-
{% elif note == "vec4-limit" %}
|
| 12 |
-
// 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
|
| 13 |
-
// per-axis workgroup fold width (the dispatch caps x and spills into y).
|
| 14 |
-
{% elif note == "element-limit" %}
|
| 15 |
-
// 2D-folded flat element index: gid.y carries the high bits past the
|
| 16 |
-
// dispatch's per-axis workgroup fold width.
|
| 17 |
-
{% elif note == "dispatch" %}
|
| 18 |
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 19 |
// per-axis workgroup fold width.
|
| 20 |
-
{
|
| 21 |
-
{% if bound == "" %}
|
| 22 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 23 |
-
{%- elif guardInline %}
|
| 24 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 25 |
-
if ({{ name }} >= {{ bound }}) { return; }
|
| 26 |
-
{%- else %}
|
| 27 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 28 |
if ({{ name }} >= {{ bound }}) {
|
| 29 |
return;
|
| 30 |
-
}
|
| 31 |
-
{%- endif %}
|
| 32 |
-
{% endmacro %}
|
| 33 |
-
|
| 34 |
{{ env.wgsl.resourceDeclarations }}
|
| 35 |
|
| 36 |
// Rank-2 [N, C] inference vec4 specialization. Vectors run across adjacent channels, so
|
| 37 |
// scale/bias/mean/var are also bound as vec4<f32> and C must be divisible by 4.
|
| 38 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 39 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 40 |
-
{{ flat_index_2d("i", "params.count4") }}
|
| 41 |
let channel4 = i % params.channels4;
|
| 42 |
let alpha = scale[channel4] * inverseSqrt(input_var[channel4] + vec4<f32>(params.epsilon));
|
| 43 |
// Subtract the mean before scaling. The expanded form
|
|
|
|
| 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 }}) {
|
| 7 |
return;
|
| 8 |
+
}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
| 9 |
{{ env.wgsl.resourceDeclarations }}
|
| 10 |
|
| 11 |
// Rank-2 [N, C] inference vec4 specialization. Vectors run across adjacent channels, so
|
| 12 |
// scale/bias/mean/var are also bound as vec4<f32> and C must be divisible by 4.
|
| 13 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 14 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 15 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE, "i", "params.count4") }}
|
| 16 |
let channel4 = i % params.channels4;
|
| 17 |
let alpha = scale[channel4] * inverseSqrt(input_var[channel4] + vec4<f32>(params.epsilon));
|
| 18 |
// Subtract the mean before scaling. The expanded form
|
build/webgpu/batch-normalization-nchw-flat-vec4.wgsl.jinja
CHANGED
|
@@ -1,36 +1,11 @@
|
|
| 1 |
-
{% macro flat_index_2d(name="i", bound="params.count", guardInline=false
|
| 2 |
-
{%
|
| 3 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 4 |
-
// per-axis workgroup fold width (outputs > 16.7M elements).
|
| 5 |
-
{% elif note == "limit" %}
|
| 6 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 7 |
-
// per-axis workgroup fold width.
|
| 8 |
-
{% elif note == "device-axis" %}
|
| 9 |
-
// The flat dispatch is folded across x/y at a fixed per-axis workgroup
|
| 10 |
-
// width; gid.y carries the high portion of the output index.
|
| 11 |
-
{% elif note == "vec4-limit" %}
|
| 12 |
-
// 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
|
| 13 |
-
// per-axis workgroup fold width (the dispatch caps x and spills into y).
|
| 14 |
-
{% elif note == "element-limit" %}
|
| 15 |
-
// 2D-folded flat element index: gid.y carries the high bits past the
|
| 16 |
-
// dispatch's per-axis workgroup fold width.
|
| 17 |
-
{% elif note == "dispatch" %}
|
| 18 |
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 19 |
// per-axis workgroup fold width.
|
| 20 |
-
{
|
| 21 |
-
{% if bound == "" %}
|
| 22 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 23 |
-
{%- elif guardInline %}
|
| 24 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 25 |
-
if ({{ name }} >= {{ bound }}) { return; }
|
| 26 |
-
{%- else %}
|
| 27 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 28 |
if ({{ name }} >= {{ bound }}) {
|
| 29 |
return;
|
| 30 |
-
}
|
| 31 |
-
{%- endif %}
|
| 32 |
-
{% endmacro %}
|
| 33 |
-
|
| 34 |
{{ env.wgsl.resourceDeclarations }}
|
| 35 |
|
| 36 |
// Inference-only vec4 specialization for channel planes whose length is not a
|
|
@@ -45,7 +20,7 @@ fn normalize_at(value: f32, channel: u32) -> f32 {
|
|
| 45 |
|
| 46 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 47 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 48 |
-
{{ flat_index_2d("i", "params.count4") }}
|
| 49 |
let base = i * 4u;
|
| 50 |
let first = (base / params.spatial) % params.channels;
|
| 51 |
let last = ((base + 3u) / params.spatial) % params.channels;
|
|
|
|
| 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 }}) {
|
| 7 |
return;
|
| 8 |
+
}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
| 9 |
{{ env.wgsl.resourceDeclarations }}
|
| 10 |
|
| 11 |
// Inference-only vec4 specialization for channel planes whose length is not a
|
|
|
|
| 20 |
|
| 21 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 22 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 23 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE, "i", "params.count4") }}
|
| 24 |
let base = i * 4u;
|
| 25 |
let first = (base / params.spatial) % params.channels;
|
| 26 |
let last = ((base + 3u) / params.spatial) % params.channels;
|
build/webgpu/batch-normalization-nchw-vec4.wgsl.jinja
CHANGED
|
@@ -1,36 +1,11 @@
|
|
| 1 |
-
{% macro flat_index_2d(name="i", bound="params.count", guardInline=false
|
| 2 |
-
{%
|
| 3 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 4 |
-
// per-axis workgroup fold width (outputs > 16.7M elements).
|
| 5 |
-
{% elif note == "limit" %}
|
| 6 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 7 |
-
// per-axis workgroup fold width.
|
| 8 |
-
{% elif note == "device-axis" %}
|
| 9 |
-
// The flat dispatch is folded across x/y at a fixed per-axis workgroup
|
| 10 |
-
// width; gid.y carries the high portion of the output index.
|
| 11 |
-
{% elif note == "vec4-limit" %}
|
| 12 |
-
// 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
|
| 13 |
-
// per-axis workgroup fold width (the dispatch caps x and spills into y).
|
| 14 |
-
{% elif note == "element-limit" %}
|
| 15 |
-
// 2D-folded flat element index: gid.y carries the high bits past the
|
| 16 |
-
// dispatch's per-axis workgroup fold width.
|
| 17 |
-
{% elif note == "dispatch" %}
|
| 18 |
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 19 |
// per-axis workgroup fold width.
|
| 20 |
-
{
|
| 21 |
-
{% if bound == "" %}
|
| 22 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 23 |
-
{%- elif guardInline %}
|
| 24 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 25 |
-
if ({{ name }} >= {{ bound }}) { return; }
|
| 26 |
-
{%- else %}
|
| 27 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 28 |
if ({{ name }} >= {{ bound }}) {
|
| 29 |
return;
|
| 30 |
-
}
|
| 31 |
-
{%- endif %}
|
| 32 |
-
{% endmacro %}
|
| 33 |
-
|
| 34 |
{{ env.wgsl.resourceDeclarations }}
|
| 35 |
|
| 36 |
// Inference-only vec4 specialization: 128-bit loads/stores over x/y. Gated to
|
|
@@ -39,7 +14,7 @@
|
|
| 39 |
// (x - mean) * inverseSqrt(var + epsilon) * scale + bias
|
| 40 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 41 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 42 |
-
{{ flat_index_2d("i", "params.count4") }}
|
| 43 |
let channel = (i / params.spatial4) % params.channels;
|
| 44 |
let normalized = (x[i] - vec4<f32>(input_mean[channel])) * inverseSqrt(input_var[channel] + params.epsilon);
|
| 45 |
y[i] = normalized * vec4<f32>(scale[channel]) + vec4<f32>(bias[channel]);
|
|
|
|
| 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 }}) {
|
| 7 |
return;
|
| 8 |
+
}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
| 9 |
{{ env.wgsl.resourceDeclarations }}
|
| 10 |
|
| 11 |
// Inference-only vec4 specialization: 128-bit loads/stores over x/y. Gated to
|
|
|
|
| 14 |
// (x - mean) * inverseSqrt(var + epsilon) * scale + bias
|
| 15 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 16 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 17 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE, "i", "params.count4") }}
|
| 18 |
let channel = (i / params.spatial4) % params.channels;
|
| 19 |
let normalized = (x[i] - vec4<f32>(input_mean[channel])) * inverseSqrt(input_var[channel] + params.epsilon);
|
| 20 |
y[i] = normalized * vec4<f32>(scale[channel]) + vec4<f32>(bias[channel]);
|
build/webgpu/batch-normalization-nchw.wgsl.jinja
CHANGED
|
@@ -1,40 +1,15 @@
|
|
| 1 |
-
{% macro flat_index_2d(name="i", bound="params.count", guardInline=false
|
| 2 |
-
{%
|
| 3 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 4 |
-
// per-axis workgroup fold width (outputs > 16.7M elements).
|
| 5 |
-
{% elif note == "limit" %}
|
| 6 |
-
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 7 |
-
// per-axis workgroup fold width.
|
| 8 |
-
{% elif note == "device-axis" %}
|
| 9 |
-
// The flat dispatch is folded across x/y at a fixed per-axis workgroup
|
| 10 |
-
// width; gid.y carries the high portion of the output index.
|
| 11 |
-
{% elif note == "vec4-limit" %}
|
| 12 |
-
// 2D-folded flat vec4 index: gid.y carries the high bits past the dispatch's
|
| 13 |
-
// per-axis workgroup fold width (the dispatch caps x and spills into y).
|
| 14 |
-
{% elif note == "element-limit" %}
|
| 15 |
-
// 2D-folded flat element index: gid.y carries the high bits past the
|
| 16 |
-
// dispatch's per-axis workgroup fold width.
|
| 17 |
-
{% elif note == "dispatch" %}
|
| 18 |
// 2D-folded flat index: gid.y carries the high bits past the dispatch's
|
| 19 |
// per-axis workgroup fold width.
|
| 20 |
-
{
|
| 21 |
-
{% if bound == "" %}
|
| 22 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 23 |
-
{%- elif guardInline %}
|
| 24 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 25 |
-
if ({{ name }} >= {{ bound }}) { return; }
|
| 26 |
-
{%- else %}
|
| 27 |
-
let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
|
| 28 |
if ({{ name }} >= {{ bound }}) {
|
| 29 |
return;
|
| 30 |
-
}
|
| 31 |
-
{%- endif %}
|
| 32 |
-
{% endmacro %}
|
| 33 |
-
|
| 34 |
{{ env.wgsl.resourceDeclarations }}
|
| 35 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 36 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 37 |
-
{{ flat_index_2d("index") }}
|
| 38 |
let spatial = params.height * params.width;
|
| 39 |
let channel = (index / spatial) % params.channels;
|
| 40 |
{% if usesF16 %}
|
|
|
|
| 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 }}) {
|
| 7 |
return;
|
| 8 |
+
}{% endmacro %}
|
|
|
|
|
|
|
|
|
|
| 9 |
{{ env.wgsl.resourceDeclarations }}
|
| 10 |
@compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
|
| 11 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 12 |
+
{{ flat_index_2d(tunables.WORKGROUP_SIZE, "index") }}
|
| 13 |
let spatial = params.height * params.width;
|
| 14 |
let channel = (index / spatial) % params.channels;
|
| 15 |
{% if usesF16 %}
|
build/webgpu/manifest.json
CHANGED
|
@@ -20,33 +20,26 @@
|
|
| 20 |
},
|
| 21 |
"when": ["inferenceContractOk"],
|
| 22 |
"bindings": {
|
| 23 |
-
"scale": { "
|
| 24 |
-
"bias": { "arg": "b", "
|
| 25 |
-
"input_mean": { "arg": "inputMean", "
|
| 26 |
-
"input_var": { "arg": "inputVar", "
|
| 27 |
-
"
|
| 28 |
-
"
|
| 29 |
-
"
|
| 30 |
-
"
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
"elementType": "vec4<f32>"
|
| 35 |
-
},
|
| 36 |
-
"input_var_2": { "arg": "inputVar", "name": "input_var", "buffer": "read-only-storage", "elementType": "vec4<f32>" },
|
| 37 |
-
"y_2": { "name": "y", "buffer": "storage", "elementType": "vec4<f32>" },
|
| 38 |
-
"params_2": {
|
| 39 |
"name": "params",
|
| 40 |
-
"buffer": "uniform",
|
| 41 |
"struct": [
|
| 42 |
{ "name": "count4", "type": "u32", "value": "numel(shapes.y) / 4" },
|
| 43 |
{ "name": "channels4", "type": "u32", "value": "dim(shapes.x, 1) / 4" },
|
| 44 |
{ "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
|
| 45 |
]
|
| 46 |
},
|
| 47 |
-
"
|
| 48 |
"name": "params",
|
| 49 |
-
"buffer": "uniform",
|
| 50 |
"struct": [
|
| 51 |
{ "name": "count4", "type": "u32", "value": "numel(shapes.y) / 4" },
|
| 52 |
{ "name": "spatial", "type": "u32", "value": "inner(shapes.x, 1)" },
|
|
@@ -54,9 +47,8 @@
|
|
| 54 |
{ "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
|
| 55 |
]
|
| 56 |
},
|
| 57 |
-
"
|
| 58 |
"name": "params",
|
| 59 |
-
"buffer": "uniform",
|
| 60 |
"struct": [
|
| 61 |
{ "name": "count4", "type": "u32", "value": "numel(shapes.y) / 4" },
|
| 62 |
{ "name": "spatial4", "type": "u32", "value": "inner(shapes.x, 1) / 4" },
|
|
@@ -109,7 +101,7 @@
|
|
| 109 |
"id": "main",
|
| 110 |
"name": "BatchNormalization.NcInferenceVec4",
|
| 111 |
"shader": "batch-normalization-nc-vec4.wgsl.jinja",
|
| 112 |
-
"bindings": ["
|
| 113 |
"dispatch": {
|
| 114 |
"x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
| 115 |
"y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
|
@@ -127,7 +119,7 @@
|
|
| 127 |
"id": "main",
|
| 128 |
"name": "BatchNormalization.InferenceFlatVec4",
|
| 129 |
"shader": "batch-normalization-nchw-flat-vec4.wgsl.jinja",
|
| 130 |
-
"bindings": ["
|
| 131 |
"dispatch": {
|
| 132 |
"x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
| 133 |
"y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
|
@@ -145,7 +137,7 @@
|
|
| 145 |
"id": "main",
|
| 146 |
"name": "BatchNormalization.InferenceVec4",
|
| 147 |
"shader": "batch-normalization-nchw-vec4.wgsl.jinja",
|
| 148 |
-
"bindings": ["
|
| 149 |
"dispatch": {
|
| 150 |
"x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
| 151 |
"y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
|
|
|
| 20 |
},
|
| 21 |
"when": ["inferenceContractOk"],
|
| 22 |
"bindings": {
|
| 23 |
+
"scale": { "elementType": "$T" },
|
| 24 |
+
"bias": { "arg": "b", "elementType": "$T" },
|
| 25 |
+
"input_mean": { "arg": "inputMean", "elementType": "$T" },
|
| 26 |
+
"input_var": { "arg": "inputVar", "elementType": "$T" },
|
| 27 |
+
"x_main": { "name": "x", "elementType": "vec4<f32>" },
|
| 28 |
+
"scale_main": { "name": "scale", "elementType": "vec4<f32>" },
|
| 29 |
+
"bias_b": { "arg": "b", "name": "bias", "elementType": "vec4<f32>" },
|
| 30 |
+
"input_mean_main": { "arg": "inputMean", "name": "input_mean", "elementType": "vec4<f32>" },
|
| 31 |
+
"input_var_main": { "arg": "inputVar", "name": "input_var", "elementType": "vec4<f32>" },
|
| 32 |
+
"y_main": { "name": "y", "elementType": "vec4<f32>" },
|
| 33 |
+
"params_main": {
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 34 |
"name": "params",
|
|
|
|
| 35 |
"struct": [
|
| 36 |
{ "name": "count4", "type": "u32", "value": "numel(shapes.y) / 4" },
|
| 37 |
{ "name": "channels4", "type": "u32", "value": "dim(shapes.x, 1) / 4" },
|
| 38 |
{ "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
|
| 39 |
]
|
| 40 |
},
|
| 41 |
+
"params__uniform": {
|
| 42 |
"name": "params",
|
|
|
|
| 43 |
"struct": [
|
| 44 |
{ "name": "count4", "type": "u32", "value": "numel(shapes.y) / 4" },
|
| 45 |
{ "name": "spatial", "type": "u32", "value": "inner(shapes.x, 1)" },
|
|
|
|
| 47 |
{ "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
|
| 48 |
]
|
| 49 |
},
|
| 50 |
+
"params_nchw_inference_vec4": {
|
| 51 |
"name": "params",
|
|
|
|
| 52 |
"struct": [
|
| 53 |
{ "name": "count4", "type": "u32", "value": "numel(shapes.y) / 4" },
|
| 54 |
{ "name": "spatial4", "type": "u32", "value": "inner(shapes.x, 1) / 4" },
|
|
|
|
| 101 |
"id": "main",
|
| 102 |
"name": "BatchNormalization.NcInferenceVec4",
|
| 103 |
"shader": "batch-normalization-nc-vec4.wgsl.jinja",
|
| 104 |
+
"bindings": ["x_main", "scale_main", "bias_b", "input_mean_main", "input_var_main", "y_main", "params_main"],
|
| 105 |
"dispatch": {
|
| 106 |
"x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
| 107 |
"y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
|
|
|
| 119 |
"id": "main",
|
| 120 |
"name": "BatchNormalization.InferenceFlatVec4",
|
| 121 |
"shader": "batch-normalization-nchw-flat-vec4.wgsl.jinja",
|
| 122 |
+
"bindings": ["x_main", "scale", "bias", "input_mean", "input_var", "y_main", "params__uniform"],
|
| 123 |
"dispatch": {
|
| 124 |
"x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
| 125 |
"y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
|
|
|
| 137 |
"id": "main",
|
| 138 |
"name": "BatchNormalization.InferenceVec4",
|
| 139 |
"shader": "batch-normalization-nchw-vec4.wgsl.jinja",
|
| 140 |
+
"bindings": ["x_main", "scale", "bias", "input_mean", "input_var", "y_main", "params_nchw_inference_vec4"],
|
| 141 |
"dispatch": {
|
| 142 |
"x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
| 143 |
"y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,24 +1,24 @@
|
|
| 1 |
{
|
| 2 |
"name": "ai.onnx.BatchNormalization",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
-
"batch-normalization-nc-vec4.wgsl.jinja": "
|
| 11 |
-
"batch-normalization-nchw-flat-vec4.wgsl.jinja": "
|
| 12 |
-
"batch-normalization-nchw-vec4.wgsl.jinja": "
|
| 13 |
-
"batch-normalization-nchw.wgsl.jinja": "
|
| 14 |
"bench.json": "yPXcRY22XiVLwm5V3NsowDKaM6BKq2z0Batz3+SR+cI=",
|
| 15 |
-
"manifest.json": "
|
| 16 |
"test.json": "AAN+Vhz8gU2l1tY8tLPNyRAYMpP+8dGalRvLSrzR/5M="
|
| 17 |
}
|
| 18 |
},
|
| 19 |
-
"provenance": { "kernel": { "sha": "
|
| 20 |
"webgpu": {
|
| 21 |
-
"manifestSpec": "2.
|
| 22 |
"variants": {
|
| 23 |
"inference_scalar": ["batch-normalization-nchw.wgsl.jinja"],
|
| 24 |
"nc_inference_vec4": ["batch-normalization-nc-vec4.wgsl.jinja"],
|
|
|
|
| 1 |
{
|
| 2 |
"name": "ai.onnx.BatchNormalization",
|
| 3 |
+
"id": "_ai_onnx_batchnormalization_webgpu_5d34230",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
+
"batch-normalization-nc-vec4.wgsl.jinja": "8jefLjDLqeMBpbJ+Y634+YmYPLNef41p2INh8Rnw1IQ=",
|
| 11 |
+
"batch-normalization-nchw-flat-vec4.wgsl.jinja": "SYG1mXn0UqWSj3Ke0OoYE83te55U9RfifmNRcjsHWvw=",
|
| 12 |
+
"batch-normalization-nchw-vec4.wgsl.jinja": "FW7YWctHdrjfgK+eme5TzzeEIJsq5U9eJgtsC4B6cak=",
|
| 13 |
+
"batch-normalization-nchw.wgsl.jinja": "BQ28jcj3rCy5YKJOyOZ/aj9gWVUOu0CKB4PInGOfNIY=",
|
| 14 |
"bench.json": "yPXcRY22XiVLwm5V3NsowDKaM6BKq2z0Batz3+SR+cI=",
|
| 15 |
+
"manifest.json": "JnANb6SyQMZS8uWm5kq0iAsuiL0YQ/4AObgi6+T0tCc=",
|
| 16 |
"test.json": "AAN+Vhz8gU2l1tY8tLPNyRAYMpP+8dGalRvLSrzR/5M="
|
| 17 |
}
|
| 18 |
},
|
| 19 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 20 |
"webgpu": {
|
| 21 |
+
"manifestSpec": "2.1",
|
| 22 |
"variants": {
|
| 23 |
"inference_scalar": ["batch-normalization-nchw.wgsl.jinja"],
|
| 24 |
"nc_inference_vec4": ["batch-normalization-nc-vec4.wgsl.jinja"],
|