Xenova HF Staff commited on
Commit
6b64365
·
verified ·
1 Parent(s): f0e6200

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -48,13 +48,13 @@ Attributes and default values (overridable per request):
48
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
49
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
50
  - [`test.json`](build/webgpu/test.json) — correctness cases
51
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
52
  - [`bitcast.wgsl.jinja`](build/webgpu/bitcast.wgsl.jinja)
53
 
54
  ## Use with `@huggingface/kernels`
55
 
56
  ```sh
57
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
58
  ```
59
 
60
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
@@ -68,7 +68,7 @@ Replace each `*Data` placeholder with a typed array containing the corresponding
68
  import { getKernel } from "@huggingface/kernels";
69
 
70
  const kernel = await getKernel("webgpu-kernels/ai.onnx.BitCast", { version: 1 });
71
- const { output } = await kernel({ input: { data: inputData, shape: [] } }, {
72
- attrs: { to: 6 },
73
  });
74
  ```
 
48
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
49
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
50
  - [`test.json`](build/webgpu/test.json) — correctness cases
51
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
52
  - [`bitcast.wgsl.jinja`](build/webgpu/bitcast.wgsl.jinja)
53
 
54
  ## Use with `@huggingface/kernels`
55
 
56
  ```sh
57
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
58
  ```
59
 
60
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
68
  import { getKernel } from "@huggingface/kernels";
69
 
70
  const kernel = await getKernel("webgpu-kernels/ai.onnx.BitCast", { version: 1 });
71
+ const { output } = await kernel({ input: { data: inputData, shape: [3] } }, {
72
+ attrs: { to: 1 },
73
  });
74
  ```
build/webgpu/bitcast.wgsl.jinja CHANGED
@@ -1,3 +1,11 @@
 
 
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  {% set vectorized = vectorizedSpec if vectorizedSpec is defined else false %}
@@ -6,12 +14,7 @@
6
  // their explicit low-byte masking and sign extension.
7
  @compute @workgroup_size({{ workgroupSizeSpec }})
8
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
9
- // The flat dispatch is folded across x/y at the device limit, so gid.y
10
- // carries the high portion of the slot index.
11
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ workgroupSizeSpec }}u;
12
- if (i >= params.count) {
13
- return;
14
- }
15
  {% if inScalar == outScalar %}
16
  output[i] = input[i];
17
  {% elif inputIsInt8 and outputIsUint8 %}
 
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
  {% set vectorized = vectorizedSpec if vectorizedSpec is defined else false %}
 
14
  // their explicit low-byte masking and sign extension.
15
  @compute @workgroup_size({{ workgroupSizeSpec }})
16
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
17
+ {{ flat_index_2d(workgroupSizeSpec) }}
 
 
 
 
 
18
  {% if inScalar == outScalar %}
19
  output[i] = input[i];
20
  {% elif inputIsInt8 and outputIsUint8 %}
build/webgpu/manifest.json CHANGED
@@ -35,7 +35,7 @@
35
  "passes": [
36
  {
37
  "id": "main",
38
- "name": "BitCast.vec4",
39
  "shader": "bitcast.wgsl.jinja",
40
  "derive": { "vectorizedSpec": true, "workgroupSizeSpec": "bitcastWorkgroupSize" },
41
  "bindings": [
 
35
  "passes": [
36
  {
37
  "id": "main",
38
+ "name": "BitCast.Vec4",
39
  "shader": "bitcast.wgsl.jinja",
40
  "derive": { "vectorizedSpec": true, "workgroupSizeSpec": "bitcastWorkgroupSize" },
41
  "bindings": [
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.BitCast",
3
- "id": "_ai_onnx_bitcast_webgpu_ab905c6",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,14 +8,14 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "7aNcJ8LB0eyjvFh7jXkwAw1p2vgZejuWw6H9iwqbgqY=",
11
- "bitcast.wgsl.jinja": "1vfOUaYUxSx36bWOnKDPyXUPGprDxhPVEv10804xfVI=",
12
- "manifest.json": "UplDa7horbJakWMkW0Gkocun6e1kcewHPZAvi6weNHs=",
13
- "test.json": "FXueeEbCTZUCVYOfr0NjJrQAoVhhvtwqADXKZI6829c="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
  "variants": { "slot32_vec4": ["bitcast.wgsl.jinja"], "slot32": ["bitcast.wgsl.jinja"] }
20
  }
21
  }
 
1
  {
2
  "name": "ai.onnx.BitCast",
3
+ "id": "_ai_onnx_bitcast_webgpu_31d09ea",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "7aNcJ8LB0eyjvFh7jXkwAw1p2vgZejuWw6H9iwqbgqY=",
11
+ "bitcast.wgsl.jinja": "ibrqjbQlHMWRA/FhpRHG1KZup5FFZtRSRIUG9hTcXiQ=",
12
+ "manifest.json": "h6XBC5d/eqbX9wJxoyVa7PiJ7aNXwEsZCXr91+F9F0c=",
13
+ "test.json": "TQvH4fE/V77N+0mk4mwR5PxcXsRctMnGBvn1S0LYlqQ="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
  "webgpu": {
18
+ "manifestSpec": "2.1",
19
  "variants": { "slot32_vec4": ["bitcast.wgsl.jinja"], "slot32": ["bitcast.wgsl.jinja"] }
20
  }
21
  }
build/webgpu/test.json CHANGED
@@ -314,7 +314,7 @@
314
  {
315
  "name": "float32_to_int32_rank7",
316
  "provenance": {
317
- "source": "ONNX BitCast-26 contract and ONNX Runtime CPUExecutionProvider",
318
  "notes": "BitCast preserves arbitrary tensor rank and only reinterprets same-width element bits."
319
  },
320
  "attrs": { "to": 6 },
 
314
  {
315
  "name": "float32_to_int32_rank7",
316
  "provenance": {
317
+ "source": "ONNX BitCast-26 contract and ONNX Runtime's CPU provider",
318
  "notes": "BitCast preserves arbitrary tensor rank and only reinterprets same-width element bits."
319
  },
320
  "attrs": { "to": 6 },