Xenova HF Staff commited on
Commit
e3afd9e
·
verified ·
1 Parent(s): 2e68766

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -44,13 +44,13 @@ See the [ONNX Runtime `LinearAttentionGate` contrib-operator spec](https://githu
44
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
45
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
46
  - [`test.json`](build/webgpu/test.json) — correctness cases
47
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
48
  - [`linear-attention-gate.wgsl.jinja`](build/webgpu/linear-attention-gate.wgsl.jinja)
49
 
50
  ## Use with `@huggingface/kernels`
51
 
52
  ```sh
53
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
54
  ```
55
 
56
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
44
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
45
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
46
  - [`test.json`](build/webgpu/test.json) — correctness cases
47
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
48
  - [`linear-attention-gate.wgsl.jinja`](build/webgpu/linear-attention-gate.wgsl.jinja)
49
 
50
  ## Use with `@huggingface/kernels`
51
 
52
  ```sh
53
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
54
  ```
55
 
56
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/linear-attention-gate.wgsl.jinja CHANGED
@@ -1,6 +1,11 @@
1
- {% if usesF16 %}
2
- enable f16;
3
- {% endif -%}
 
 
 
 
 
4
  {{ env.wgsl.resourceDeclarations }}
5
 
6
  {% if vectorized %}
@@ -53,11 +58,7 @@ fn softplus(x: f32) -> f32 {
53
  fn main(
54
  @builtin(global_invocation_id) gid: vec3<u32>
55
  ) {
56
- // Rebuild the flat invocation index after the 2D dispatch fold.
57
- let item = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WORKGROUP_SIZE;
58
- if (item >= GATE_ITEMS) {
59
- return;
60
- }
61
 
62
  // The last axis is the head axis, so the per-head parameter index is the flat index
63
  // modulo the head count. The vectorized path holds because the head count is a multiple
 
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
  {% if vectorized %}
 
58
  fn main(
59
  @builtin(global_invocation_id) gid: vec3<u32>
60
  ) {
61
+ {{ flat_index_2d("WORKGROUP_SIZE", "item", "GATE_ITEMS") }}
 
 
 
 
62
 
63
  // The last axis is the head axis, so the per-head parameter index is the flat index
64
  // modulo the head count. The vectorized path holds because the head count is a multiple
build/webgpu/manifest.json CHANGED
@@ -29,26 +29,17 @@
29
  "decayOnlyContract": "tensorContract and not present.betaT",
30
  "workgroupFits": "tunables.WORKGROUP_SIZE > 0 and tunables.WORKGROUP_SIZE <= deviceWorkgroupCap",
31
  "scalarDispatchFits": "ceilDiv(gateCount, tunables.WORKGROUP_SIZE) <= foldedDispatchCapacity",
32
- "vec4DispatchFits": "numHeads % 4 == 0 and ceilDiv(gateVec4Count, tunables.WORKGROUP_SIZE) <= foldedDispatchCapacity"
 
33
  },
34
  "when": ["workgroupFits"],
35
  "bindings": {
36
- "a": { "arg": "aT", "buffer": "read-only-storage", "elementType": "$gateElement", "length": "$gateItems" },
37
- "dt_bias": {
38
- "arg": "dtBiasT",
39
- "buffer": "read-only-storage",
40
- "elementType": "$paramElement",
41
- "length": "$headItems"
42
- },
43
- "decay_scale": {
44
- "arg": "decayScaleT",
45
- "buffer": "read-only-storage",
46
- "elementType": "$paramElement",
47
- "length": "$headItems"
48
- },
49
- "b": { "arg": "bT", "buffer": "read-only-storage", "elementType": "$gateElement", "length": "$gateItems" },
50
- "decay": { "arg": "decayT", "buffer": "storage", "elementType": "$gateElement", "length": "$gateItems" },
51
- "beta": { "arg": "betaT", "buffer": "storage", "elementType": "$gateElement", "length": "$gateItems" }
52
  },
53
  "variants": [
54
  {
@@ -57,13 +48,11 @@
57
  "when": ["betaContract", "vec4DispatchFits"],
58
  "derive": {
59
  "vectorized": true,
60
- "hasBeta": true,
61
- "usesF16": "gateDtype == \"float16\"",
62
  "gateElement": "\"vec4<f16>\" if gateDtype == \"float16\" else \"vec4<f32>\"",
63
  "paramElement": "\"vec4<f32>\"",
64
  "headItems": "headsVec4",
65
- "gateItems": "gateVec4Count",
66
- "workgroupSize": "tunables.WORKGROUP_SIZE"
67
  },
68
  "passes": [
69
  {
@@ -72,8 +61,8 @@
72
  "shader": "linear-attention-gate.wgsl.jinja",
73
  "bindings": ["a", "dt_bias", "decay_scale", "b", "decay", "beta"],
74
  "dispatch": {
75
- "x": "min(ceilDiv((gateVec4Count), (tunables.WORKGROUP_SIZE)), 65535)",
76
- "y": "ceilDiv(ceilDiv((gateVec4Count), (tunables.WORKGROUP_SIZE)), 65535)",
77
  "z": 1
78
  }
79
  }
@@ -85,13 +74,11 @@
85
  "when": ["decayOnlyContract", "vec4DispatchFits"],
86
  "derive": {
87
  "vectorized": true,
88
- "hasBeta": false,
89
- "usesF16": "gateDtype == \"float16\"",
90
  "gateElement": "\"vec4<f16>\" if gateDtype == \"float16\" else \"vec4<f32>\"",
91
  "paramElement": "\"vec4<f32>\"",
92
  "headItems": "headsVec4",
93
- "gateItems": "gateVec4Count",
94
- "workgroupSize": "tunables.WORKGROUP_SIZE"
95
  },
96
  "passes": [
97
  {
@@ -100,8 +87,8 @@
100
  "shader": "linear-attention-gate.wgsl.jinja",
101
  "bindings": ["a", "dt_bias", "decay_scale", "decay"],
102
  "dispatch": {
103
- "x": "min(ceilDiv((gateVec4Count), (tunables.WORKGROUP_SIZE)), 65535)",
104
- "y": "ceilDiv(ceilDiv((gateVec4Count), (tunables.WORKGROUP_SIZE)), 65535)",
105
  "z": 1
106
  }
107
  }
@@ -113,13 +100,11 @@
113
  "when": ["betaContract", "scalarDispatchFits"],
114
  "derive": {
115
  "vectorized": false,
116
- "hasBeta": true,
117
- "usesF16": "gateDtype == \"float16\"",
118
  "gateElement": "\"f16\" if gateDtype == \"float16\" else \"f32\"",
119
  "paramElement": "\"f32\"",
120
  "headItems": "numHeads",
121
- "gateItems": "gateCount",
122
- "workgroupSize": "tunables.WORKGROUP_SIZE"
123
  },
124
  "passes": [
125
  {
@@ -128,8 +113,8 @@
128
  "shader": "linear-attention-gate.wgsl.jinja",
129
  "bindings": ["a", "dt_bias", "decay_scale", "b", "decay", "beta"],
130
  "dispatch": {
131
- "x": "min(ceilDiv((gateCount), (tunables.WORKGROUP_SIZE)), 65535)",
132
- "y": "ceilDiv(ceilDiv((gateCount), (tunables.WORKGROUP_SIZE)), 65535)",
133
  "z": 1
134
  }
135
  }
@@ -141,13 +126,11 @@
141
  "when": ["decayOnlyContract", "scalarDispatchFits"],
142
  "derive": {
143
  "vectorized": false,
144
- "hasBeta": false,
145
- "usesF16": "gateDtype == \"float16\"",
146
  "gateElement": "\"f16\" if gateDtype == \"float16\" else \"f32\"",
147
  "paramElement": "\"f32\"",
148
  "headItems": "numHeads",
149
- "gateItems": "gateCount",
150
- "workgroupSize": "tunables.WORKGROUP_SIZE"
151
  },
152
  "passes": [
153
  {
@@ -156,8 +139,8 @@
156
  "shader": "linear-attention-gate.wgsl.jinja",
157
  "bindings": ["a", "dt_bias", "decay_scale", "decay"],
158
  "dispatch": {
159
- "x": "min(ceilDiv((gateCount), (tunables.WORKGROUP_SIZE)), 65535)",
160
- "y": "ceilDiv(ceilDiv((gateCount), (tunables.WORKGROUP_SIZE)), 65535)",
161
  "z": 1
162
  }
163
  }
 
29
  "decayOnlyContract": "tensorContract and not present.betaT",
30
  "workgroupFits": "tunables.WORKGROUP_SIZE > 0 and tunables.WORKGROUP_SIZE <= deviceWorkgroupCap",
31
  "scalarDispatchFits": "ceilDiv(gateCount, tunables.WORKGROUP_SIZE) <= foldedDispatchCapacity",
32
+ "vec4DispatchFits": "numHeads % 4 == 0 and ceilDiv(gateVec4Count, tunables.WORKGROUP_SIZE) <= foldedDispatchCapacity",
33
+ "workgroupSize": "tunables.WORKGROUP_SIZE"
34
  },
35
  "when": ["workgroupFits"],
36
  "bindings": {
37
+ "a": { "arg": "aT", "elementType": "$gateElement", "length": "$gateItems" },
38
+ "dt_bias": { "arg": "dtBiasT", "elementType": "$paramElement", "length": "$headItems" },
39
+ "decay_scale": { "arg": "decayScaleT", "elementType": "$paramElement", "length": "$headItems" },
40
+ "b": { "arg": "bT", "elementType": "$gateElement", "length": "$gateItems" },
41
+ "decay": { "arg": "decayT", "elementType": "$gateElement", "length": "$gateItems" },
42
+ "beta": { "arg": "betaT", "elementType": "$gateElement", "length": "$gateItems" }
 
 
 
 
 
 
 
 
 
 
43
  },
44
  "variants": [
45
  {
 
48
  "when": ["betaContract", "vec4DispatchFits"],
49
  "derive": {
50
  "vectorized": true,
51
+ "hasBeta": "present.betaT",
 
52
  "gateElement": "\"vec4<f16>\" if gateDtype == \"float16\" else \"vec4<f32>\"",
53
  "paramElement": "\"vec4<f32>\"",
54
  "headItems": "headsVec4",
55
+ "gateItems": "gateVec4Count"
 
56
  },
57
  "passes": [
58
  {
 
61
  "shader": "linear-attention-gate.wgsl.jinja",
62
  "bindings": ["a", "dt_bias", "decay_scale", "b", "decay", "beta"],
63
  "dispatch": {
64
+ "x": "min(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
65
+ "y": "ceilDiv(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
66
  "z": 1
67
  }
68
  }
 
74
  "when": ["decayOnlyContract", "vec4DispatchFits"],
75
  "derive": {
76
  "vectorized": true,
77
+ "hasBeta": "present.betaT",
 
78
  "gateElement": "\"vec4<f16>\" if gateDtype == \"float16\" else \"vec4<f32>\"",
79
  "paramElement": "\"vec4<f32>\"",
80
  "headItems": "headsVec4",
81
+ "gateItems": "gateVec4Count"
 
82
  },
83
  "passes": [
84
  {
 
87
  "shader": "linear-attention-gate.wgsl.jinja",
88
  "bindings": ["a", "dt_bias", "decay_scale", "decay"],
89
  "dispatch": {
90
+ "x": "min(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
91
+ "y": "ceilDiv(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
92
  "z": 1
93
  }
94
  }
 
100
  "when": ["betaContract", "scalarDispatchFits"],
101
  "derive": {
102
  "vectorized": false,
103
+ "hasBeta": "present.betaT",
 
104
  "gateElement": "\"f16\" if gateDtype == \"float16\" else \"f32\"",
105
  "paramElement": "\"f32\"",
106
  "headItems": "numHeads",
107
+ "gateItems": "gateCount"
 
108
  },
109
  "passes": [
110
  {
 
113
  "shader": "linear-attention-gate.wgsl.jinja",
114
  "bindings": ["a", "dt_bias", "decay_scale", "b", "decay", "beta"],
115
  "dispatch": {
116
+ "x": "min(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
117
+ "y": "ceilDiv(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
118
  "z": 1
119
  }
120
  }
 
126
  "when": ["decayOnlyContract", "scalarDispatchFits"],
127
  "derive": {
128
  "vectorized": false,
129
+ "hasBeta": "present.betaT",
 
130
  "gateElement": "\"f16\" if gateDtype == \"float16\" else \"f32\"",
131
  "paramElement": "\"f32\"",
132
  "headItems": "numHeads",
133
+ "gateItems": "gateCount"
 
134
  },
135
  "passes": [
136
  {
 
139
  "shader": "linear-attention-gate.wgsl.jinja",
140
  "bindings": ["a", "dt_bias", "decay_scale", "decay"],
141
  "dispatch": {
142
+ "x": "min(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
143
+ "y": "ceilDiv(ceilDiv((gateItems), (tunables.WORKGROUP_SIZE)), 65535)",
144
  "z": 1
145
  }
146
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.LinearAttentionGate",
3
- "id": "_com_microsoft_linearattentiongate_webgpu_4bb5397",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,14 +8,14 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "nSMkCE+adLNKSBkCYEFz8YeOJf6YjUsw4vEzFToXs5U=",
11
- "linear-attention-gate.wgsl.jinja": "Gi935dqD4NLjZbH6v4gzTYZElbzpeL5mC+ZD91H5pvo=",
12
- "manifest.json": "kEskRGoQGGG0erg57qnQc/Sa4hJq5cNnduwqwU5zuDk=",
13
- "test.json": "rlpI/FCSMX4yxQoveUCaj13GqK8JA+dadMzKFo129I8="
14
  }
15
  },
16
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
17
  "webgpu": {
18
- "manifestSpec": "2.0",
19
  "variants": {
20
  "vec4_gate_beta": ["linear-attention-gate.wgsl.jinja"],
21
  "vec4_gate": ["linear-attention-gate.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.LinearAttentionGate",
3
+ "id": "_com_microsoft_linearattentiongate_webgpu_502eb17",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "nSMkCE+adLNKSBkCYEFz8YeOJf6YjUsw4vEzFToXs5U=",
11
+ "linear-attention-gate.wgsl.jinja": "/e9pU3UP1X3ScwIjKeEPLVvqcznnSXVYQAvll/xqigE=",
12
+ "manifest.json": "SjgmqEe/oy28CeKd29htvOfS0sW/t7WqolTgm5l1U6k=",
13
+ "test.json": "NLi6ptg6CNprO6zHHud1Vxkf7JZduFkizDYvb6CttXA="
14
  }
15
  },
16
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
17
  "webgpu": {
18
+ "manifestSpec": "2.1",
19
  "variants": {
20
  "vec4_gate_beta": ["linear-attention-gate.wgsl.jinja"],
21
  "vec4_gate": ["linear-attention-gate.wgsl.jinja"],
build/webgpu/test.json CHANGED
@@ -3,7 +3,7 @@
3
  {
4
  "name": "rank3_h8_vec4_gate_beta",
5
  "provenance": {
6
- "notes": "The schema's (B,T,H) layout uses a head count divisible by four, exercising vectorized beta gating. A 2e-6 tolerance covers f32 Softplus rounding near the log1p series crossover."
7
  },
8
  "inputs": {
9
  "aT": {
@@ -27,7 +27,7 @@
27
  {
28
  "name": "rank2_h6_scalar_gate_beta",
29
  "provenance": {
30
- "notes": "Head count 6 is not a multiple of four, so the vectorized head-to-parameter mapping does not hold and the scalar path is the only eligible one. Covers scalar_gate_beta."
31
  },
32
  "inputs": {
33
  "aT": {
@@ -262,6 +262,124 @@
262
  "decayT": { "dtype": "float16", "shape": [5, 6], "tolerance": 0, "relTolerance": 0.002 },
263
  "betaT": { "dtype": "float16", "shape": [5, 6], "tolerance": 0, "relTolerance": 0.002 }
264
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
265
  }
266
  ]
267
  }
 
3
  {
4
  "name": "rank3_h8_vec4_gate_beta",
5
  "provenance": {
6
+ "notes": "Head count 8 (divisible by four) exercises four-wide beta computation in the (B,T,H) layout; a 2e-6 tolerance covers float32 Softplus rounding near where its log1p-based series approximation switches."
7
  },
8
  "inputs": {
9
  "aT": {
 
27
  {
28
  "name": "rank2_h6_scalar_gate_beta",
29
  "provenance": {
30
+ "notes": "Head count 6 is not a multiple of four, so four-wide parameter mapping does not apply; a rank-2 (no batch dimension) input checks per-head computation one head at a time."
31
  },
32
  "inputs": {
33
  "aT": {
 
262
  "decayT": { "dtype": "float16", "shape": [5, 6], "tolerance": 0, "relTolerance": 0.002 },
263
  "betaT": { "dtype": "float16", "shape": [5, 6], "tolerance": 0, "relTolerance": 0.002 }
264
  }
265
+ },
266
+ {
267
+ "name": "ort_gate_float_decode_h32",
268
+ "inputs": {
269
+ "aT": {
270
+ "dtype": "float32",
271
+ "shape": [1, 1, 32],
272
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.17, "scale": 6.0 }
273
+ },
274
+ "dtBiasT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
275
+ "decayScaleT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -4.0, "end": -0.1 } },
276
+ "bT": {
277
+ "dtype": "float32",
278
+ "shape": [1, 1, 32],
279
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.13, "scale": 6.0 }
280
+ }
281
+ },
282
+ "outputs": {
283
+ "decayT": { "dtype": "float32", "shape": [1, 1, 32], "tolerance": 0, "relTolerance": 0.000002 },
284
+ "betaT": { "dtype": "float32", "shape": [1, 1, 32], "tolerance": 0, "relTolerance": 0.000002 }
285
+ }
286
+ },
287
+ {
288
+ "name": "ort_gate_float_speculative_decode_tile_h32",
289
+ "inputs": {
290
+ "aT": {
291
+ "dtype": "float32",
292
+ "shape": [1, 4, 32],
293
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.17, "scale": 6.0 }
294
+ },
295
+ "dtBiasT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
296
+ "decayScaleT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -4.0, "end": -0.1 } },
297
+ "bT": {
298
+ "dtype": "float32",
299
+ "shape": [1, 4, 32],
300
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.13, "scale": 6.0 }
301
+ }
302
+ },
303
+ "outputs": {
304
+ "decayT": { "dtype": "float32", "shape": [1, 4, 32], "tolerance": 0, "relTolerance": 0.000002 },
305
+ "betaT": { "dtype": "float32", "shape": [1, 4, 32], "tolerance": 0, "relTolerance": 0.000002 }
306
+ }
307
+ },
308
+ {
309
+ "name": "ort_gate_float_decay_only_h16",
310
+ "inputs": {
311
+ "aT": {
312
+ "dtype": "float32",
313
+ "shape": [2, 3, 16],
314
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.17, "scale": 6.0 }
315
+ },
316
+ "dtBiasT": { "dtype": "float32", "shape": [16], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
317
+ "decayScaleT": { "dtype": "float32", "shape": [16], "data": { "kind": "linspace", "start": -4.0, "end": -0.1 } }
318
+ },
319
+ "outputs": { "decayT": { "dtype": "float32", "shape": [2, 3, 16], "tolerance": 0, "relTolerance": 0.000002 } }
320
+ },
321
+ {
322
+ "name": "ort_gate_float16_speculative_decode_tile_h32",
323
+ "inputs": {
324
+ "aT": {
325
+ "dtype": "float16",
326
+ "shape": [1, 4, 32],
327
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.17, "scale": 6.0 }
328
+ },
329
+ "dtBiasT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
330
+ "decayScaleT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -4.0, "end": -0.1 } },
331
+ "bT": {
332
+ "dtype": "float16",
333
+ "shape": [1, 4, 32],
334
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.13, "scale": 6.0 }
335
+ }
336
+ },
337
+ "outputs": {
338
+ "decayT": { "dtype": "float16", "shape": [1, 4, 32], "tolerance": 0, "relTolerance": 0.002 },
339
+ "betaT": { "dtype": "float16", "shape": [1, 4, 32], "tolerance": 0, "relTolerance": 0.002 }
340
+ }
341
+ },
342
+ {
343
+ "name": "ort_gate_float16_prefill_h32",
344
+ "inputs": {
345
+ "aT": {
346
+ "dtype": "float16",
347
+ "shape": [2, 37, 32],
348
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.17, "scale": 6.0 }
349
+ },
350
+ "dtBiasT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
351
+ "decayScaleT": { "dtype": "float32", "shape": [32], "data": { "kind": "linspace", "start": -4.0, "end": -0.1 } },
352
+ "bT": {
353
+ "dtype": "float16",
354
+ "shape": [2, 37, 32],
355
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.13, "scale": 6.0 }
356
+ }
357
+ },
358
+ "outputs": {
359
+ "decayT": { "dtype": "float16", "shape": [2, 37, 32], "tolerance": 0, "relTolerance": 0.002 },
360
+ "betaT": { "dtype": "float16", "shape": [2, 37, 32], "tolerance": 0, "relTolerance": 0.002 }
361
+ }
362
+ },
363
+ {
364
+ "name": "ort_gate_float16_ragged_tail_h7",
365
+ "inputs": {
366
+ "aT": {
367
+ "dtype": "float16",
368
+ "shape": [1, 5, 7],
369
+ "data": { "kind": "fillFloat32", "sinStep": 0.23, "cosStep": 0.17, "scale": 6.0 }
370
+ },
371
+ "dtBiasT": { "dtype": "float32", "shape": [7], "data": { "kind": "linspace", "start": -2.0, "end": 2.0 } },
372
+ "decayScaleT": { "dtype": "float32", "shape": [7], "data": { "kind": "linspace", "start": -4.0, "end": -0.1 } },
373
+ "bT": {
374
+ "dtype": "float16",
375
+ "shape": [1, 5, 7],
376
+ "data": { "kind": "fillFloat32", "sinStep": 0.31, "cosStep": 0.13, "scale": 6.0 }
377
+ }
378
+ },
379
+ "outputs": {
380
+ "decayT": { "dtype": "float16", "shape": [1, 5, 7], "tolerance": 0, "relTolerance": 0.002 },
381
+ "betaT": { "dtype": "float16", "shape": [1, 5, 7], "tolerance": 0, "relTolerance": 0.002 }
382
+ }
383
  }
384
  ]
385
  }