sync 6fdf6301e2bb
Browse files- README.md +8 -2
- build/webgpu/gemma-rotary-embedding.wgsl.jinja +9 -7
- build/webgpu/manifest.json +17 -56
- build/webgpu/metadata.json +6 -9
- build/webgpu/test.json +1 -1
README.md
CHANGED
|
@@ -40,6 +40,12 @@ See the [ONNX Runtime `GemmaRotaryEmbedding` contrib-operator spec](https://gith
|
|
| 40 |
| `T` | `float16` |
|
| 41 |
| `U` | `float32` |
|
| 42 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 43 |
## Device requirements
|
| 44 |
|
| 45 |
Every implementation variant requires `shader-f16`; the package has no variant-level fallback without that capability.
|
|
@@ -49,13 +55,13 @@ Every implementation variant requires `shader-f16`; the package has no variant-l
|
|
| 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
|
| 53 |
- [`gemma-rotary-embedding.wgsl.jinja`](build/webgpu/gemma-rotary-embedding.wgsl.jinja)
|
| 54 |
|
| 55 |
## Use with `@huggingface/kernels`
|
| 56 |
|
| 57 |
```sh
|
| 58 |
-
npm install --save-exact @huggingface/kernels@0.0.1-preview.
|
| 59 |
```
|
| 60 |
|
| 61 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
|
|
|
| 40 |
| `T` | `float16` |
|
| 41 |
| `U` | `float32` |
|
| 42 |
|
| 43 |
+
## Implementation variants
|
| 44 |
+
|
| 45 |
+
One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
|
| 46 |
+
|
| 47 |
+
- `elementwise` — Elementwise rotation. Streams `vec4` words when each head's `(seq, dim)` plane holds a multiple of four elements, so no word straddles two heads; scalar elements otherwise.
|
| 48 |
+
|
| 49 |
## Device requirements
|
| 50 |
|
| 51 |
Every implementation variant requires `shader-f16`; the package has no variant-level fallback without that capability.
|
|
|
|
| 55 |
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
|
| 56 |
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
|
| 57 |
- [`test.json`](build/webgpu/test.json) — correctness cases
|
| 58 |
+
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
|
| 59 |
- [`gemma-rotary-embedding.wgsl.jinja`](build/webgpu/gemma-rotary-embedding.wgsl.jinja)
|
| 60 |
|
| 61 |
## Use with `@huggingface/kernels`
|
| 62 |
|
| 63 |
```sh
|
| 64 |
+
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
|
| 65 |
```
|
| 66 |
|
| 67 |
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
|
build/webgpu/gemma-rotary-embedding.wgsl.jinja
CHANGED
|
@@ -1,4 +1,11 @@
|
|
| 1 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2 |
{{ env.wgsl.resourceDeclarations }}
|
| 3 |
|
| 4 |
// output1 = q * cos(emb) + q_rot * sin(emb); output2 applies the same
|
|
@@ -13,12 +20,7 @@ const ZERO: {{ scalar }} = {{ scalar }}(0.0);
|
|
| 13 |
|
| 14 |
@compute @workgroup_size(WG, 1, 1)
|
| 15 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 16 |
-
|
| 17 |
-
// per-axis dispatch fold width. Reduces to gid.x when the dispatch does not fold.
|
| 18 |
-
let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
|
| 19 |
-
if (index >= params.count) {
|
| 20 |
-
return;
|
| 21 |
-
}
|
| 22 |
// (batch, head, seq, dim) -> (batch, seq, dim): divide out the head-major stride and
|
| 23 |
// keep the remainder, which is exactly the (seq, dim) offset shared by every head.
|
| 24 |
let emb_index = (index / params.headSeqDim) * params.seqDim + index % params.seqDim;
|
|
|
|
| 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 |
// output1 = q * cos(emb) + q_rot * sin(emb); output2 applies the same
|
|
|
|
| 20 |
|
| 21 |
@compute @workgroup_size(WG, 1, 1)
|
| 22 |
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 23 |
+
{{ flat_index_2d("WG", "index") }}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 24 |
// (batch, head, seq, dim) -> (batch, seq, dim): divide out the head-major stride and
|
| 25 |
// keep the remainder, which is exactly the (seq, dim) offset shared by every head.
|
| 26 |
let emb_index = (index / params.headSeqDim) * params.seqDim + index % params.seqDim;
|
build/webgpu/manifest.json
CHANGED
|
@@ -28,54 +28,15 @@
|
|
| 28 |
"when": ["contract", "tunables.workgroupSize >= 1", "floor(tunables.workgroupSize) == tunables.workgroupSize", "tunables.workgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.workgroupSize <= device.limits.maxComputeWorkgroupSizeX"],
|
| 29 |
"variants": [
|
| 30 |
{
|
| 31 |
-
"id": "
|
| 32 |
-
"priority": 10,
|
| 33 |
-
"when": ["vec4Ok"],
|
| 34 |
"requires": { "features": ["shader-f16"] },
|
| 35 |
"derive": {
|
| 36 |
-
"vec4":
|
| 37 |
-
"
|
| 38 |
-
"vector": "\"vec4<f16>\"",
|
| 39 |
-
"workgroupSize": "tunables.workgroupSize"
|
| 40 |
-
},
|
| 41 |
-
"passes": [
|
| 42 |
-
{
|
| 43 |
-
"id": "main",
|
| 44 |
-
"name": "GemmaRotaryEmbedding.Vec4",
|
| 45 |
-
"shader": "gemma-rotary-embedding.wgsl.jinja",
|
| 46 |
-
"bindings": [
|
| 47 |
-
{ "arg": "embT", "name": "emb", "elementType": "vec4<f32>" },
|
| 48 |
-
{ "arg": "qT", "name": "q", "elementType": "$vector" },
|
| 49 |
-
{ "arg": "qRotT", "name": "q_rot", "elementType": "$vector" },
|
| 50 |
-
{ "arg": "kT", "name": "k", "elementType": "$vector" },
|
| 51 |
-
{ "arg": "kRotT", "name": "k_rot", "elementType": "$vector" },
|
| 52 |
-
{ "arg": "output1T", "name": "output1", "elementType": "$vector" },
|
| 53 |
-
{ "arg": "output2T", "name": "output2", "elementType": "$vector" },
|
| 54 |
-
{
|
| 55 |
-
"name": "params",
|
| 56 |
-
"struct": [
|
| 57 |
-
{ "name": "count", "type": "u32", "value": "numel(shapes.qT) / 4" },
|
| 58 |
-
{ "name": "seqDim", "type": "u32", "value": "(seqLen * headDim) / 4" },
|
| 59 |
-
{ "name": "headSeqDim", "type": "u32", "value": "(numHeads * seqLen * headDim) / 4" }
|
| 60 |
-
]
|
| 61 |
-
}
|
| 62 |
-
],
|
| 63 |
-
"dispatch": {
|
| 64 |
-
"x": "min(ceilDiv((numel(shapes.qT) / 4), (workgroupSize)), 65535)",
|
| 65 |
-
"y": "ceilDiv(ceilDiv((numel(shapes.qT) / 4), (workgroupSize)), 65535)",
|
| 66 |
-
"z": 1
|
| 67 |
-
}
|
| 68 |
-
}
|
| 69 |
-
]
|
| 70 |
-
},
|
| 71 |
-
{
|
| 72 |
-
"id": "scalar",
|
| 73 |
-
"priority": 0,
|
| 74 |
-
"requires": { "features": ["shader-f16"] },
|
| 75 |
-
"derive": {
|
| 76 |
-
"vec4": false,
|
| 77 |
"scalar": "dtypes.T",
|
| 78 |
"vector": "\"vec4<f16>\"",
|
|
|
|
|
|
|
| 79 |
"workgroupSize": "tunables.workgroupSize"
|
| 80 |
},
|
| 81 |
"passes": [
|
|
@@ -84,25 +45,25 @@
|
|
| 84 |
"name": "GemmaRotaryEmbedding",
|
| 85 |
"shader": "gemma-rotary-embedding.wgsl.jinja",
|
| 86 |
"bindings": [
|
| 87 |
-
{ "arg": "embT", "name": "emb" },
|
| 88 |
-
{ "arg": "qT", "name": "q", "elementType": "$
|
| 89 |
-
{ "arg": "qRotT", "name": "q_rot", "elementType": "$
|
| 90 |
-
{ "arg": "kT", "name": "k", "elementType": "$
|
| 91 |
-
{ "arg": "kRotT", "name": "k_rot", "elementType": "$
|
| 92 |
-
{ "arg": "output1T", "name": "output1", "elementType": "$
|
| 93 |
-
{ "arg": "output2T", "name": "output2", "elementType": "$
|
| 94 |
{
|
| 95 |
"name": "params",
|
| 96 |
"struct": [
|
| 97 |
-
{ "name": "count", "type": "u32", "value": "numel(shapes.qT)" },
|
| 98 |
-
{ "name": "seqDim", "type": "u32", "value": "(seqLen * headDim)" },
|
| 99 |
-
{ "name": "headSeqDim", "type": "u32", "value": "(numHeads * seqLen * headDim)" }
|
| 100 |
]
|
| 101 |
}
|
| 102 |
],
|
| 103 |
"dispatch": {
|
| 104 |
-
"x": "min(ceilDiv((max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
|
| 105 |
-
"y": "ceilDiv(ceilDiv((max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
|
| 106 |
"z": 1
|
| 107 |
}
|
| 108 |
}
|
|
|
|
| 28 |
"when": ["contract", "tunables.workgroupSize >= 1", "floor(tunables.workgroupSize) == tunables.workgroupSize", "tunables.workgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.workgroupSize <= device.limits.maxComputeWorkgroupSizeX"],
|
| 29 |
"variants": [
|
| 30 |
{
|
| 31 |
+
"id": "elementwise",
|
|
|
|
|
|
|
| 32 |
"requires": { "features": ["shader-f16"] },
|
| 33 |
"derive": {
|
| 34 |
+
"vec4": "vec4Ok",
|
| 35 |
+
"lanes": "4 if vec4Ok else 1",
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
"scalar": "dtypes.T",
|
| 37 |
"vector": "\"vec4<f16>\"",
|
| 38 |
+
"element": "vector if vec4Ok else scalar",
|
| 39 |
+
"embElement": "\"vec4<f32>\" if vec4Ok else \"f32\"",
|
| 40 |
"workgroupSize": "tunables.workgroupSize"
|
| 41 |
},
|
| 42 |
"passes": [
|
|
|
|
| 45 |
"name": "GemmaRotaryEmbedding",
|
| 46 |
"shader": "gemma-rotary-embedding.wgsl.jinja",
|
| 47 |
"bindings": [
|
| 48 |
+
{ "arg": "embT", "name": "emb", "elementType": "$embElement" },
|
| 49 |
+
{ "arg": "qT", "name": "q", "elementType": "$element" },
|
| 50 |
+
{ "arg": "qRotT", "name": "q_rot", "elementType": "$element" },
|
| 51 |
+
{ "arg": "kT", "name": "k", "elementType": "$element" },
|
| 52 |
+
{ "arg": "kRotT", "name": "k_rot", "elementType": "$element" },
|
| 53 |
+
{ "arg": "output1T", "name": "output1", "elementType": "$element" },
|
| 54 |
+
{ "arg": "output2T", "name": "output2", "elementType": "$element" },
|
| 55 |
{
|
| 56 |
"name": "params",
|
| 57 |
"struct": [
|
| 58 |
+
{ "name": "count", "type": "u32", "value": "numel(shapes.qT) / lanes" },
|
| 59 |
+
{ "name": "seqDim", "type": "u32", "value": "(seqLen * headDim) / lanes" },
|
| 60 |
+
{ "name": "headSeqDim", "type": "u32", "value": "(numHeads * seqLen * headDim) / lanes" }
|
| 61 |
]
|
| 62 |
}
|
| 63 |
],
|
| 64 |
"dispatch": {
|
| 65 |
+
"x": "min(ceilDiv((numel(shapes.qT) / 4 if vec4Ok else max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
|
| 66 |
+
"y": "ceilDiv(ceilDiv((numel(shapes.qT) / 4 if vec4Ok else max(1, numel(shapes.qT))), (workgroupSize)), 65535)",
|
| 67 |
"z": 1
|
| 68 |
}
|
| 69 |
}
|
build/webgpu/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.GemmaRotaryEmbedding",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
@@ -8,14 +8,11 @@
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "E1026CMMv9Pppl2+JpGoj3dWCSmk8tZUEwBYo3UCiD4=",
|
| 11 |
-
"gemma-rotary-embedding.wgsl.jinja": "
|
| 12 |
-
"manifest.json": "
|
| 13 |
-
"test.json": "
|
| 14 |
}
|
| 15 |
},
|
| 16 |
-
"provenance": { "kernel": { "sha": "
|
| 17 |
-
"webgpu": {
|
| 18 |
-
"manifestSpec": "2.0",
|
| 19 |
-
"variants": { "vec4": ["gemma-rotary-embedding.wgsl.jinja"], "scalar": ["gemma-rotary-embedding.wgsl.jinja"] }
|
| 20 |
-
}
|
| 21 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "com.microsoft.GemmaRotaryEmbedding",
|
| 3 |
+
"id": "_com_microsoft_gemmarotaryembedding_webgpu_bb6bdf4",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
|
|
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
"bench.json": "E1026CMMv9Pppl2+JpGoj3dWCSmk8tZUEwBYo3UCiD4=",
|
| 11 |
+
"gemma-rotary-embedding.wgsl.jinja": "J7OaLWdNRVSQJN7334F9gTgV5OrdWUywfjXx5rPnHr4=",
|
| 12 |
+
"manifest.json": "QJJdYbfJe3E3SRzXSXLz3IuLFfsthXglsTWLtCQWh4U=",
|
| 13 |
+
"test.json": "0qR3IsgdZC/lwA1WbH1QDFYTijEjsyo3t++UOJnt6VY="
|
| 14 |
}
|
| 15 |
},
|
| 16 |
+
"provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
|
| 17 |
+
"webgpu": { "manifestSpec": "2.1", "variants": { "elementwise": ["gemma-rotary-embedding.wgsl.jinja"] } }
|
|
|
|
|
|
|
|
|
|
| 18 |
}
|
build/webgpu/test.json
CHANGED
|
@@ -201,7 +201,7 @@
|
|
| 201 |
{
|
| 202 |
"name": "f16_vec4_dim6_seq4_straddles_heads",
|
| 203 |
"provenance": {
|
| 204 |
-
"notes": "
|
| 205 |
},
|
| 206 |
"inputs": {
|
| 207 |
"embT": {
|
|
|
|
| 201 |
{
|
| 202 |
"name": "f16_vec4_dim6_seq4_straddles_heads",
|
| 203 |
"provenance": {
|
| 204 |
+
"notes": "Head dim 6 is not a multiple of four, but seq_len x head_dim = 4x6 = 24 is; checks four-element grouping that straddles head-dim boundaries within one (batch, seq) block rather than aligning to them."
|
| 205 |
},
|
| 206 |
"inputs": {
|
| 207 |
"embT": {
|