Xenova HF Staff commited on
Commit
7d3dd73
·
verified ·
1 Parent(s): 1238ce8

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -49,7 +49,7 @@ Attributes and default values (overridable per request):
49
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
50
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
51
  - [`test.json`](build/webgpu/test.json) — correctness cases
52
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
53
  - [`eyelike-clear-vec4.wgsl.jinja`](build/webgpu/eyelike-clear-vec4.wgsl.jinja)
54
  - [`eyelike-diagonal.wgsl.jinja`](build/webgpu/eyelike-diagonal.wgsl.jinja)
55
  - [`eyelike.wgsl.jinja`](build/webgpu/eyelike.wgsl.jinja)
@@ -57,7 +57,7 @@ Attributes and default values (overridable per request):
57
  ## Use with `@huggingface/kernels`
58
 
59
  ```sh
60
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
61
  ```
62
 
63
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
49
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
50
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
51
  - [`test.json`](build/webgpu/test.json) — correctness cases
52
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
53
  - [`eyelike-clear-vec4.wgsl.jinja`](build/webgpu/eyelike-clear-vec4.wgsl.jinja)
54
  - [`eyelike-diagonal.wgsl.jinja`](build/webgpu/eyelike-diagonal.wgsl.jinja)
55
  - [`eyelike.wgsl.jinja`](build/webgpu/eyelike.wgsl.jinja)
 
57
  ## Use with `@huggingface/kernels`
58
 
59
  ```sh
60
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
61
  ```
62
 
63
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/bench.json CHANGED
@@ -6,7 +6,7 @@
6
  "outputs": { "output": { "dtype": "float32", "shape": [2048, 2048] } }
7
  },
8
  {
9
- "name": "eyelike-f32-4097-square-tail-healthy",
10
  "preset": "smoke",
11
  "vars": { "dtype": "float32", "count": 16785409 },
12
  "inputs": { "input": { "dtype": "float32", "shape": [4097, 4097] } },
 
6
  "outputs": { "output": { "dtype": "float32", "shape": [2048, 2048] } }
7
  },
8
  {
9
+ "name": "eyelike-f32-4097-square-tail-control",
10
  "preset": "smoke",
11
  "vars": { "dtype": "float32", "count": 16785409 },
12
  "inputs": { "input": { "dtype": "float32", "shape": [4097, 4097] } },
build/webgpu/eyelike-clear-vec4.wgsl.jinja CHANGED
@@ -1,36 +1,8 @@
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
  const COUNT4: u32 = {{ count4 }}u;
@@ -40,7 +12,7 @@ const COUNT: u32 = {{ count }}u;
40
  {% endif %}
41
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
42
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
43
- {{ flat_index_2d("q", "", note="") }}
44
  if (q < COUNT4) {
45
  {% if tailSafe %}
46
  let base = q * 4u;
 
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 }};{% endmacro %}
 
 
 
 
 
 
 
 
 
 
 
 
 
6
  {{ env.wgsl.resourceDeclarations }}
7
 
8
  const COUNT4: u32 = {{ count4 }}u;
 
12
  {% endif %}
13
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
14
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
15
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "q", "") }}
16
  if (q < COUNT4) {
17
  {% if tailSafe %}
18
  let base = q * 4u;
build/webgpu/eyelike-diagonal.wgsl.jinja CHANGED
@@ -1,43 +1,19 @@
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
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
37
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
38
- {{ flat_index_2d("row", "params.rows", guardInline=true, note="") }}
39
- let col = i32(row) + params.k;
40
- if (col >= 0 && col < i32(params.cols)) {
41
- output[row * params.cols + u32(col)] = {{ outScalar }}(1);
 
 
 
42
  }
 
43
  }
 
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 }};{% endmacro %}
 
 
 
 
 
 
 
 
 
 
 
 
 
6
  {{ env.wgsl.resourceDeclarations }}
7
 
8
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
9
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
10
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE, "lane", "") }}
11
+ // Only the rows the diagonal meets are dispatched, [max(0, -k), min(rows, cols - k)), so
12
+ // lane 0 starts on the first of them and col = row + k is never negative.
13
+ let row = i32(lane) + max(0, -params.k);
14
+ let col = row + params.k;
15
+ if (row >= i32(params.rows) || col >= i32(params.cols)) {
16
+ return;
17
  }
18
+ output[u32(row) * params.cols + u32(col)] = {{ outScalar }}(1);
19
  }
build/webgpu/eyelike.wgsl.jinja CHANGED
@@ -1,41 +1,16 @@
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
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
37
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
38
- {{ flat_index_2d() }}
39
  let row = i32(i / params.cols);
40
  let col = i32(i % params.cols);
41
  output[i] = select({{ outScalar }}(0), {{ outScalar }}(1), col - row == params.k);
 
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
  @compute @workgroup_size({{ tunables.WORKGROUP_SIZE }})
12
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
13
+ {{ flat_index_2d(tunables.WORKGROUP_SIZE) }}
14
  let row = i32(i / params.cols);
15
  let col = i32(i % params.cols);
16
  output[i] = select({{ outScalar }}(0), {{ outScalar }}(1), col - row == params.k);
build/webgpu/manifest.json CHANGED
@@ -14,8 +14,9 @@
14
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
15
  "reportedNonWave32Adapter": "not wave32Adapter and (has(device.adapterInfo, \"subgroupMinSize\") or has(device.adapterInfo, \"subgroupMaxSize\"))",
16
  "dtypeContract": "(has(attrs, \"dtype\") and attrs.dtype == onnxDtypeCode(logicalDtypes.T2)) or (not has(attrs, \"dtype\") and logicalDtypes.T2 == logicalDtypes.T1)",
17
- "packedClearPreferred": "numel(shapes.output) % 4 == 0 or not reportedNonWave32Adapter",
18
- "outScalar": "dtypes.T2"
 
19
  },
20
  "when": ["dtypeContract", "ranks.input == 2", "ranks.output == 2", "dim(shapes.output, 0) == dim(shapes.input, 0)", "dim(shapes.output, 1) == dim(shapes.input, 1)", "f16Ok(tensorDtypes.output)"],
21
  "variants": [
@@ -26,7 +27,8 @@
26
  "derive": {
27
  "tailSafe": "numel(shapes.output) % 4 != 0",
28
  "outVector": "\"vec4<\" ~ dtypes.T2 ~ \">\"",
29
- "clearElement": "dtypes.T2 if numel(shapes.output) % 4 != 0 else (\"vec4<\" ~ dtypes.T2 ~ \">\")"
 
30
  },
31
  "passes": [
32
  {
@@ -52,17 +54,13 @@
52
  "struct": [
53
  { "name": "rows", "type": "u32", "value": "dim(shapes.output, 0)" },
54
  { "name": "cols", "type": "u32", "value": "dim(shapes.output, 1)" },
55
- {
56
- "name": "k",
57
- "type": "i32",
58
- "value": "max(min(attrs.k, dim(shapes.output, 1)), 0 - dim(shapes.output, 0))"
59
- }
60
  ]
61
  }
62
  ],
63
  "dispatch": {
64
- "x": "min(ceilDiv((dim(shapes.output, 0)), (tunables.WORKGROUP_SIZE)), 65535)",
65
- "y": "ceilDiv(ceilDiv((dim(shapes.output, 0)), (tunables.WORKGROUP_SIZE)), 65535)",
66
  "z": 1
67
  }
68
  }
@@ -82,11 +80,7 @@
82
  "struct": [
83
  { "name": "count", "type": "u32", "value": "numel(shapes.output)" },
84
  { "name": "cols", "type": "u32", "value": "dim(shapes.output, 1)" },
85
- {
86
- "name": "k",
87
- "type": "i32",
88
- "value": "max(min(attrs.k, dim(shapes.output, 1)), 0 - dim(shapes.output, 0))"
89
- }
90
  ]
91
  }
92
  ],
 
14
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
15
  "reportedNonWave32Adapter": "not wave32Adapter and (has(device.adapterInfo, \"subgroupMinSize\") or has(device.adapterInfo, \"subgroupMaxSize\"))",
16
  "dtypeContract": "(has(attrs, \"dtype\") and attrs.dtype == onnxDtypeCode(logicalDtypes.T2)) or (not has(attrs, \"dtype\") and logicalDtypes.T2 == logicalDtypes.T1)",
17
+ "packedClearPreferred": "numel(shapes.output) % 4 == 0 or not (reportedNonWave32Adapter and device.features.has(\"subgroups\"))",
18
+ "outScalar": "dtypes.T2",
19
+ "diagonalK": "max(min(attrs.k, dim(shapes.output, 1)), 0 - dim(shapes.output, 0))"
20
  },
21
  "when": ["dtypeContract", "ranks.input == 2", "ranks.output == 2", "dim(shapes.output, 0) == dim(shapes.input, 0)", "dim(shapes.output, 1) == dim(shapes.input, 1)", "f16Ok(tensorDtypes.output)"],
22
  "variants": [
 
27
  "derive": {
28
  "tailSafe": "numel(shapes.output) % 4 != 0",
29
  "outVector": "\"vec4<\" ~ dtypes.T2 ~ \">\"",
30
+ "clearElement": "dtypes.T2 if numel(shapes.output) % 4 != 0 else (\"vec4<\" ~ dtypes.T2 ~ \">\")",
31
+ "diagonalRows": "max(0, min(dim(shapes.output, 0), dim(shapes.output, 1) - diagonalK) - max(0, 0 - diagonalK))"
32
  },
33
  "passes": [
34
  {
 
54
  "struct": [
55
  { "name": "rows", "type": "u32", "value": "dim(shapes.output, 0)" },
56
  { "name": "cols", "type": "u32", "value": "dim(shapes.output, 1)" },
57
+ { "name": "k", "type": "i32", "value": "diagonalK" }
 
 
 
 
58
  ]
59
  }
60
  ],
61
  "dispatch": {
62
+ "x": "min(ceilDiv((max(1, diagonalRows)), (tunables.WORKGROUP_SIZE)), 65535)",
63
+ "y": "ceilDiv(ceilDiv((max(1, diagonalRows)), (tunables.WORKGROUP_SIZE)), 65535)",
64
  "z": 1
65
  }
66
  }
 
80
  "struct": [
81
  { "name": "count", "type": "u32", "value": "numel(shapes.output)" },
82
  { "name": "cols", "type": "u32", "value": "dim(shapes.output, 1)" },
83
+ { "name": "k", "type": "i32", "value": "diagonalK" }
 
 
 
 
84
  ]
85
  }
86
  ],
build/webgpu/metadata.json CHANGED
@@ -1,23 +1,23 @@
1
  {
2
  "name": "ai.onnx.EyeLike",
3
- "id": "_ai_onnx_eyelike_webgpu_a68b51c",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "7iac6y3brK8RDhh+aJr8jCIzFVUfuKX1UXqRT+QXc2Q=",
11
- "eyelike-clear-vec4.wgsl.jinja": "23ASRrA6Hz2EJZfucMYwY0OqHWQbzRcfDDPk16owtjM=",
12
- "eyelike-diagonal.wgsl.jinja": "lP7UaipIa60v3nfrOigD0jjr1tPtf3iaVskNyw4PCgY=",
13
- "eyelike.wgsl.jinja": "SzeY2UkJqE/aFOH9I0JMqRWUuazCPy6f7/cAQKAEybI=",
14
- "manifest.json": "+Z0Tu+zA0k5ppZNM1g+cAyv+xbq9rfyW0rPO5NLevlg=",
15
  "test.json": "KG4MR6xICXdAoULdHEiJ6KbYoE5TxHpz/6Aj2t59qfk="
16
  }
17
  },
18
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
19
  "webgpu": {
20
- "manifestSpec": "2.0",
21
  "variants": {
22
  "rank2_vec4": ["eyelike-clear-vec4.wgsl.jinja", "eyelike-diagonal.wgsl.jinja"],
23
  "rank2": ["eyelike.wgsl.jinja"]
 
1
  {
2
  "name": "ai.onnx.EyeLike",
3
+ "id": "_ai_onnx_eyelike_webgpu_2f86257",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "JQ/JUgD9zXqUA7m9lgKBtVIS+qgjcrP5+DUwc2oc73I=",
11
+ "eyelike-clear-vec4.wgsl.jinja": "UZF0MVFhtM7htAwmcaz0T/dK8Cklzu6SuJUb2FIoD20=",
12
+ "eyelike-diagonal.wgsl.jinja": "lhT517EU9eozeFmvsA3beETxwXFdlbbGaKN3QGYZFOs=",
13
+ "eyelike.wgsl.jinja": "iMfiIGW74i4s5GpFox1RcHVCEDqV35MHLAC0lvRhyPU=",
14
+ "manifest.json": "+QdLObVKzKnn2kVV+1zGDEajqNrJpESRaCiPrkVLgyk=",
15
  "test.json": "KG4MR6xICXdAoULdHEiJ6KbYoE5TxHpz/6Aj2t59qfk="
16
  }
17
  },
18
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
19
  "webgpu": {
20
+ "manifestSpec": "2.1",
21
  "variants": {
22
  "rank2_vec4": ["eyelike-clear-vec4.wgsl.jinja", "eyelike-diagonal.wgsl.jinja"],
23
  "rank2": ["eyelike.wgsl.jinja"]