Xenova HF Staff commited on
Commit
f7870ad
·
verified ·
1 Parent(s): 78115a9

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -1,3 +1,91 @@
1
  ---
 
2
  license: apache-2.0
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ library_name: kernels
3
  license: apache-2.0
4
+ tags:
5
+ - kernel
6
+ - webgpu
7
+ - wgsl
8
  ---
9
+ # com.microsoft.NGramHashMapping
10
+
11
+ `com.microsoft` · ONNX Runtime contrib operator · contrib since_version 1
12
+
13
+ ## Description
14
+
15
+ Engram n-gram hash ids from compressed tokenizer ids. For every order `n` in `[2, max_ngram_size]` it mixes the `n` causal shifts of `input_ids` as `mix = ids[t] * multipliers[0] xor ... xor ids[t-n+1] * multipliers[n-1]`, the products wrapping in two's complement, then emits `mix` modulo each of that order's head vocabulary sizes. `past_ids`/`present_ids` carry the `max_ngram_size - 1` preceding ids, so chunked prefill and decode agree with one full-sequence run. Only `int32` ids are implemented: `int64` is a different answer, not a wider one, since its products wrap at another width.
16
+
17
+ See the [ONNX Runtime `NGramHashMapping` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.NGramHashMapping) for the reference semantics.
18
+
19
+ ## Inputs
20
+
21
+ | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
22
+ | --- | --- | --- | --- | --- | --- | --- |
23
+ | `inputIdsT` | `input_ids` | `M` | `2` | — | Compressed tokenizer ids with shape `(batch_size, sequence_length)`. A zero-length sequence is accepted: it emits no hash ids and passes the carry state through unchanged. | required |
24
+ | `multipliersT` | `multipliers` | `M` | `1` | — | Per-shift hash multipliers with shape `(max_ngram_size)`. Conventionally odd, but any value is accepted; the product wraps in two's complement rather than saturating. | required |
25
+ | `vocabSizesT` | `vocab_sizes` | `M` | `1` | — | Per-output-head vocabulary sizes with shape `((max_ngram_size - 1) * n_head_per_ngram)`, conventionally prime and strictly positive. A non-positive entry has no meaningful modulo; this kernel emits a hash id of `0` for that head, which is what every GPU implementation of the operator does, while the reference CPU implementation rejects it instead. | required |
26
+ | `pastIdsT` | `past_ids` | `M` | `2` | — | Optional ids for the `max_ngram_size - 1` positions preceding this call, with shape `(batch_size, max_ngram_size - 1)`. Right-aligned, so the last slot is the most recent id. When it is absent the missing history is `pad_id`, which is what a fresh sequence means. | optional |
27
+
28
+ ## Outputs
29
+
30
+ | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
31
+ | --- | --- | --- | --- | --- | --- | --- |
32
+ | `hashIdsT` | `hash_ids` | `M` | `3` | derived | Hash ids with shape `(batch_size, sequence_length, (max_ngram_size - 1) * n_head_per_ngram)`. The heads of order `n = 2` come first, then `n = 3`, and so on. | required |
33
+ | `presentIdsT` | `present_ids` | `M` | `2` | derived | The trailing `max_ngram_size - 1` ids of `past_ids` followed by `input_ids`, with shape `(batch_size, max_ngram_size - 1)`. Feed it back as `past_ids` on the next call. It is always written, into a buffer distinct from `past_ids`. | required |
34
+
35
+ ## Attributes
36
+
37
+ Attributes and default values (overridable per request):
38
+
39
+ | Attribute | Default | Description |
40
+ | --- | --- | --- |
41
+ | `max_ngram_size` | — | Highest n-gram order, at least 2. It is baked into the rendered kernel, so the shift walk and the head arithmetic are compile-time; this package renders orders up to 16. |
42
+ | `n_head_per_ngram` | — | Number of hash heads emitted for each n-gram order, at least 1. It is baked into the rendered kernel alongside `max_ngram_size`, and their product with `max_ngram_size - 1` is capped at 64 heads. |
43
+ | `pad_id` | — | Compressed tokenizer id used for causal-shift positions before the beginning of the whole sequence. It must be representable as `int32`. |
44
+
45
+ ## Type constraints
46
+
47
+ | Variable | Allowed dtypes |
48
+ | --- | --- |
49
+ | `M` | `int32` |
50
+
51
+ ## Implementation variants
52
+
53
+ One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
54
+
55
+ - `past_shared_prefix` — Stage each causal prefix once per token in workgroup memory, then distribute contiguous head outputs across lanes. The tile obeys device storage and invocation limits. Small combined coefficient tables retain the compact token-owned path to avoid staging and barrier overhead.
56
+ - `fresh_shared_prefix` — Stage each causal prefix once per token in workgroup memory, then distribute contiguous head outputs across lanes. The tile obeys device storage and invocation limits. Small combined coefficient tables retain the compact token-owned path to avoid staging and barrier overhead.
57
+
58
+ ## Files
59
+
60
+ - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
61
+ - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
62
+ - [`test.json`](build/webgpu/test.json) — correctness cases
63
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
64
+ - [`ngram-hash-mapping-shared-prefix.wgsl.jinja`](build/webgpu/ngram-hash-mapping-shared-prefix.wgsl.jinja)
65
+ - [`ngram-hash-mapping.wgsl.jinja`](build/webgpu/ngram-hash-mapping.wgsl.jinja)
66
+
67
+ ## Use with `@huggingface/kernels`
68
+
69
+ ```sh
70
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
71
+ ```
72
+
73
+ Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
74
+
75
+ The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
76
+ It follows the `v1` branch as fixes land. To pin exact artifact bytes, pass a 40-character commit `revision` instead of `version`.
77
+
78
+ Replace each `*Data` placeholder with a typed array containing the corresponding input data.
79
+
80
+ ```js
81
+ import { getKernel } from "@huggingface/kernels";
82
+
83
+ const kernel = await getKernel("webgpu-kernels/com.microsoft.NGramHashMapping", { version: 1 });
84
+ const { hashIdsT, presentIdsT } = await kernel({
85
+ inputIdsT: { data: inputIdsTData, shape: [1, 1] },
86
+ multipliersT: { data: multipliersTData, shape: [3] },
87
+ vocabSizesT: { data: vocabSizesTData, shape: [4] },
88
+ }, {
89
+ attrs: { max_ngram_size: 3, n_head_per_ngram: 2, pad_id: 9 },
90
+ });
91
+ ```
build/webgpu/bench.json ADDED
The diff for this file is too large to render. See raw diff
 
build/webgpu/manifest.json ADDED
@@ -0,0 +1,148 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "domain": "com.microsoft",
3
+ "name": "NGramHashMapping",
4
+ "sinceVersion": 1,
5
+ "inputs": {
6
+ "inputIdsT": { "onnx": "input_ids", "dtype": "M", "rank": 2 },
7
+ "multipliersT": { "onnx": "multipliers", "dtype": "M", "rank": 1 },
8
+ "vocabSizesT": { "onnx": "vocab_sizes", "dtype": "M", "rank": 1 },
9
+ "pastIdsT": { "onnx": "past_ids", "dtype": "M", "rank": 2, "optional": true }
10
+ },
11
+ "outputs": {
12
+ "hashIdsT": {
13
+ "onnx": "hash_ids",
14
+ "dtype": "M",
15
+ "rank": 3,
16
+ "shape": "[dim(shapes.inputIdsT, 0), dim(shapes.inputIdsT, 1), (attrs.max_ngram_size - 1) * attrs.n_head_per_ngram]"
17
+ },
18
+ "presentIdsT": {
19
+ "onnx": "present_ids",
20
+ "dtype": "M",
21
+ "rank": 2,
22
+ "shape": "[dim(shapes.inputIdsT, 0), attrs.max_ngram_size - 1]"
23
+ }
24
+ },
25
+ "attributes": { "max_ngram_size": {}, "n_head_per_ngram": {}, "pad_id": {} },
26
+ "attributeConstraints": {
27
+ "max_ngram_size": { "required": true },
28
+ "n_head_per_ngram": { "required": true },
29
+ "pad_id": { "required": true }
30
+ },
31
+ "typeConstraints": { "M": ["int32"] },
32
+ "tunables": { "WORKGROUP_SIZE": { "default": 256 }, "ITEMS_PER_LANE": { "default": 4 } },
33
+ "derive": {
34
+ "foldedDispatchCapacity": "min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
35
+ "maxNgramSize": "attrs.max_ngram_size",
36
+ "headsPerNgram": "attrs.n_head_per_ngram",
37
+ "stateLength": "maxNgramSize - 1",
38
+ "numHeads": "stateLength * headsPerNgram",
39
+ "stagedTableScalars": "maxNgramSize + numHeads",
40
+ "padId": "attrs.pad_id",
41
+ "idsRankOk": "ranks.inputIdsT == 2",
42
+ "batchSize": "dim(shapes.inputIdsT, 0) if idsRankOk else 0",
43
+ "seqLength": "dim(shapes.inputIdsT, 1) if idsRankOk else 0",
44
+ "spanLength": "seqLength + stateLength",
45
+ "sharedTileSlots": "max(1, min(tunables.WORKGROUP_SIZE, spanLength, floor(device.limits.maxComputeWorkgroupStorageSize / (max(1, stateLength) * 4))))",
46
+ "sharedSpanTiles": "ceilDiv(spanLength, sharedTileSlots)",
47
+ "sharedDispatchOk": "sharedSpanTiles <= foldedDispatchCapacity",
48
+ "sharedMemoryOk": "sharedTileSlots * stateLength * 4 <= device.limits.maxComputeWorkgroupStorageSize",
49
+ "tileSlots": "tunables.WORKGROUP_SIZE * tunables.ITEMS_PER_LANE",
50
+ "spanTiles": "ceilDiv(spanLength, max(1, tileSlots))",
51
+ "idScalar": "dtypes.M",
52
+ "attrsOk": "maxNgramSize >= 2 and maxNgramSize <= 16 and headsPerNgram >= 1 and numHeads <= 64 and padId >= -2147483648 and padId <= 2147483647",
53
+ "dtypeOk": "tensorDtypes.inputIdsT == \"int32\" and tensorDtypes.multipliersT == \"int32\" and tensorDtypes.vocabSizesT == \"int32\" and tensorDtypes.hashIdsT == \"int32\" and tensorDtypes.presentIdsT == \"int32\"",
54
+ "shapeOk": "idsRankOk and ranks.multipliersT == 1 and dim(shapes.multipliersT, 0) == maxNgramSize and ranks.vocabSizesT == 1 and dim(shapes.vocabSizesT, 0) == numHeads and ranks.hashIdsT == 3 and sameShape(shapes.hashIdsT, [batchSize, seqLength, numHeads]) and ranks.presentIdsT == 2 and sameShape(shapes.presentIdsT, [batchSize, stateLength])",
55
+ "dispatchOk": "batchSize <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and spanTiles <= foldedDispatchCapacity",
56
+ "workgroupOk": "tunables.WORKGROUP_SIZE > 0 and tunables.WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX and tunables.ITEMS_PER_LANE >= 1 and floor(tunables.ITEMS_PER_LANE) == tunables.ITEMS_PER_LANE",
57
+ "baseContract": "attrsOk and dtypeOk and shapeOk and dispatchOk and workgroupOk",
58
+ "noPastContract": "baseContract and not present.pastIdsT",
59
+ "pastContract": "baseContract and present.pastIdsT and tensorDtypes.pastIdsT == \"int32\" and ranks.pastIdsT == 2 and sameShape(shapes.pastIdsT, [batchSize, stateLength])"
60
+ },
61
+ "bindings": {
62
+ "input_ids": { "arg": "inputIdsT", "elementType": "$idScalar" },
63
+ "multipliers": { "arg": "multipliersT", "elementType": "$idScalar", "length": "$maxNgramSize" },
64
+ "vocab_sizes": { "arg": "vocabSizesT", "elementType": "$idScalar", "length": "$numHeads" },
65
+ "past_ids": { "arg": "pastIdsT", "elementType": "$idScalar" },
66
+ "hash_ids": { "arg": "hashIdsT", "elementType": "$idScalar" },
67
+ "present_ids": { "arg": "presentIdsT", "elementType": "$idScalar" },
68
+ "params": { "struct": [{ "name": "seqLength", "type": "u32", "value": "seqLength" }] }
69
+ },
70
+ "variants": [
71
+ {
72
+ "id": "past",
73
+ "when": ["pastContract"],
74
+ "derive": { "hasPast": true },
75
+ "passes": [
76
+ {
77
+ "id": "main",
78
+ "name": "NGramHashMapping.Past",
79
+ "shader": "ngram-hash-mapping.wgsl.jinja",
80
+ "bindings": ["input_ids", "multipliers", "vocab_sizes", "past_ids", "hash_ids", "present_ids", "params"],
81
+ "dispatch": {
82
+ "x": "min(spanTiles, DISPATCH_FOLD_WIDTH)",
83
+ "y": "batchSize",
84
+ "z": "ceilDiv(spanTiles, DISPATCH_FOLD_WIDTH)"
85
+ }
86
+ }
87
+ ]
88
+ },
89
+ {
90
+ "id": "past_shared_prefix",
91
+ "when": ["pastContract", "sharedDispatchOk", "sharedMemoryOk"],
92
+ "derive": { "hasPast": true },
93
+ "passes": [
94
+ {
95
+ "id": "main",
96
+ "name": "NGramHashMapping.Past.SharedPrefix",
97
+ "shader": "ngram-hash-mapping-shared-prefix.wgsl.jinja",
98
+ "bindings": ["input_ids", "multipliers", "vocab_sizes", "past_ids", "hash_ids", "present_ids", "params"],
99
+ "dispatch": {
100
+ "x": "min(sharedSpanTiles, DISPATCH_FOLD_WIDTH)",
101
+ "y": "batchSize",
102
+ "z": "ceilDiv(sharedSpanTiles, DISPATCH_FOLD_WIDTH)"
103
+ }
104
+ }
105
+ ],
106
+ "priority": 10,
107
+ "demoteWhen": ["stagedTableScalars < 59"]
108
+ },
109
+ {
110
+ "id": "fresh",
111
+ "when": ["noPastContract"],
112
+ "derive": { "hasPast": false },
113
+ "passes": [
114
+ {
115
+ "id": "main",
116
+ "name": "NGramHashMapping.Fresh",
117
+ "shader": "ngram-hash-mapping.wgsl.jinja",
118
+ "bindings": ["input_ids", "multipliers", "vocab_sizes", "hash_ids", "present_ids", "params"],
119
+ "dispatch": {
120
+ "x": "min(spanTiles, DISPATCH_FOLD_WIDTH)",
121
+ "y": "batchSize",
122
+ "z": "ceilDiv(spanTiles, DISPATCH_FOLD_WIDTH)"
123
+ }
124
+ }
125
+ ]
126
+ },
127
+ {
128
+ "id": "fresh_shared_prefix",
129
+ "when": ["noPastContract", "sharedDispatchOk", "sharedMemoryOk"],
130
+ "derive": { "hasPast": false },
131
+ "passes": [
132
+ {
133
+ "id": "main",
134
+ "name": "NGramHashMapping.Fresh.SharedPrefix",
135
+ "shader": "ngram-hash-mapping-shared-prefix.wgsl.jinja",
136
+ "bindings": ["input_ids", "multipliers", "vocab_sizes", "hash_ids", "present_ids", "params"],
137
+ "dispatch": {
138
+ "x": "min(sharedSpanTiles, DISPATCH_FOLD_WIDTH)",
139
+ "y": "batchSize",
140
+ "z": "ceilDiv(sharedSpanTiles, DISPATCH_FOLD_WIDTH)"
141
+ }
142
+ }
143
+ ],
144
+ "priority": 10,
145
+ "demoteWhen": ["stagedTableScalars < 59"]
146
+ }
147
+ ]
148
+ }
build/webgpu/metadata.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "com.microsoft.NGramHashMapping",
3
+ "id": "_com_microsoft_ngramhashmapping_webgpu_2772cbc",
4
+ "version": 1,
5
+ "license": "Apache-2.0",
6
+ "backend": { "type": "webgpu" },
7
+ "digest": {
8
+ "algorithm": "sha256",
9
+ "files": {
10
+ "bench.json": "9TYxNQMMA7MFXU6lsmktq1DuTX+F96McWtO10Wi6Kow=",
11
+ "manifest.json": "GIvp96ai3vrBQBgRUoQPH44v0GP+JsH/ps8PEWjC0io=",
12
+ "ngram-hash-mapping-shared-prefix.wgsl.jinja": "AxbE9+l3w7VS9v44/AxMP+ncC4yxdss9T1A44ZQ+9Fw=",
13
+ "ngram-hash-mapping.wgsl.jinja": "2PNqcut+20bhW9G5cOGwNKBRjzpnSIErvj3ZquDQX14=",
14
+ "test.json": "zzm2j7sopidQ0yqpDfJEr3BWCj0Z+Vl0okbjsSeEbfw="
15
+ }
16
+ },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
+ "webgpu": {
19
+ "manifestSpec": "2.1",
20
+ "variants": {
21
+ "past": ["ngram-hash-mapping.wgsl.jinja"],
22
+ "past_shared_prefix": ["ngram-hash-mapping-shared-prefix.wgsl.jinja"],
23
+ "fresh": ["ngram-hash-mapping.wgsl.jinja"],
24
+ "fresh_shared_prefix": ["ngram-hash-mapping-shared-prefix.wgsl.jinja"]
25
+ }
26
+ }
27
+ }
build/webgpu/ngram-hash-mapping-shared-prefix.wgsl.jinja ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {{ env.wgsl.resourceDeclarations }}
2
+
3
+ const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
4
+ const TILE: u32 = {{ sharedTileSlots }}u;
5
+ const STATE_LENGTH: u32 = {{ stateLength }}u;
6
+ const HEADS: u32 = {{ numHeads }}u;
7
+ const HEADS_PER_ORDER: u32 = {{ headsPerNgram }}u;
8
+ {% if not hasPast %}
9
+ const PAD_ID: i32 = {{ padId }};
10
+ {% endif %}
11
+ // Order-major storage coalesces prefix writes and shares each value across heads.
12
+ var<workgroup> prefixes: array<i32, {{ sharedTileSlots * stateLength }}>;
13
+
14
+ fn positive_mod(value: i32, divisor: i32) -> i32 {
15
+ let r = value % max(divisor, 1);
16
+ return select(r, r + divisor, r < 0);
17
+ }
18
+ @compute @workgroup_size(WG)
19
+ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
20
+ let b = wg.y;
21
+ let tile = wg.x + wg.z * {{ DISPATCH_FOLD_WIDTH }}u;
22
+ let seq = params.seqLength;
23
+ let tileBase = tile * TILE;
24
+ let inputBase = b * seq;
25
+ {% if hasPast %}
26
+ let pastBase = b * STATE_LENGTH;
27
+ {% endif %}
28
+ if (lid.x < TILE) {
29
+ let t = tileBase + lid.x;
30
+ if (t < seq) {
31
+ // Unsigned products preserve the specified two's-complement wrapping.
32
+ var mix = bitcast<u32>(input_ids[inputBase + t]) * bitcast<u32>(multipliers[0u]);
33
+ {% for k in range(1, maxNgramSize) %}
34
+ {% if hasPast %}
35
+ var token{{ k }}: i32;
36
+ if (t >= {{ k }}u) { token{{ k }} = input_ids[inputBase + t - {{ k }}u]; }
37
+ else { token{{ k }} = past_ids[pastBase + STATE_LENGTH + t - {{ k }}u]; }
38
+ {% else %}
39
+ var token{{ k }} = PAD_ID;
40
+ if (t >= {{ k }}u) { token{{ k }} = input_ids[inputBase + t - {{ k }}u]; }
41
+ {% endif %}
42
+ mix = mix ^ (bitcast<u32>(token{{ k }}) * bitcast<u32>(multipliers[{{ k }}u]));
43
+ prefixes[{{ k - 1 }}u * TILE + lid.x] = bitcast<i32>(mix);
44
+ {% endfor %}
45
+ } else if (t < seq + STATE_LENGTH) {
46
+ {% if hasPast %}
47
+ var carry: i32;
48
+ if (t >= STATE_LENGTH) { carry = input_ids[inputBase + t - STATE_LENGTH]; }
49
+ else { carry = past_ids[pastBase + t]; }
50
+ {% else %}
51
+ var carry = PAD_ID;
52
+ if (t >= STATE_LENGTH) { carry = input_ids[inputBase + t - STATE_LENGTH]; }
53
+ {% endif %}
54
+ present_ids[b * STATE_LENGTH + t - seq] = carry;
55
+ }
56
+ }
57
+ // Every lane reaches the barrier, including carry-only and partial tiles.
58
+ workgroupBarrier();
59
+ var tokens = 0u;
60
+ if (tileBase < seq) { tokens = min(TILE, seq - tileBase); }
61
+ for (var index = lid.x; index < tokens * HEADS; index = index + WG) {
62
+ let token = index / HEADS;
63
+ let head = index % HEADS;
64
+ let mix = prefixes[(head / HEADS_PER_ORDER) * TILE + token];
65
+ hash_ids[(inputBase + tileBase) * HEADS + index] = positive_mod(mix, vocab_sizes[head]);
66
+ }
67
+ }
build/webgpu/ngram-hash-mapping.wgsl.jinja ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {{ env.wgsl.resourceDeclarations }}
2
+
3
+ // Engram n-gram hash ids. For every order n in [2, MAX_NGRAM] the kernel mixes the n causal
4
+ // shifts of the compressed tokenizer ids,
5
+ // mix_n = ids[t] * m[0] ^ ids[t-1] * m[1] ^ ... ^ ids[t-n+1] * m[n-1],
6
+ // and emits mix_n modulo each of that order's head vocabulary sizes. mix_n extends mix_{n-1}
7
+ // by a single term, so one walk over the window produces every order.
8
+ //
9
+ // Each product is formed in u32 so that overflow wraps in two's complement, which is the
10
+ // arithmetic the operator specifies. A signed multiply that overflows has an indeterminate
11
+ // result in WGSL, so the bitcasts here are load-bearing rather than cosmetic.
12
+ //
13
+ // One dispatch covers both outputs. Workgroup row wg.y owns batch row b, and the tile index
14
+ // folded across wg.x/wg.z walks a span of SEQ + STATE_LENGTH slots: the first SEQ slots hash a
15
+ // token position, and the trailing STATE_LENGTH slots emit the carry state that the next call
16
+ // feeds back as past_ids. A slot past the end of the span exits, so a short sequence costs one
17
+ // workgroup rather than a second dispatch.
18
+ const WG: u32 = {{ tunables.WORKGROUP_SIZE }}u;
19
+ const ITEMS: u32 = {{ tunables.ITEMS_PER_LANE }}u;
20
+ const STATE_LENGTH: u32 = {{ stateLength }}u;
21
+ const PAD_ID: i32 = {{ padId }};
22
+
23
+ // Euclidean remainder for a positive divisor: WGSL `%` truncates toward zero exactly as C
24
+ // does, so a negative dividend needs the one correction the operator specifies. A head whose
25
+ // vocabulary size is not strictly positive has no meaningful modulo, and clamping the divisor
26
+ // to 1 makes its remainder 0 -- the value every GPU implementation of this operator emits for
27
+ // that head, and the reason it stays a device-side division rather than a rejected input.
28
+ fn positive_mod(value: i32, divisor: i32) -> i32 {
29
+ let r = value % max(divisor, 1);
30
+ return select(r, r + divisor, r < 0);
31
+ }
32
+ {% macro hash_position(interior) %}
33
+ let idBase = inputBase + t;
34
+ var mix = bitcast<u32>(input_ids[idBase]) * mult0;
35
+ {% for k in range(1, maxNgramSize) %}
36
+ {% if interior %}
37
+ mix = mix ^ (bitcast<u32>(input_ids[idBase - {{ k }}u]) * mult{{ k }});
38
+ {% else %}
39
+ var token{{ k }} = PAD_ID;
40
+ if (t >= {{ k }}u) {
41
+ token{{ k }} = input_ids[idBase - {{ k }}u];
42
+ }
43
+ {% if hasPast %}
44
+ else {
45
+ // past_ids is right-aligned, so position -1 is its last slot. t < k bounds the index
46
+ // inside the window, which is why no second range test is needed here.
47
+ token{{ k }} = past_ids[pastBase + STATE_LENGTH + t - {{ k }}u];
48
+ }
49
+ {% endif %}
50
+ mix = mix ^ (bitcast<u32>(token{{ k }}) * mult{{ k }});
51
+ {% endif %}
52
+ let mixOrder{{ k }} = bitcast<i32>(mix);
53
+ {% endfor %}
54
+ {% for n in range(1, maxNgramSize) %}
55
+ {% for h in range(headsPerNgram) %}
56
+ let head{{ (n - 1) * headsPerNgram + h }} = positive_mod(mixOrder{{ n }}, vocab{{ (n - 1) * headsPerNgram + h }});
57
+ {% endfor %}
58
+ {% endfor %}
59
+ let outBase = idBase * {{ numHeads }}u;
60
+ {% for h in range(numHeads) %}
61
+ hash_ids[outBase + {{ h }}u] = head{{ h }};
62
+ {% endfor %}
63
+ {% endmacro %}
64
+
65
+ @compute @workgroup_size(WG)
66
+ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
67
+ // wg.y is a real axis -- one workgroup row per batch row -- so only the span tiles are
68
+ // folded, across wg.x and wg.z at the per-axis workgroup fold width.
69
+ let b = wg.y;
70
+ let tile = wg.x + wg.z * {{ DISPATCH_FOLD_WIDTH }}u;
71
+ let seq = params.seqLength;
72
+ let span = seq + STATE_LENGTH;
73
+ let inputBase = b * seq;
74
+ let tileBase = tile * (WG * ITEMS);
75
+ {% if hasPast %}
76
+ let pastBase = b * STATE_LENGTH;
77
+ {% endif %}
78
+
79
+ // The multiplier and vocabulary tables are attribute-sized and every position reads all of
80
+ // them, so an invocation stages them once instead of reloading them inside the mix walk.
81
+ {% for k in range(maxNgramSize) %}
82
+ let mult{{ k }} = bitcast<u32>(multipliers[{{ k }}u]);
83
+ {% endfor %}
84
+ {% for h in range(numHeads) %}
85
+ let vocab{{ h }} = vocab_sizes[{{ h }}u];
86
+ {% endfor %}
87
+
88
+ // An interior tile is one whose every slot hashes a token position and whose every causal
89
+ // shift lands inside input_ids. Both of the per-slot tests below, and the per-shift history
90
+ // fallback, are then statically false, which is worth one workgroup-uniform branch.
91
+ if (tileBase >= STATE_LENGTH && tileBase + WG * ITEMS <= seq) {
92
+ for (var item = 0u; item < ITEMS; item = item + 1u) {
93
+ let t = tileBase + lid.x + item * WG;
94
+ {{ hash_position(true) }}
95
+ }
96
+ return;
97
+ }
98
+
99
+ for (var item = 0u; item < ITEMS; item = item + 1u) {
100
+ let slot = tileBase + lid.x + item * WG;
101
+ if (slot >= span) {
102
+ break;
103
+ }
104
+ if (slot < seq) {
105
+ let t = slot;
106
+ {{ hash_position(false) }}
107
+ } else {
108
+ // Carry slot j holds the id at position seq - STATE_LENGTH + j of this call, and slot j
109
+ // of past_ids when the call is shorter than the window. Writing it from this dispatch is
110
+ // safe because present_ids never shares its buffer with past_ids here.
111
+ var carry = PAD_ID;
112
+ if (slot >= STATE_LENGTH) {
113
+ carry = input_ids[inputBase + slot - STATE_LENGTH];
114
+ }
115
+ {% if hasPast %}
116
+ else {
117
+ carry = past_ids[pastBase + slot];
118
+ }
119
+ {% endif %}
120
+ present_ids[b * STATE_LENGTH + slot - seq] = carry;
121
+ }
122
+ }
123
+ }
build/webgpu/test.json ADDED
The diff for this file is too large to render. See raw diff