Xenova HF Staff commited on
Commit
27ad418
·
verified ·
1 Parent(s): ab820a8

sync 6fdf6301e2bb

Browse files
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 + tuning 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,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.2
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, note="dispatch-limit") %}
2
- {% if note == "dispatch-limit" %}
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
- {% endif %}
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, note="dispatch-limit") %}
2
- {% if note == "dispatch-limit" %}
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
- {% endif %}
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, note="dispatch-limit") %}
2
- {% if note == "dispatch-limit" %}
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
- {% endif %}
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, note="dispatch-limit") %}
2
- {% if note == "dispatch-limit" %}
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
- {% endif %}
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": { "buffer": "read-only-storage", "elementType": "$T" },
24
- "bias": { "arg": "b", "buffer": "read-only-storage", "elementType": "$T" },
25
- "input_mean": { "arg": "inputMean", "buffer": "read-only-storage", "elementType": "$T" },
26
- "input_var": { "arg": "inputVar", "buffer": "read-only-storage", "elementType": "$T" },
27
- "x_2": { "name": "x", "buffer": "read-only-storage", "elementType": "vec4<f32>" },
28
- "scale_2": { "name": "scale", "buffer": "read-only-storage", "elementType": "vec4<f32>" },
29
- "bias_2": { "arg": "b", "name": "bias", "buffer": "read-only-storage", "elementType": "vec4<f32>" },
30
- "input_mean_2": {
31
- "arg": "inputMean",
32
- "name": "input_mean",
33
- "buffer": "read-only-storage",
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
- "params_3": {
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
- "params_4": {
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": ["x_2", "scale_2", "bias_2", "input_mean_2", "input_var_2", "y_2", "params_2"],
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": ["x_2", "scale", "bias", "input_mean", "input_var", "y_2", "params_3"],
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": ["x_2", "scale", "bias", "input_mean", "input_var", "y_2", "params_4"],
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": "_ai_onnx_batchnormalization_webgpu_4d2bfed",
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": "yZ3Hf6SM+D6EW4GvvG5Q0N5tuxQlys5l78zAfqwH0IE=",
11
- "batch-normalization-nchw-flat-vec4.wgsl.jinja": "ZU77NvidV7PpSQC2Fhc/pDgtSlG6EypswF7MWBi+qzw=",
12
- "batch-normalization-nchw-vec4.wgsl.jinja": "L4z18z+y+0qESVx8FkmHJhTORNeEAOPz1KoV/e7WkCE=",
13
- "batch-normalization-nchw.wgsl.jinja": "bv0TT2Ry2691M8bUjXYC1tlsrPZ+vPsGgdcUKSj59pg=",
14
  "bench.json": "yPXcRY22XiVLwm5V3NsowDKaM6BKq2z0Batz3+SR+cI=",
15
- "manifest.json": "PRmol3fOKZDuk/f21hi3uxRSbxsOhP7RQd+KA2zkW0E=",
16
  "test.json": "AAN+Vhz8gU2l1tY8tLPNyRAYMpP+8dGalRvLSrzR/5M="
17
  }
18
  },
19
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
20
  "webgpu": {
21
- "manifestSpec": "2.0",
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"],