Xenova HF Staff commited on
Commit
739291e
·
verified ·
1 Parent(s): 53e6a1d

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -40,7 +40,7 @@ See the [ONNX `Sub` spec](https://onnx.ai/onnx/operators/onnx__Sub.html) for the
40
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
41
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
42
  - [`test.json`](build/webgpu/test.json) — correctness cases
43
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
44
  - [`binary-broadcast-vec4.wgsl.jinja`](build/webgpu/binary-broadcast-vec4.wgsl.jinja)
45
  - [`binary-broadcast.wgsl.jinja`](build/webgpu/binary-broadcast.wgsl.jinja)
46
  - [`binary-vec4.wgsl.jinja`](build/webgpu/binary-vec4.wgsl.jinja)
@@ -48,7 +48,7 @@ See the [ONNX `Sub` spec](https://onnx.ai/onnx/operators/onnx__Sub.html) for the
48
  ## Use with `@huggingface/kernels`
49
 
50
  ```sh
51
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
52
  ```
53
 
54
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
40
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
41
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
42
  - [`test.json`](build/webgpu/test.json) — correctness cases
43
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
44
  - [`binary-broadcast-vec4.wgsl.jinja`](build/webgpu/binary-broadcast-vec4.wgsl.jinja)
45
  - [`binary-broadcast.wgsl.jinja`](build/webgpu/binary-broadcast.wgsl.jinja)
46
  - [`binary-vec4.wgsl.jinja`](build/webgpu/binary-vec4.wgsl.jinja)
 
48
  ## Use with `@huggingface/kernels`
49
 
50
  ```sh
51
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
52
  ```
53
 
54
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/binary-broadcast-vec4.wgsl.jinja CHANGED
@@ -33,15 +33,25 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
33
  {% endfor %}
34
  return offset;
35
  {% endif %}
36
- }
37
- {%- endmacro %}{% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
38
  {% set op_numel = namespace(value=1) %}
39
- {% for d in opShape %}{% set op_numel.value = op_numel.value * d %}{% endfor %}
 
 
40
  {% set out_numel = namespace(value=1) %}
41
- {% for d in outShape %}{% set out_numel.value = out_numel.value * d %}{% endfor %}
42
- {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %})
43
- {%- endmacro %}
44
-
 
 
 
 
 
 
 
 
45
  {{ env.wgsl.resourceDeclarations }}
46
 
47
  // Vec4 broadcast binary op. Offsets use compile-time strides, and each
@@ -99,12 +109,12 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
99
 
100
  {% if not a_same.value and a_mode != "scalar" %}
101
  {{ offset_fn("a_offset", aShape, aRank, false, a_numel.value, cShape, cRank, c_numel.value) }}
102
- {% endif %}
103
 
 
104
  {% if not b_same.value and b_mode != "scalar" %}
105
  {{ offset_fn("b_offset", bShape, bRank, false, b_numel.value, cShape, cRank, c_numel.value) }}
106
- {% endif %}
107
 
 
108
  {% set is_int = scalar == "i32" or scalar == "u32" %}
109
  {% set acc = scalar if is_int else "f32" %}
110
  // Narrow integer operations wrap modulo the logical dtype width; int8/uint8
@@ -120,12 +130,7 @@ fn wrap_dtype(v: vec4<u32>) -> vec4<u32> { return v & vec4<u32>(0xFFu); }
120
  {% endif %}
121
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
122
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
123
- // 2D-folded flat index: gid.y carries the high bits when the element count exceeds the
124
- // per-axis dispatch fold width (the dispatch caps x and spills the rest into y).
125
- let i4 = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
126
- if (i4 >= params.count) {
127
- return;
128
- }
129
  {% set needsBase = c_numel.value != 0 and ((not a_same.value and a_mode != "scalar") or (not b_same.value and b_mode != "scalar")) %}
130
  {% if needsBase %}
131
  let base = i4 * 4u;
 
33
  {% endfor %}
34
  return offset;
35
  {% endif %}
36
+ }{% endmacro %}
37
+ {% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
38
  {% set op_numel = namespace(value=1) %}
39
+ {% for d in opShape %}
40
+ {% set op_numel.value = op_numel.value * d %}
41
+ {% endfor %}
42
  {% set out_numel = namespace(value=1) %}
43
+ {% for d in outShape %}
44
+ {% set out_numel.value = out_numel.value * d %}
45
+ {% endfor %}
46
+ {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %}){% endmacro %}
47
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
48
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
49
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
50
+ // per-axis workgroup fold width.
51
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
52
+ if ({{ name }} >= {{ bound }}) {
53
+ return;
54
+ }{% endmacro %}
55
  {{ env.wgsl.resourceDeclarations }}
56
 
57
  // Vec4 broadcast binary op. Offsets use compile-time strides, and each
 
109
 
110
  {% if not a_same.value and a_mode != "scalar" %}
111
  {{ offset_fn("a_offset", aShape, aRank, false, a_numel.value, cShape, cRank, c_numel.value) }}
 
112
 
113
+ {% endif %}
114
  {% if not b_same.value and b_mode != "scalar" %}
115
  {{ offset_fn("b_offset", bShape, bRank, false, b_numel.value, cShape, cRank, c_numel.value) }}
 
116
 
117
+ {% endif %}
118
  {% set is_int = scalar == "i32" or scalar == "u32" %}
119
  {% set acc = scalar if is_int else "f32" %}
120
  // Narrow integer operations wrap modulo the logical dtype width; int8/uint8
 
130
  {% endif %}
131
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
132
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
133
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "i4") }}
 
 
 
 
 
134
  {% set needsBase = c_numel.value != 0 and ((not a_same.value and a_mode != "scalar") or (not b_same.value and b_mode != "scalar")) %}
135
  {% if needsBase %}
136
  let base = i4 * 4u;
build/webgpu/binary-broadcast.wgsl.jinja CHANGED
@@ -12,9 +12,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
  for (var i = begin; i < end; i = i + 1u) {
13
  {%- endmacro %}
14
  {% macro flat_tail_close() %}
15
- }
16
- {% endmacro %}
17
-
18
  {% macro offset_fn(fn_name, opShape, opRank, op_same, op_numel, outShape, outRank, out_numel) %}
19
  fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif %}) -> u32 {
20
  {% if out_numel == 0 %}
@@ -50,14 +48,18 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
50
  {% endfor %}
51
  return offset;
52
  {% endif %}
53
- }
54
- {%- endmacro %}{% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
55
  {% set op_numel = namespace(value=1) %}
56
- {% for d in opShape %}{% set op_numel.value = op_numel.value * d %}{% endfor %}
 
 
57
  {% set out_numel = namespace(value=1) %}
58
- {% for d in outShape %}{% set out_numel.value = out_numel.value * d %}{% endfor %}
59
- {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %})
60
- {%- endmacro %}{% macro broadcast_offset_fn(fn_name, opShape, opRank, outShape, outRank) %}
 
 
61
  {% set op_numel = namespace(value=1) %}
62
  {% for d in opShape %}
63
  {% set op_numel.value = op_numel.value * d %}
@@ -74,8 +76,8 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
74
  {% endif %}
75
  {% endfor %}
76
  {% endif %}
77
- {{ offset_fn(fn_name, opShape, opRank, op_same.value, op_numel.value, outShape, outRank, out_numel.value) }}
78
- {%- endmacro %}{% macro binary_broadcast_offsets() %}
79
  {% set aShape = aShape | default([]) %}
80
  {% set aRank = aRank | default(0) %}
81
  {% set bShape = bShape | default([]) %}
@@ -84,12 +86,9 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
84
  {% set cRank = cRank | default(0) %}
85
  {{ broadcast_offset_fn("a_offset", aShape, aRank, cShape, cRank) }}
86
 
87
- {{ broadcast_offset_fn("b_offset", bShape, bRank, cShape, cRank) }}
88
- {%- endmacro %}
89
-
90
  {{ env.wgsl.resourceDeclarations }}
91
 
92
-
93
  {{ binary_broadcast_offsets() }}
94
 
95
  {{ flat_tail_open() }}
@@ -111,5 +110,5 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
111
  let bv = f32(b[{{ broadcast_offset_call("b_offset", bShape, cShape, "i") }}]);
112
  c[i] = {{ scalar }}(av - bv);
113
  {% endif %}
114
- {{ flat_tail_close() -}}
115
  }
 
12
  for (var i = begin; i < end; i = i + 1u) {
13
  {%- endmacro %}
14
  {% macro flat_tail_close() %}
15
+ }{% endmacro %}
 
 
16
  {% macro offset_fn(fn_name, opShape, opRank, op_same, op_numel, outShape, outRank, out_numel) %}
17
  fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif %}) -> u32 {
18
  {% if out_numel == 0 %}
 
48
  {% endfor %}
49
  return offset;
50
  {% endif %}
51
+ }{% endmacro %}
52
+ {% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
53
  {% set op_numel = namespace(value=1) %}
54
+ {% for d in opShape %}
55
+ {% set op_numel.value = op_numel.value * d %}
56
+ {% endfor %}
57
  {% set out_numel = namespace(value=1) %}
58
+ {% for d in outShape %}
59
+ {% set out_numel.value = out_numel.value * d %}
60
+ {% endfor %}
61
+ {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %}){% endmacro %}
62
+ {% macro broadcast_offset_fn(fn_name, opShape, opRank, outShape, outRank) %}
63
  {% set op_numel = namespace(value=1) %}
64
  {% for d in opShape %}
65
  {% set op_numel.value = op_numel.value * d %}
 
76
  {% endif %}
77
  {% endfor %}
78
  {% endif %}
79
+ {{ offset_fn(fn_name, opShape, opRank, op_same.value, op_numel.value, outShape, outRank, out_numel.value) }}{% endmacro %}
80
+ {% macro binary_broadcast_offsets() %}
81
  {% set aShape = aShape | default([]) %}
82
  {% set aRank = aRank | default(0) %}
83
  {% set bShape = bShape | default([]) %}
 
86
  {% set cRank = cRank | default(0) %}
87
  {{ broadcast_offset_fn("a_offset", aShape, aRank, cShape, cRank) }}
88
 
89
+ {{ broadcast_offset_fn("b_offset", bShape, bRank, cShape, cRank) }}{% endmacro %}
 
 
90
  {{ env.wgsl.resourceDeclarations }}
91
 
 
92
  {{ binary_broadcast_offsets() }}
93
 
94
  {{ flat_tail_open() }}
 
110
  let bv = f32(b[{{ broadcast_offset_call("b_offset", bShape, cShape, "i") }}]);
111
  c[i] = {{ scalar }}(av - bv);
112
  {% endif %}
113
+ {{ flat_tail_close() }}
114
  }
build/webgpu/binary-vec4.wgsl.jinja CHANGED
@@ -1,32 +1,41 @@
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
- {% set vec4PerThread = vec4PerThread %}
 
 
 
 
 
 
 
 
 
4
  {% if vec4PerThread > 1 %}
5
  const ITEMS: u32 = {{ vec4PerThread }}u;
6
- {% endif %}
7
 
 
8
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
9
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
10
- // 2D-folded flat index: gid.y carries the high bits when the element count exceeds the
11
- // per-axis dispatch fold width (the dispatch caps x and spills the rest into y).
12
  {% if vec4PerThread > 1 %}
13
  // Each invocation walks ITEMS vec4 groups a span apart. Consecutive lanes
14
  // access consecutive words on every step, while each lane can keep several
15
  // independent loads in flight.
16
- let tid = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
17
  let span = (params.count + ITEMS - 1u) / ITEMS;
 
 
 
 
 
 
18
  for (var j = 0u; j < ITEMS; j = j + 1u) {
19
  let i = tid + j * span;
20
  if (i >= params.count) {
21
  break;
22
  }
23
  {% else %}
24
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ tunables.WORKGROUP_SIZE }}u;
25
- if (i >= params.count) {
26
- return;
27
- }
28
  {% endif %}
29
-
30
  {% set scalarOperand = scalarOperand if scalarOperand is defined else "" %}
31
  {% if scalarOperand == "a" %}
32
  // One-element operand: read once and splat across the vector.
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
4
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
5
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
6
+ // per-axis workgroup fold width.
7
+ {% if bound == "" %}
8
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% else %}
9
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
10
+ if ({{ name }} >= {{ bound }}) {
11
+ return;
12
+ }{% endif %}{% endmacro %}
13
  {% if vec4PerThread > 1 %}
14
  const ITEMS: u32 = {{ vec4PerThread }}u;
 
15
 
16
+ {% endif %}
17
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
18
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
 
 
19
  {% if vec4PerThread > 1 %}
20
  // Each invocation walks ITEMS vec4 groups a span apart. Consecutive lanes
21
  // access consecutive words on every step, while each lane can keep several
22
  // independent loads in flight.
23
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "tid", "") }}
24
  let span = (params.count + ITEMS - 1u) / ITEMS;
25
+ // Lanes at or past span exist only because the dispatch rounds up to whole
26
+ // workgroups. Lane span + k would start on lane k's second group and rewrite
27
+ // up to ITEMS - 1 groups another lane already stored.
28
+ if (tid >= span) {
29
+ return;
30
+ }
31
  for (var j = 0u; j < ITEMS; j = j + 1u) {
32
  let i = tid + j * span;
33
  if (i >= params.count) {
34
  break;
35
  }
36
  {% else %}
37
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE) }}
 
 
 
38
  {% endif %}
 
39
  {% set scalarOperand = scalarOperand if scalarOperand is defined else "" %}
40
  {% if scalarOperand == "a" %}
41
  // One-element operand: read once and splat across the vector.
build/webgpu/manifest.json CHANGED
@@ -17,7 +17,7 @@
17
  "passes": [
18
  {
19
  "id": "main",
20
- "name": "Sub.vec4",
21
  "shader": "binary-vec4.wgsl.jinja",
22
  "derive": {
23
  "op": "\"sub\"",
@@ -41,7 +41,7 @@
41
  "passes": [
42
  {
43
  "id": "main",
44
- "name": "Sub.scalarBVec4",
45
  "shader": "binary-vec4.wgsl.jinja",
46
  "derive": {
47
  "op": "\"sub\"",
@@ -49,7 +49,7 @@
49
  "scalarOperand": "\"b\"",
50
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
51
  },
52
- "bindings": ["a", "b_2", "c_binary", "params"],
53
  "dispatch": {
54
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
55
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -66,7 +66,7 @@
66
  "passes": [
67
  {
68
  "id": "main",
69
- "name": "Sub.scalarAVec4",
70
  "shader": "binary-vec4.wgsl.jinja",
71
  "derive": {
72
  "op": "\"sub\"",
@@ -74,7 +74,7 @@
74
  "scalarOperand": "\"a\"",
75
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
76
  },
77
- "bindings": ["a_2", "b", "c_binary", "params"],
78
  "dispatch": {
79
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
80
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -108,7 +108,7 @@
108
  "op": "\"sub\"",
109
  "cDtype": "tensorDtypes.c"
110
  },
111
- "bindings": ["a_3", "b_3", "c_2_binary", "params"],
112
  "dispatch": {
113
  "x": "min(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
114
  "y": "ceilDiv(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -117,36 +117,6 @@
117
  }
118
  ]
119
  },
120
- {
121
- "id": "same_shape_scalar_x4",
122
- "priority": 15,
123
- "when": ["sameShape(shapes.a, shapes.c)", "sameShape(shapes.b, shapes.c)", "numel(shapes.c) > 0", "numel(shapes.c) % 4 != 0", "f16Ok(dtypes.T)"],
124
- "derive": { "scalar": "dtypes.T" },
125
- "passes": [
126
- {
127
- "id": "main",
128
- "name": "Sub",
129
- "shader": "binary-broadcast.wgsl.jinja",
130
- "derive": {
131
- "aShape": "shapes.a",
132
- "bShape": "shapes.b",
133
- "cShape": "shapes.c",
134
- "aRank": "ranks.a",
135
- "bRank": "ranks.b",
136
- "cRank": "ranks.c",
137
- "op": "\"sub\"",
138
- "cDtype": "tensorDtypes.c",
139
- "itemsPerInvocation": 4
140
- },
141
- "bindings": ["a_2", "b_2", "c_3", "params_2"],
142
- "dispatch": {
143
- "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
144
- "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
145
- "z": 1
146
- }
147
- }
148
- ]
149
- },
150
  {
151
  "id": "broadcast",
152
  "when": ["ranks.a <= ranks.c", "ranks.b <= ranks.c", "f16Ok(dtypes.T)"],
@@ -167,7 +137,7 @@
167
  "cDtype": "tensorDtypes.c",
168
  "itemsPerInvocation": 4
169
  },
170
- "bindings": ["a_2", "b_2", "c_3", "params_2"],
171
  "dispatch": {
172
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
173
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -178,20 +148,16 @@
178
  }
179
  ],
180
  "bindings": {
181
- "a": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
182
- "b": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
183
- "c_binary": { "buffer": "storage", "elementType": "$vectorScalar", "name": "c" },
184
- "params": { "buffer": "uniform", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c) / 4" }] },
185
- "b_2": { "buffer": "read-only-storage", "name": "b", "elementType": "$scalar" },
186
- "a_2": { "buffer": "read-only-storage", "name": "a", "elementType": "$scalar" },
187
- "a_3": { "buffer": "read-only-storage", "name": "a", "elementType": "$aElement" },
188
- "b_3": { "buffer": "read-only-storage", "name": "b", "elementType": "$bElement" },
189
- "c_2_binary": { "buffer": "storage", "name": "c", "elementType": "$vec4Scalar" },
190
- "c_3": { "buffer": "storage", "name": "c", "elementType": "$scalar" },
191
- "params_2": {
192
- "buffer": "uniform",
193
- "name": "params",
194
- "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c)" }]
195
- }
196
  }
197
  }
 
17
  "passes": [
18
  {
19
  "id": "main",
20
+ "name": "Sub.Vec4",
21
  "shader": "binary-vec4.wgsl.jinja",
22
  "derive": {
23
  "op": "\"sub\"",
 
41
  "passes": [
42
  {
43
  "id": "main",
44
+ "name": "Sub.ScalarBVec4",
45
  "shader": "binary-vec4.wgsl.jinja",
46
  "derive": {
47
  "op": "\"sub\"",
 
49
  "scalarOperand": "\"b\"",
50
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
51
  },
52
+ "bindings": ["a", "b_scalar", "c_binary", "params"],
53
  "dispatch": {
54
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
55
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
 
66
  "passes": [
67
  {
68
  "id": "main",
69
+ "name": "Sub.ScalarAVec4",
70
  "shader": "binary-vec4.wgsl.jinja",
71
  "derive": {
72
  "op": "\"sub\"",
 
74
  "scalarOperand": "\"a\"",
75
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
76
  },
77
+ "bindings": ["a_scalar", "b", "c_binary", "params"],
78
  "dispatch": {
79
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
80
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
 
108
  "op": "\"sub\"",
109
  "cDtype": "tensorDtypes.c"
110
  },
111
+ "bindings": ["a_element", "b_element", "c_2_binary", "params"],
112
  "dispatch": {
113
  "x": "min(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
114
  "y": "ceilDiv(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), 65535)",
 
117
  }
118
  ]
119
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
  {
121
  "id": "broadcast",
122
  "when": ["ranks.a <= ranks.c", "ranks.b <= ranks.c", "f16Ok(dtypes.T)"],
 
137
  "cDtype": "tensorDtypes.c",
138
  "itemsPerInvocation": 4
139
  },
140
+ "bindings": ["a_scalar", "b_scalar", "c_scalar", "params_count"],
141
  "dispatch": {
142
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
143
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
 
148
  }
149
  ],
150
  "bindings": {
151
+ "a": { "elementType": "$vectorScalar" },
152
+ "b": { "elementType": "$vectorScalar" },
153
+ "c_binary": { "elementType": "$vectorScalar", "name": "c" },
154
+ "params": { "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c) / 4" }] },
155
+ "b_scalar": { "name": "b", "elementType": "$scalar" },
156
+ "a_scalar": { "name": "a", "elementType": "$scalar" },
157
+ "a_element": { "name": "a", "elementType": "$aElement" },
158
+ "b_element": { "name": "b", "elementType": "$bElement" },
159
+ "c_2_binary": { "name": "c", "elementType": "$vec4Scalar" },
160
+ "c_scalar": { "name": "c", "elementType": "$scalar" },
161
+ "params_count": { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c)" }] }
 
 
 
 
162
  }
163
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.Sub",
3
- "id": "_ai_onnx_sub_webgpu_b71720b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,22 +8,21 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "pKpR5KsoFpsBCy+SiB1GHNRrl1O74Sa6HHerqGDL8fM=",
11
- "binary-broadcast-vec4.wgsl.jinja": "7kHsU4qoI/bt9ckuLZGQxH/Q4fyZjwwDatAHRXkxkzM=",
12
- "binary-broadcast.wgsl.jinja": "tE1ZKbRZIFMRWoVAsweV25u/rFzTtJKaUF2MfhksWic=",
13
- "binary-vec4.wgsl.jinja": "J2E/4vcnoHkFlLOEpndT2gScbjlTaTqpS1CsoX/82WE=",
14
- "manifest.json": "FcUv055haswyZeZqv8cky0RBjSsyuLEcBGNDK/eUl/o=",
15
- "test.json": "q38v2+2C+ycRyWR3OOq+k2NL1ylkfrRQvdDZKwcrRYw="
16
  }
17
  },
18
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
19
  "webgpu": {
20
- "manifestSpec": "2.0",
21
  "variants": {
22
  "same_shape_vec4": ["binary-vec4.wgsl.jinja"],
23
  "scalar_b_vec4": ["binary-vec4.wgsl.jinja"],
24
  "scalar_a_vec4": ["binary-vec4.wgsl.jinja"],
25
  "broadcast_vec4": ["binary-broadcast-vec4.wgsl.jinja"],
26
- "same_shape_scalar_x4": ["binary-broadcast.wgsl.jinja"],
27
  "broadcast": ["binary-broadcast.wgsl.jinja"]
28
  }
29
  }
 
1
  {
2
  "name": "ai.onnx.Sub",
3
+ "id": "_ai_onnx_sub_webgpu_dfb9100",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "pKpR5KsoFpsBCy+SiB1GHNRrl1O74Sa6HHerqGDL8fM=",
11
+ "binary-broadcast-vec4.wgsl.jinja": "s8UMM9TiMY+birA4Ev2w45U4sJtyWv2ee2URz5Espho=",
12
+ "binary-broadcast.wgsl.jinja": "iZID+Tcs+bryrYhWORo/l7IsHC1sLwj3xBUEC5Yz43I=",
13
+ "binary-vec4.wgsl.jinja": "F4kTq6LWQ1LR93GSXrBwMUJP+ATaLqB3Xk30nxSPr5w=",
14
+ "manifest.json": "fnNmWeriT/i5ilO7vPgr2F8mzG5jdi2QAsoVbgON5N8=",
15
+ "test.json": "+USjswIWfolYJ4B+Q3jtQatqxyM/7yihRSg9iyL7ThY="
16
  }
17
  },
18
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
19
  "webgpu": {
20
+ "manifestSpec": "2.1",
21
  "variants": {
22
  "same_shape_vec4": ["binary-vec4.wgsl.jinja"],
23
  "scalar_b_vec4": ["binary-vec4.wgsl.jinja"],
24
  "scalar_a_vec4": ["binary-vec4.wgsl.jinja"],
25
  "broadcast_vec4": ["binary-broadcast-vec4.wgsl.jinja"],
 
26
  "broadcast": ["binary-broadcast.wgsl.jinja"]
27
  }
28
  }
build/webgpu/test.json CHANGED
@@ -524,7 +524,7 @@
524
  {
525
  "name": "scalar_b_vec4_route",
526
  "provenance": {
527
- "notes": "Route lock for the one-element-operand vec4 kernel: the other operand matches the output and the vector count is a multiple of four, so the scalar reads once and splats."
528
  },
529
  "inputs": {
530
  "a": {
@@ -543,7 +543,7 @@
543
  {
544
  "name": "scalar_a_vec4_route",
545
  "provenance": {
546
- "notes": "Route lock for the one-element-operand vec4 kernel: the other operand matches the output and the vector count is a multiple of four, so the scalar reads once and splats."
547
  },
548
  "inputs": {
549
  "a": {
 
524
  {
525
  "name": "scalar_b_vec4_route",
526
  "provenance": {
527
+ "notes": "One scalar operand broadcasts across a vec4-aligned output; the other operand already matches the output shape."
528
  },
529
  "inputs": {
530
  "a": {
 
543
  {
544
  "name": "scalar_a_vec4_route",
545
  "provenance": {
546
+ "notes": "One scalar operand broadcasts across a vec4-aligned output; the other operand already matches the output shape."
547
  },
548
  "inputs": {
549
  "a": {