Xenova HF Staff commited on
Commit
ee208cc
·
verified ·
1 Parent(s): 0356b3a

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -41,7 +41,7 @@ See the [ONNX `Less` spec](https://onnx.ai/onnx/operators/onnx__Less.html) for t
41
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
42
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
43
  - [`test.json`](build/webgpu/test.json) — correctness cases
44
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
45
  - [`compare-broadcast-vec4.wgsl.jinja`](build/webgpu/compare-broadcast-vec4.wgsl.jinja)
46
  - [`compare-broadcast.wgsl.jinja`](build/webgpu/compare-broadcast.wgsl.jinja)
47
  - [`compare-vec4.wgsl.jinja`](build/webgpu/compare-vec4.wgsl.jinja)
@@ -49,7 +49,7 @@ See the [ONNX `Less` spec](https://onnx.ai/onnx/operators/onnx__Less.html) for t
49
  ## Use with `@huggingface/kernels`
50
 
51
  ```sh
52
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
53
  ```
54
 
55
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
@@ -63,5 +63,5 @@ Replace each `*Data` placeholder with a typed array containing the corresponding
63
  import { getKernel } from "@huggingface/kernels";
64
 
65
  const kernel = await getKernel("webgpu-kernels/ai.onnx.Less", { version: 1 });
66
- const { c } = await kernel({ a: { data: aData, shape: [] }, b: { data: bData, shape: [] } });
67
  ```
 
41
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
42
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
43
  - [`test.json`](build/webgpu/test.json) — correctness cases
44
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
45
  - [`compare-broadcast-vec4.wgsl.jinja`](build/webgpu/compare-broadcast-vec4.wgsl.jinja)
46
  - [`compare-broadcast.wgsl.jinja`](build/webgpu/compare-broadcast.wgsl.jinja)
47
  - [`compare-vec4.wgsl.jinja`](build/webgpu/compare-vec4.wgsl.jinja)
 
49
  ## Use with `@huggingface/kernels`
50
 
51
  ```sh
52
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
53
  ```
54
 
55
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
63
  import { getKernel } from "@huggingface/kernels";
64
 
65
  const kernel = await getKernel("webgpu-kernels/ai.onnx.Less", { version: 1 });
66
+ const { c } = await kernel({ a: { data: aData, shape: [4] }, b: { data: bData, shape: [1] } });
67
  ```
build/webgpu/compare-broadcast-vec4.wgsl.jinja CHANGED
@@ -33,15 +33,17 @@ 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
  // Broadcast comparison vectorized over the innermost output axis. Each thread
@@ -74,12 +76,12 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
74
 
75
  {% if a_mode != "scalar" %}
76
  {{ offset_fn("a_offset", aShape, aRank, a_same.value, a_numel.value, cShape, cRank, c_numel.value) }}
77
- {% endif %}
78
 
 
79
  {% if b_mode != "scalar" %}
80
  {{ offset_fn("b_offset", bShape, bRank, b_same.value, b_numel.value, cShape, cRank, c_numel.value) }}
81
- {% endif %}
82
 
 
83
  {% set OP = {"equal": "==", "greater": ">", "greaterOrEqual": ">=", "less": "<", "lessOrEqual": "<="}[op] %}
84
  const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
85
 
 
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
  {{ env.wgsl.resourceDeclarations }}
48
 
49
  // Broadcast comparison vectorized over the innermost output axis. Each thread
 
76
 
77
  {% if a_mode != "scalar" %}
78
  {{ offset_fn("a_offset", aShape, aRank, a_same.value, a_numel.value, cShape, cRank, c_numel.value) }}
 
79
 
80
+ {% endif %}
81
  {% if b_mode != "scalar" %}
82
  {{ offset_fn("b_offset", bShape, bRank, b_same.value, b_numel.value, cShape, cRank, c_numel.value) }}
 
83
 
84
+ {% endif %}
85
  {% set OP = {"equal": "==", "greater": ">", "greaterOrEqual": ">=", "less": "<", "lessOrEqual": "<="}[op] %}
86
  const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
87
 
build/webgpu/compare-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,15 +86,12 @@ 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() }}
96
  c[i] = select(0u, 1u, a[{{ broadcast_offset_call("a_offset", aShape, cShape, "i") }}] < b[{{ broadcast_offset_call("b_offset", bShape, cShape, "i") }}]);
97
- {{ flat_tail_close() -}}
98
  }
 
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() }}
95
  c[i] = select(0u, 1u, a[{{ broadcast_offset_call("a_offset", aShape, cShape, "i") }}] < b[{{ broadcast_offset_call("b_offset", bShape, cShape, "i") }}]);
96
+ {{ flat_tail_close() }}
97
  }
build/webgpu/compare-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": "Less.vec4",
21
  "shader": "compare-vec4.wgsl.jinja",
22
  "derive": {
23
  "op": "\"less\"",
@@ -40,14 +40,14 @@
40
  "passes": [
41
  {
42
  "id": "main",
43
- "name": "Less.scalarBVec4",
44
  "shader": "compare-vec4.wgsl.jinja",
45
  "derive": {
46
  "op": "\"less\"",
47
  "scalarOperand": "\"b\"",
48
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
49
  },
50
- "bindings": ["a", "b_2", "c", "params"],
51
  "dispatch": {
52
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
53
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -64,14 +64,14 @@
64
  "passes": [
65
  {
66
  "id": "main",
67
- "name": "Less.scalarAVec4",
68
  "shader": "compare-vec4.wgsl.jinja",
69
  "derive": {
70
  "op": "\"less\"",
71
  "scalarOperand": "\"a\"",
72
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
73
  },
74
- "bindings": ["a_2", "b", "c", "params"],
75
  "dispatch": {
76
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
77
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -99,7 +99,7 @@
99
  "cRank": "ranks.c",
100
  "op": "\"less\""
101
  },
102
- "bindings": ["a_2", "b_2", "c", "params"],
103
  "dispatch": {
104
  "x": "min(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
105
  "y": 1,
@@ -127,36 +127,7 @@
127
  "op": "\"less\"",
128
  "itemsPerInvocation": 4
129
  },
130
- "bindings": ["a_2", "b_2", "c_2", "params_2"],
131
- "dispatch": {
132
- "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
133
- "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
134
- "z": 1
135
- }
136
- }
137
- ]
138
- },
139
- {
140
- "id": "same_shape_scalar_x4",
141
- "priority": 15,
142
- "when": ["sameShape(shapes.a, shapes.c)", "sameShape(shapes.b, shapes.c)", "numel(shapes.c) > 0", "numel(shapes.c) % 4 != 0", "ranks.a <= ranks.c", "ranks.b <= ranks.c", "f16Ok(dtypes.T)"],
143
- "derive": { "scalar": "dtypes.T" },
144
- "passes": [
145
- {
146
- "id": "main",
147
- "name": "Less",
148
- "shader": "compare-broadcast.wgsl.jinja",
149
- "derive": {
150
- "aShape": "shapes.a",
151
- "bShape": "shapes.b",
152
- "cShape": "shapes.c",
153
- "aRank": "ranks.a",
154
- "bRank": "ranks.b",
155
- "cRank": "ranks.c",
156
- "op": "\"less\"",
157
- "itemsPerInvocation": 4
158
- },
159
- "bindings": ["a_2", "b_2", "c_2", "params_2"],
160
  "dispatch": {
161
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
162
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
@@ -167,17 +138,13 @@
167
  }
168
  ],
169
  "bindings": {
170
- "a": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
171
- "b": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
172
- "c": { "buffer": "storage", "elementType": "vec4<u32>" },
173
- "params": { "buffer": "uniform", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c) / 4" }] },
174
- "b_2": { "buffer": "read-only-storage", "name": "b", "elementType": "$scalar" },
175
- "a_2": { "buffer": "read-only-storage", "name": "a", "elementType": "$scalar" },
176
- "c_2": { "buffer": "storage", "name": "c", "elementType": "u32" },
177
- "params_2": {
178
- "buffer": "uniform",
179
- "name": "params",
180
- "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c)" }]
181
- }
182
  }
183
  }
 
17
  "passes": [
18
  {
19
  "id": "main",
20
+ "name": "Less.Vec4",
21
  "shader": "compare-vec4.wgsl.jinja",
22
  "derive": {
23
  "op": "\"less\"",
 
40
  "passes": [
41
  {
42
  "id": "main",
43
+ "name": "Less.ScalarBVec4",
44
  "shader": "compare-vec4.wgsl.jinja",
45
  "derive": {
46
  "op": "\"less\"",
47
  "scalarOperand": "\"b\"",
48
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
49
  },
50
+ "bindings": ["a", "b_scalar", "c", "params"],
51
  "dispatch": {
52
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
53
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
 
64
  "passes": [
65
  {
66
  "id": "main",
67
+ "name": "Less.ScalarAVec4",
68
  "shader": "compare-vec4.wgsl.jinja",
69
  "derive": {
70
  "op": "\"less\"",
71
  "scalarOperand": "\"a\"",
72
  "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1"
73
  },
74
+ "bindings": ["a_scalar", "b", "c", "params"],
75
  "dispatch": {
76
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
77
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)",
 
99
  "cRank": "ranks.c",
100
  "op": "\"less\""
101
  },
102
+ "bindings": ["a_scalar", "b_scalar", "c", "params"],
103
  "dispatch": {
104
  "x": "min(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
105
  "y": 1,
 
127
  "op": "\"less\"",
128
  "itemsPerInvocation": 4
129
  },
130
+ "bindings": ["a_scalar", "b_scalar", "c_u32", "params_count"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
131
  "dispatch": {
132
  "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
133
  "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)",
 
138
  }
139
  ],
140
  "bindings": {
141
+ "a": { "elementType": "$vectorScalar" },
142
+ "b": { "elementType": "$vectorScalar" },
143
+ "c": { "elementType": "vec4<u32>" },
144
+ "params": { "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c) / 4" }] },
145
+ "b_scalar": { "name": "b", "elementType": "$scalar" },
146
+ "a_scalar": { "name": "a", "elementType": "$scalar" },
147
+ "c_u32": { "name": "c", "elementType": "u32" },
148
+ "params_count": { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c)" }] }
 
 
 
 
149
  }
150
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.Less",
3
- "id": "_ai_onnx_less_webgpu_cd5aeae",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,23 +8,22 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "DYOpGhl+2huEW8ES9/klafDCuls+BKZycGYBsG3xuXc=",
11
- "compare-broadcast-vec4.wgsl.jinja": "e9IjvIN8zaZuy3Oxyo0917Z43virxHPxIKYg4GD5DDo=",
12
- "compare-broadcast.wgsl.jinja": "149Y9wl0draPoa3sCNPbGVSvEk3Z4G0Rt1lifzBbCFM=",
13
- "compare-vec4.wgsl.jinja": "HHor4TZ32hagFcdLFOoBRtFnH6CyrYjaIcviPAUhSWM=",
14
- "manifest.json": "IhvdJxWOC3OjvHtb6mjf/7GCfUAHPm1gIjVpfWwLKNk=",
15
- "test.json": "BkFQK6dgyaPaA0W1QSVZO1MlUw5afjqO9YwzV1lymrg="
16
  }
17
  },
18
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
19
  "webgpu": {
20
- "manifestSpec": "2.0",
21
  "variants": {
22
  "same_shape_vec4": ["compare-vec4.wgsl.jinja"],
23
  "scalar_b_vec4": ["compare-vec4.wgsl.jinja"],
24
  "scalar_a_vec4": ["compare-vec4.wgsl.jinja"],
25
  "broadcast_vec4": ["compare-broadcast-vec4.wgsl.jinja"],
26
- "broadcast": ["compare-broadcast.wgsl.jinja"],
27
- "same_shape_scalar_x4": ["compare-broadcast.wgsl.jinja"]
28
  }
29
  }
30
  }
 
1
  {
2
  "name": "ai.onnx.Less",
3
+ "id": "_ai_onnx_less_webgpu_f2e56a2",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "DYOpGhl+2huEW8ES9/klafDCuls+BKZycGYBsG3xuXc=",
11
+ "compare-broadcast-vec4.wgsl.jinja": "qT6u5x2OtFJGn0Ow8c9eLFwH2eE26gs59Rr1CPL872A=",
12
+ "compare-broadcast.wgsl.jinja": "q9DaIhQCKOZKNcuckQrvcMJmk/y5PBb0COT2m27i150=",
13
+ "compare-vec4.wgsl.jinja": "oqki3FpqohhfM1TEE/jfbrviCCgCREn+iEobDFzpDI4=",
14
+ "manifest.json": "aRWof2vs86t+VPgS3036BO3lzBg7hwsXR2ACfP0VCoA=",
15
+ "test.json": "JE4ijBYG4kV05JEJ37aa2XejGEq3auxLNhFYJB+z4JQ="
16
  }
17
  },
18
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
19
  "webgpu": {
20
+ "manifestSpec": "2.1",
21
  "variants": {
22
  "same_shape_vec4": ["compare-vec4.wgsl.jinja"],
23
  "scalar_b_vec4": ["compare-vec4.wgsl.jinja"],
24
  "scalar_a_vec4": ["compare-vec4.wgsl.jinja"],
25
  "broadcast_vec4": ["compare-broadcast-vec4.wgsl.jinja"],
26
+ "broadcast": ["compare-broadcast.wgsl.jinja"]
 
27
  }
28
  }
29
  }
build/webgpu/test.json CHANGED
@@ -531,7 +531,7 @@
531
  {
532
  "name": "scalar_b_vec4_route",
533
  "provenance": {
534
- "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. Values straddle and hit the scalar exactly, so the compare emits both results."
535
  },
536
  "inputs": {
537
  "a": {
@@ -546,7 +546,7 @@
546
  {
547
  "name": "scalar_a_vec4_route",
548
  "provenance": {
549
- "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. Values straddle and hit the scalar exactly, so the compare emits both results."
550
  },
551
  "inputs": {
552
  "a": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.25] } },
 
531
  {
532
  "name": "scalar_b_vec4_route",
533
  "provenance": {
534
+ "notes": "One scalar operand broadcasts across a vec4-aligned output; the other operand already matches the output shape. Values straddle and hit the scalar exactly, so the compare emits both results."
535
  },
536
  "inputs": {
537
  "a": {
 
546
  {
547
  "name": "scalar_a_vec4_route",
548
  "provenance": {
549
+ "notes": "One scalar operand broadcasts across a vec4-aligned output; the other operand already matches the output shape. Values straddle and hit the scalar exactly, so the compare emits both results."
550
  },
551
  "inputs": {
552
  "a": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.25] } },