Xenova HF Staff commited on
Commit
767d527
·
verified ·
1 Parent(s): a3b790c

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -48,19 +48,30 @@ Default values (overridable per request):
48
  | --- | --- |
49
  | `T` | `float32`, `float16` |
50
 
 
 
 
 
 
 
 
 
 
 
 
51
  ## Files
52
 
53
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
54
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
55
  - [`test.json`](build/webgpu/test.json) — correctness cases
56
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
57
  - [`layer-normalization.wgsl.jinja`](build/webgpu/layer-normalization.wgsl.jinja)
58
  - [`norm-row-stats.wgsl.jinja`](build/webgpu/norm-row-stats.wgsl.jinja)
59
 
60
  ## Use with `@huggingface/kernels`
61
 
62
  ```sh
63
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
64
  ```
65
 
66
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
48
  | --- | --- |
49
  | `T` | `float32`, `float16` |
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
+ - `last_axis_row_vec2` — Uses paired float32 loads for even last-axis rows that cannot use four-wide storage.
56
+ - `last_axis_bias_row_vec2` — Uses paired float32 loads for even last-axis rows that cannot use four-wide storage.
57
+ - `last_axis_broadcast_row_vec4` — Normalizes packed rows with affine tensors broadcast across leading dimensions; short rows share a workgroup within device limits.
58
+ - `last_axis_broadcast_bias_row_vec4` — Normalizes packed rows with affine tensors broadcast across leading dimensions; short rows share a workgroup within device limits.
59
+ - `last_axis_broadcast_rows_vec4` — Normalizes packed rows with affine tensors broadcast across leading dimensions; short rows share a workgroup within device limits.
60
+ - `last_axis_broadcast_bias_rows_vec4` — Normalizes packed rows with affine tensors broadcast across leading dimensions; short rows share a workgroup within device limits.
61
+
62
  ## Files
63
 
64
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
65
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
66
  - [`test.json`](build/webgpu/test.json) — correctness cases
67
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
68
  - [`layer-normalization.wgsl.jinja`](build/webgpu/layer-normalization.wgsl.jinja)
69
  - [`norm-row-stats.wgsl.jinja`](build/webgpu/norm-row-stats.wgsl.jinja)
70
 
71
  ## Use with `@huggingface/kernels`
72
 
73
  ```sh
74
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
75
  ```
76
 
77
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/layer-normalization.wgsl.jinja CHANGED
@@ -33,25 +33,22 @@ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif
33
  {% endfor %}
34
  return offset;
35
  {% endif %}
36
- }
37
- {%- endmacro %}{% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
38
  {% set op_numel = namespace(value=1) %}
39
- {% for d in opShape %}{% set op_numel.value = op_numel.value * d %}{% endfor %}
 
 
40
  {% set out_numel = namespace(value=1) %}
41
- {% for d in outShape %}{% set out_numel.value = out_numel.value * d %}{% endfor %}
42
- {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %})
43
- {%- endmacro %}
44
-
45
  {{ env.wgsl.resourceDeclarations }}
46
 
47
  const HIDDEN: u32 = {{ hiddenSize }}u;
48
  const EPSILON: f32 = {{ epsilon }};
49
  const WG: u32 = {{ workgroupSize }}u;
50
-
51
- var<workgroup> partial: array<f32, WG>;
52
- var<workgroup> row_mean: f32;
53
- var<workgroup> row_inv: f32;
54
-
55
  {% set xNumel = namespace(value=1) %}
56
  {% for dim in xShape %}
57
  {% set xNumel.value = xNumel.value * dim %}
@@ -61,82 +58,223 @@ var<workgroup> row_inv: f32;
61
  {% set scaleNumel.value = scaleNumel.value * dim %}
62
  {% endfor %}
63
  {% if scaleNumel.value != 1 %}
 
64
  {{ offset_fn("scale_offset", scaleShape, scaleShape | length, scaleShape == xShape, scaleNumel.value, xShape, xShape | length, xNumel.value) }}
65
  {% endif %}
66
-
67
  {% if hasBias %}
68
  {% set biasNumel = namespace(value=1) %}
69
  {% for dim in biasShape %}
70
  {% set biasNumel.value = biasNumel.value * dim %}
71
  {% endfor %}
72
  {% if biasNumel.value != 1 %}
 
73
  {{ offset_fn("bias_offset", biasShape, biasShape | length, biasShape == xShape, biasNumel.value, xShape, xShape | length, xNumel.value) }}
74
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
76
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
78
- {% if op == "max" %}
79
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
80
- {%- else %}
81
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
82
- {%- endif %}
83
- {% endmacro %}
84
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
85
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
86
  loop {
87
- {% if form == "head" %}
88
- {% if breakInline %}
89
- if ({{ svar }} == 0u) { break; }
90
- {% else %}
91
  if ({{ svar }} == 0u) {
92
  break;
93
  }
94
- {% endif %}
95
- {% endif %}
96
- {% if bodyInline %}
97
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
98
- {% else %}
99
  if ({{ idx }} < {{ svar }}) {
100
  {% for a in arrays %}
101
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
102
  {% endfor %}
103
  }
104
- {% endif %}
105
- {% if form == "head" %}
106
- {% if barrierFirst %}
107
- workgroupBarrier();
108
  {{ svar }} = {{ svar }} / 2u;
109
- {% else %}
110
- {{ svar }} = {{ svar }} / 2u;
111
- workgroupBarrier();
112
- {% endif %}
113
- {% else %}
114
  workgroupBarrier();
115
- if ({{ svar }} == 1u) {
116
- break;
117
- }
118
- {{ svar }} = {{ svar }} / 2u;
119
- {% endif %}
120
- }
121
- {%- endmacro %}
122
-
123
  // Reusing partial after this reduction requires a barrier between the read of
124
  // partial[0] and the next write, or the next round can race the prior readers.
125
- {% set trailingBarrier = trailingBarrier is defined and trailingBarrier %}
126
  fn reduce_sum(value: f32, tid: u32) -> f32 {
127
  partial[tid] = value;
128
  workgroupBarrier();
129
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
130
- {% if trailingBarrier %}
131
- let total = partial[0];
132
- workgroupBarrier();
133
- return total;
134
- {% else %}
135
  return partial[0];
136
- {% endif %}
137
  }
138
 
139
-
140
  @compute @workgroup_size(WG, 1, 1)
141
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
142
  let row = wg.x + wg.y * params.rowStride;
@@ -187,3 +325,4 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
187
  y[index] = {{ scalar }}(value);
188
  }
189
  }
 
 
33
  {% endfor %}
34
  return offset;
35
  {% endif %}
36
+ }{% endmacro %}
37
+ {% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
38
  {% set op_numel = namespace(value=1) %}
39
+ {% for d in opShape %}
40
+ {% set op_numel.value = op_numel.value * d %}
41
+ {% endfor %}
42
  {% set out_numel = namespace(value=1) %}
43
+ {% for d in outShape %}
44
+ {% set out_numel.value = out_numel.value * d %}
45
+ {% endfor %}
46
+ {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %}){% endmacro %}
47
  {{ env.wgsl.resourceDeclarations }}
48
 
49
  const HIDDEN: u32 = {{ hiddenSize }}u;
50
  const EPSILON: f32 = {{ epsilon }};
51
  const WG: u32 = {{ workgroupSize }}u;
 
 
 
 
 
52
  {% set xNumel = namespace(value=1) %}
53
  {% for dim in xShape %}
54
  {% set xNumel.value = xNumel.value * dim %}
 
58
  {% set scaleNumel.value = scaleNumel.value * dim %}
59
  {% endfor %}
60
  {% if scaleNumel.value != 1 %}
61
+
62
  {{ offset_fn("scale_offset", scaleShape, scaleShape | length, scaleShape == xShape, scaleNumel.value, xShape, xShape | length, xNumel.value) }}
63
  {% endif %}
 
64
  {% if hasBias %}
65
  {% set biasNumel = namespace(value=1) %}
66
  {% for dim in biasShape %}
67
  {% set biasNumel.value = biasNumel.value * dim %}
68
  {% endfor %}
69
  {% if biasNumel.value != 1 %}
70
+
71
  {{ offset_fn("bias_offset", biasShape, biasShape | length, biasShape == xShape, biasNumel.value, xShape, xShape | length, xNumel.value) }}
72
  {% endif %}
73
+ {% endif %}
74
+
75
+ {% if scalar == "f16" %}
76
+ {% set modeSpec = "layer" %}
77
+ {% set halfWriteMean = writeMean %}
78
+ {% set halfWriteInv = writeInvStdDev %}
79
+ {% set halfScaleOffset = "0u" if scaleNumel.value == 1 else broadcast_offset_call("scale_offset", scaleShape, xShape, "base + i") %}
80
+ {% if hasBias %}
81
+ {% set halfBiasOffset = "0u" if biasNumel.value == 1 else broadcast_offset_call("bias_offset", biasShape, xShape, "base + i") %}
82
+ {% endif %}
83
+ {% set halfOutputScalar = "f16" %}
84
+ fn round_f16_bits_rte(value: f32) -> u32 {
85
+ let bits = bitcast<u32>(value);
86
+ let sign = (bits >> 16u) & 0x8000u;
87
+ let exponent_f32 = (bits >> 23u) & 0xffu;
88
+ let mantissa_f32 = bits & 0x7fffffu;
89
+
90
+ if (exponent_f32 == 0xffu) {
91
+ if (mantissa_f32 != 0u) {
92
+ return 0x7e00u;
93
+ }
94
+ return sign | 0x7c00u;
95
+ }
96
+
97
+ var exponent_f16 = i32(exponent_f32) - 127 + 15;
98
+ if (exponent_f16 >= 0x1f) {
99
+ return sign | 0x7c00u;
100
+ }
101
+
102
+ if (exponent_f16 <= 0) {
103
+ if (exponent_f16 < -10) {
104
+ return sign;
105
+ }
106
+ let significand = mantissa_f32 | 0x800000u;
107
+ let shift = u32(14 - exponent_f16);
108
+ let halfway = 1u << (shift - 1u);
109
+ let discarded = significand & ((1u << shift) - 1u);
110
+ var fraction = significand >> shift;
111
+ if (discarded > halfway || (discarded == halfway && (fraction & 1u) == 1u)) {
112
+ fraction = fraction + 1u;
113
+ }
114
+ return sign | fraction;
115
+ }
116
+
117
+ let halfway = 1u << 12u;
118
+ let discarded = mantissa_f32 & 0x1fffu;
119
+ var mantissa_f16 = mantissa_f32 >> 13u;
120
+ if (discarded > halfway || (discarded == halfway && (mantissa_f16 & 1u) == 1u)) {
121
+ mantissa_f16 = mantissa_f16 + 1u;
122
+ if (mantissa_f16 == 0x400u) {
123
+ mantissa_f16 = 0u;
124
+ exponent_f16 = exponent_f16 + 1;
125
+ }
126
+ }
127
+ if (exponent_f16 >= 0x1f) {
128
+ return sign | 0x7c00u;
129
+ }
130
+ return sign | (u32(exponent_f16) << 10u) | mantissa_f16;
131
+ }
132
+
133
+ fn widen_f16_bits(value: u32) -> f32 {
134
+ return unpack2x16float(value & 0xffffu).x;
135
+ }
136
+
137
+ // Typed ONNX edges must survive arithmetic fusion and narrow/wide casts.
138
+ // Integer rounding also fixes the ties-to-even rule independently of the
139
+ // implementation's floating-point conversion rounding mode.
140
+ fn half_stage(value: f32) -> f32 {
141
+ return widen_f16_bits(round_f16_bits_rte(value));
142
+ }
143
+
144
+ // Half output magnifies statistics errors at rounding midpoints. Keep a low
145
+ // residual through the reduction and normalization, then round the typed
146
+ // float32-normalized/half-scale/half-bias edges explicitly. This is the standard ONNX half path;
147
+ // other normalization contracts retain their existing arithmetic.
148
+ {% set halfWriteMean = halfWriteMean if halfWriteMean is defined else (writeStats and modeSpec == "layer") %}
149
+ {% set halfWriteInv = halfWriteInv if halfWriteInv is defined else writeStats %}
150
+
151
+ fn pair_add(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> {
152
+ let s = fma(a.x, 1.0, b.x);
153
+ // Materialize each rounded subtraction in the error-free transform. Plain
154
+ // cancellation expressions do not preserve the intended evaluation tree on
155
+ // every shader backend.
156
+ let bv = fma(-1.0, a.x, s);
157
+ let av = fma(-1.0, bv, s);
158
+ let a_error = fma(-1.0, av, a.x);
159
+ let b_error = fma(-1.0, bv, b.x);
160
+ let error = fma(a_error, 1.0, b_error);
161
+ let e = fma(fma(error, 1.0, a.y), 1.0, b.y);
162
+ let hi = fma(s, 1.0, e);
163
+ return vec2<f32>(hi, fma(-1.0, fma(-1.0, s, hi), e));
164
+ }
165
+
166
+ fn pair_mul(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> {
167
+ let p = fma(a.x, b.x, 0.0);
168
+ let error = fma(a.x, b.x, -p);
169
+ return pair_add(vec2<f32>(p, 0.0), vec2<f32>(error + (a.x * b.y + b.x * a.y), 0.0));
170
+ }
171
+
172
+ fn pair_div(a: vec2<f32>, b: f32) -> vec2<f32> {
173
+ let q = a.x / b;
174
+ let residual = pair_add(a, -pair_mul(vec2<f32>(q, 0.0), vec2<f32>(b, 0.0)));
175
+ return pair_add(vec2<f32>(q, 0.0), vec2<f32>((residual.x + residual.y) / b, 0.0));
176
+ }
177
+
178
+ fn pair_inverse_sqrt(a: vec2<f32>) -> vec2<f32> {
179
+ let r = inverseSqrt(a.x);
180
+ let rr = pair_mul(vec2<f32>(r, 0.0), vec2<f32>(r, 0.0));
181
+ let residual = pair_add(vec2<f32>(1.0, 0.0), -pair_mul(a, rr));
182
+ return pair_add(vec2<f32>(r, 0.0), vec2<f32>((0.5 * r) * (residual.x + residual.y), 0.0));
183
+ }
184
+
185
+ fn half_normalized(value: vec2<f32>) -> f32 {
186
+ // stash_type=1 materializes Normalized as float32 before its cast to half.
187
+ // Collapse the compensated residual at that typed edge; rounding the pair
188
+ // directly to half can choose a different result at a float32 midpoint.
189
+ return widen_f16_bits(round_f16_bits_rte(fma(value.x, 1.0, value.y)));
190
+ }
191
+
192
+ var<workgroup> partial: array<vec2<f32>, WG>;
193
+ fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
194
+ partial[tid] = value;
195
+ workgroupBarrier();
196
+ for (var step = WG / 2u; step > 0u; step /= 2u) {
197
+ if (tid < step) { partial[tid] = pair_add(partial[tid], partial[tid + step]); }
198
+ workgroupBarrier();
199
+ }
200
+ let total = partial[0];
201
+ workgroupBarrier();
202
+ return total;
203
+ }
204
+ {% set reduceArgs = "tid" %}
205
 
206
+ fn load_value(index: u32) -> f32 {
207
+ return f32(x[index]);
208
+ }
209
+
210
+ fn normalize_half_row(row: u32, tid: u32
211
+ ) {
212
+ if (row >= params.rows) { return; }
213
+ let base = row * HIDDEN;
214
+ var local_sum = vec2<f32>(0.0);
215
+ for (var i = tid; i < HIDDEN; i += WG) {
216
+ local_sum = pair_add(local_sum, vec2<f32>(load_value(base + i), 0.0));
217
+ }
218
+ let mean = pair_div(reduce_pair(local_sum, {{ reduceArgs }}), f32(HIDDEN));
219
+ var local_square = vec2<f32>(0.0);
220
+ for (var i = tid; i < HIDDEN; i += WG) {
221
+ let centered = pair_add(vec2<f32>(load_value(base + i), 0.0), -mean);
222
+ local_square = pair_add(local_square, pair_mul(centered, centered));
223
+ }
224
+ let variance = pair_div(reduce_pair(local_square, {{ reduceArgs }}), f32(HIDDEN));
225
+ let inv = pair_inverse_sqrt(pair_add(variance, vec2<f32>(EPSILON, 0.0)));
226
+ {% if halfWriteMean %}
227
+ if (tid == 0u) { mean_out[row] = mean.x; }
228
+ {% endif %}
229
+ {% if halfWriteInv %}
230
+ if (tid == 0u) { inv_std_out[row] = inv.x; }
231
  {% endif %}
232
+ for (var i = tid; i < HIDDEN; i += WG) {
233
+ let normalized = half_normalized(pair_mul(pair_add(vec2<f32>(load_value(base + i), 0.0), -mean), inv));
234
+ var value = half_stage(normalized * f32(scale[{{ halfScaleOffset | default("i") }}]));
235
+ {% if modeSpec == "layer" and hasBias %}
236
+ value = half_stage(value + f32(bias[{{ halfBiasOffset | default("i") }}]));
237
+ {% endif %}
238
+ y[base + i] = {{ halfOutputScalar }}(value);
239
+ }
240
+ }
241
+ @compute @workgroup_size(WG, 1, 1)
242
+ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
243
+ let row = wg.x + wg.y * params.rowStride;
244
+ normalize_half_row(row, lid.x);
245
+ }
246
+ {% else %}
247
+ var<workgroup> partial: array<f32, WG>;
248
+ var<workgroup> row_mean: f32;
249
+ var<workgroup> row_inv: f32;
250
+
251
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
252
+ {% if op == "max" or op == "min" %}
253
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
254
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
255
+ {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false, reuse=false) %}
 
 
 
256
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
257
  loop {
 
 
 
 
258
  if ({{ svar }} == 0u) {
259
  break;
260
  }
 
 
 
 
 
261
  if ({{ idx }} < {{ svar }}) {
262
  {% for a in arrays %}
263
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
264
  {% endfor %}
265
  }
 
 
 
 
266
  {{ svar }} = {{ svar }} / 2u;
 
 
 
 
 
267
  workgroupBarrier();
268
+ }{% endmacro %}
 
 
 
 
 
 
 
269
  // Reusing partial after this reduction requires a barrier between the read of
270
  // partial[0] and the next write, or the next round can race the prior readers.
 
271
  fn reduce_sum(value: f32, tid: u32) -> f32 {
272
  partial[tid] = value;
273
  workgroupBarrier();
274
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
 
 
 
 
 
275
  return partial[0];
 
276
  }
277
 
 
278
  @compute @workgroup_size(WG, 1, 1)
279
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
280
  let row = wg.x + wg.y * params.rowStride;
 
325
  y[index] = {{ scalar }}(value);
326
  }
327
  }
328
+ {% endif %}
build/webgpu/manifest.json CHANGED
@@ -15,7 +15,7 @@
15
  "attributes": { "axis": { "default": -1 }, "epsilon": { "default": 0.00001 }, "stash_type": { "default": 1 } },
16
  "attributeConstraints": { "stash_type": { "values": [1] } },
17
  "typeConstraints": { "T": ["float32", "float16"] },
18
- "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 }, "SCALAR_FAST_MAX_HIDDEN": { "default": 1024 } },
19
  "derive": {
20
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
21
  "normWorkgroupCap": "min(tunables.MAX_WORKGROUP_SIZE, deviceWorkgroupCap)",
@@ -54,25 +54,82 @@
54
  },
55
  "when": ["f16Ok(dtypes.T)"],
56
  "bindings": {
57
- "x": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
58
- "scale": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
59
- "y": { "buffer": "storage", "elementType": "$vectorScalar" },
60
  "params": {
61
- "buffer": "uniform",
62
  "struct": [
63
  { "name": "rows", "type": "u32", "value": "normRows" },
64
  { "name": "rowStride", "type": "u32", "value": "normRowStride" }
65
  ]
66
  },
67
- "bias": { "arg": "b", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
68
- "mean_out": { "arg": "mean", "buffer": "storage", "elementType": "f32" },
69
- "inv_std_out": { "arg": "invStdDev", "buffer": "storage", "elementType": "f32" },
70
- "x_2": { "name": "x", "buffer": "read-only-storage", "elementType": "$scalar" },
71
- "scale_2": { "name": "scale", "buffer": "read-only-storage", "elementType": "$scalar" },
72
- "y_2": { "name": "y", "buffer": "storage", "elementType": "$scalar" },
73
- "bias_2": { "arg": "b", "name": "bias", "buffer": "read-only-storage", "elementType": "$scalar" }
74
  },
75
  "variants": [
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
76
  {
77
  "id": "last_axis_row_vec4",
78
  "priority": 110,
@@ -85,11 +142,10 @@
85
  "shader": "norm-row-stats.wgsl.jinja",
86
  "derive": {
87
  "modeSpec": "\"layer\"",
 
88
  "vec4": true,
89
- "hasBias": false,
90
- "writeStats": false,
91
- "scalar": "dtypes.T",
92
- "usesF16Spec": "dtypes.T == \"f16\"",
93
  "hidden": "dim(shapes.x, -1)",
94
  "wg": "lastAxisWgVec4",
95
  "epsilon": "attrs.epsilon",
@@ -98,8 +154,7 @@
98
  "combineSubgroups": "hasSubgroupId"
99
  },
100
  "bindings": ["x", "scale", "y", "params"],
101
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
102
- "subgroupCollectivesWidth": "portable"
103
  }
104
  ]
105
  },
@@ -116,19 +171,17 @@
116
  "shader": "norm-row-stats.wgsl.jinja",
117
  "derive": {
118
  "modeSpec": "\"layer\"",
 
119
  "vec4": false,
120
- "hasBias": false,
121
- "writeStats": false,
122
- "scalar": "dtypes.T",
123
- "usesF16Spec": "dtypes.T == \"f16\"",
124
  "hidden": "dim(shapes.x, -1)",
125
  "wg": "lastAxisWg",
126
  "epsilon": "attrs.epsilon",
127
  "combineSubgroups": "hasSubgroupId"
128
  },
129
  "bindings": ["x", "scale", "y", "params"],
130
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
131
- "subgroupCollectivesWidth": "portable"
132
  }
133
  ]
134
  },
@@ -144,11 +197,10 @@
144
  "shader": "norm-row-stats.wgsl.jinja",
145
  "derive": {
146
  "modeSpec": "\"layer\"",
 
147
  "vec4": true,
148
- "hasBias": true,
149
- "writeStats": false,
150
- "scalar": "dtypes.T",
151
- "usesF16Spec": "dtypes.T == \"f16\"",
152
  "hidden": "dim(shapes.x, -1)",
153
  "wg": "lastAxisWgVec4",
154
  "epsilon": "attrs.epsilon",
@@ -157,8 +209,7 @@
157
  "combineSubgroups": "hasSubgroupId"
158
  },
159
  "bindings": ["x", "scale", "bias", "y", "params"],
160
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
161
- "subgroupCollectivesWidth": "portable"
162
  }
163
  ]
164
  },
@@ -175,19 +226,17 @@
175
  "shader": "norm-row-stats.wgsl.jinja",
176
  "derive": {
177
  "modeSpec": "\"layer\"",
 
178
  "vec4": false,
179
- "hasBias": true,
180
- "writeStats": false,
181
- "scalar": "dtypes.T",
182
- "usesF16Spec": "dtypes.T == \"f16\"",
183
  "hidden": "dim(shapes.x, -1)",
184
  "wg": "lastAxisWg",
185
  "epsilon": "attrs.epsilon",
186
  "combineSubgroups": "hasSubgroupId"
187
  },
188
  "bindings": ["x", "scale", "bias", "y", "params"],
189
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
190
- "subgroupCollectivesWidth": "portable"
191
  }
192
  ]
193
  },
@@ -203,11 +252,10 @@
203
  "shader": "norm-row-stats.wgsl.jinja",
204
  "derive": {
205
  "modeSpec": "\"layer\"",
 
206
  "vec4": true,
207
- "hasBias": false,
208
- "writeStats": true,
209
- "scalar": "dtypes.T",
210
- "usesF16Spec": "dtypes.T == \"f16\"",
211
  "hidden": "dim(shapes.x, -1)",
212
  "wg": "lastAxisWgVec4",
213
  "epsilon": "attrs.epsilon",
@@ -216,8 +264,7 @@
216
  "combineSubgroups": "hasSubgroupId"
217
  },
218
  "bindings": ["x", "scale", "y", "mean_out", "inv_std_out", "params"],
219
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
220
- "subgroupCollectivesWidth": "portable"
221
  }
222
  ]
223
  },
@@ -234,19 +281,17 @@
234
  "shader": "norm-row-stats.wgsl.jinja",
235
  "derive": {
236
  "modeSpec": "\"layer\"",
 
237
  "vec4": false,
238
- "hasBias": false,
239
- "writeStats": true,
240
- "scalar": "dtypes.T",
241
- "usesF16Spec": "dtypes.T == \"f16\"",
242
  "hidden": "dim(shapes.x, -1)",
243
  "wg": "lastAxisWg",
244
  "epsilon": "attrs.epsilon",
245
  "combineSubgroups": "hasSubgroupId"
246
  },
247
  "bindings": ["x", "scale", "y", "mean_out", "inv_std_out", "params"],
248
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
249
- "subgroupCollectivesWidth": "portable"
250
  }
251
  ]
252
  },
@@ -262,11 +307,10 @@
262
  "shader": "norm-row-stats.wgsl.jinja",
263
  "derive": {
264
  "modeSpec": "\"layer\"",
 
265
  "vec4": true,
266
- "hasBias": true,
267
- "writeStats": true,
268
- "scalar": "dtypes.T",
269
- "usesF16Spec": "dtypes.T == \"f16\"",
270
  "hidden": "dim(shapes.x, -1)",
271
  "wg": "lastAxisWgVec4",
272
  "epsilon": "attrs.epsilon",
@@ -275,8 +319,7 @@
275
  "combineSubgroups": "hasSubgroupId"
276
  },
277
  "bindings": ["x", "scale", "bias", "y", "mean_out", "inv_std_out", "params"],
278
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
279
- "subgroupCollectivesWidth": "portable"
280
  }
281
  ]
282
  },
@@ -293,19 +336,17 @@
293
  "shader": "norm-row-stats.wgsl.jinja",
294
  "derive": {
295
  "modeSpec": "\"layer\"",
 
296
  "vec4": false,
297
- "hasBias": true,
298
- "writeStats": true,
299
- "scalar": "dtypes.T",
300
- "usesF16Spec": "dtypes.T == \"f16\"",
301
  "hidden": "dim(shapes.x, -1)",
302
  "wg": "lastAxisWg",
303
  "epsilon": "attrs.epsilon",
304
  "combineSubgroups": "hasSubgroupId"
305
  },
306
  "bindings": ["x", "scale", "bias", "y", "mean_out", "inv_std_out", "params"],
307
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
308
- "subgroupCollectivesWidth": "portable"
309
  }
310
  ]
311
  },
@@ -321,11 +362,10 @@
321
  "shader": "norm-row-stats.wgsl.jinja",
322
  "derive": {
323
  "modeSpec": "\"layer\"",
 
324
  "vec4": true,
325
  "hasBias": true,
326
  "writeStats": false,
327
- "scalar": "dtypes.T",
328
- "usesF16Spec": "dtypes.T == \"f16\"",
329
  "hidden": "suffixAxisSize",
330
  "wg": "suffixAxisWgVec4",
331
  "epsilon": "attrs.epsilon",
@@ -334,8 +374,7 @@
334
  "combineSubgroups": "hasSubgroupId"
335
  },
336
  "bindings": ["x", "scale", "bias", "y", "params"],
337
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
338
- "subgroupCollectivesWidth": "portable"
339
  }
340
  ]
341
  },
@@ -544,9 +583,9 @@
544
  "priority": 31,
545
  "when": ["not present.b and meanOnlyOutputs and meanRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
546
  "derive": {
547
- "hasBias": false,
548
- "writeMean": true,
549
- "writeInvStdDev": false,
550
  "scalar": "dtypes.T",
551
  "hiddenSize": "genericHiddenSize",
552
  "workgroupSize": "genericWorkgroupSize",
@@ -558,7 +597,7 @@
558
  "name": "LayerNormalization.MeanOnly",
559
  "shader": "layer-normalization.wgsl.jinja",
560
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale" },
561
- "bindings": ["x_2", "scale_2", "y_2", "mean_out", "params"],
562
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
563
  }
564
  ]
@@ -568,9 +607,9 @@
568
  "priority": 32,
569
  "when": ["present.b and meanOnlyOutputs and biasBroadcastOk and meanRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
570
  "derive": {
571
- "hasBias": true,
572
- "writeMean": true,
573
- "writeInvStdDev": false,
574
  "scalar": "dtypes.T",
575
  "hiddenSize": "genericHiddenSize",
576
  "workgroupSize": "genericWorkgroupSize",
@@ -582,7 +621,7 @@
582
  "name": "LayerNormalization.BiasMeanOnly",
583
  "shader": "layer-normalization.wgsl.jinja",
584
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale", "biasShape": "shapes.b" },
585
- "bindings": ["x_2", "scale_2", "bias_2", "y_2", "mean_out", "params"],
586
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
587
  }
588
  ]
@@ -592,9 +631,9 @@
592
  "priority": 33,
593
  "when": ["not present.b and invStdOnlyOutputs and invStdRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
594
  "derive": {
595
- "hasBias": false,
596
- "writeMean": false,
597
- "writeInvStdDev": true,
598
  "scalar": "dtypes.T",
599
  "hiddenSize": "genericHiddenSize",
600
  "workgroupSize": "genericWorkgroupSize",
@@ -606,7 +645,7 @@
606
  "name": "LayerNormalization.InvStdDevOnly",
607
  "shader": "layer-normalization.wgsl.jinja",
608
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale" },
609
- "bindings": ["x_2", "scale_2", "y_2", "inv_std_out", "params"],
610
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
611
  }
612
  ]
@@ -616,9 +655,9 @@
616
  "priority": 34,
617
  "when": ["present.b and invStdOnlyOutputs and biasBroadcastOk and invStdRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
618
  "derive": {
619
- "hasBias": true,
620
- "writeMean": false,
621
- "writeInvStdDev": true,
622
  "scalar": "dtypes.T",
623
  "hiddenSize": "genericHiddenSize",
624
  "workgroupSize": "genericWorkgroupSize",
@@ -630,10 +669,154 @@
630
  "name": "LayerNormalization.BiasInvStdDevOnly",
631
  "shader": "layer-normalization.wgsl.jinja",
632
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale", "biasShape": "shapes.b" },
633
- "bindings": ["x_2", "scale_2", "bias_2", "y_2", "inv_std_out", "params"],
634
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
635
  }
636
  ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
637
  }
638
  ]
639
  }
 
15
  "attributes": { "axis": { "default": -1 }, "epsilon": { "default": 0.00001 }, "stash_type": { "default": 1 } },
16
  "attributeConstraints": { "stash_type": { "values": [1] } },
17
  "typeConstraints": { "T": ["float32", "float16"] },
18
+ "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 }, "SCALAR_FAST_MAX_HIDDEN": { "default": 4096 } },
19
  "derive": {
20
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
21
  "normWorkgroupCap": "min(tunables.MAX_WORKGROUP_SIZE, deviceWorkgroupCap)",
 
54
  },
55
  "when": ["f16Ok(dtypes.T)"],
56
  "bindings": {
57
+ "x": { "elementType": "$vectorScalar" },
58
+ "scale": { "elementType": "$vectorScalar" },
59
+ "y": { "elementType": "$vectorScalar" },
60
  "params": {
 
61
  "struct": [
62
  { "name": "rows", "type": "u32", "value": "normRows" },
63
  { "name": "rowStride", "type": "u32", "value": "normRowStride" }
64
  ]
65
  },
66
+ "bias": { "arg": "b", "elementType": "$vectorScalar" },
67
+ "mean_out": { "arg": "mean", "elementType": "f32" },
68
+ "inv_std_out": { "arg": "invStdDev", "elementType": "f32" },
69
+ "x_main": { "name": "x", "elementType": "$scalar" },
70
+ "scale_main": { "name": "scale", "elementType": "$scalar" },
71
+ "y_main": { "name": "y", "elementType": "$scalar" },
72
+ "bias_b": { "arg": "b", "name": "bias", "elementType": "$scalar" }
73
  },
74
  "variants": [
75
+ {
76
+ "id": "last_axis_row_vec2",
77
+ "priority": 105,
78
+ "when": ["dtypes.T == \"f32\"", "lastAxisExactScaleOk", "dim(shapes.x, -1) % 4 == 2", "noStatsOutputs", "not present.b"],
79
+ "derive": { "scalar": "dtypes.T", "vectorScalar": "\"vec2<f32>\"" },
80
+ "passes": [
81
+ {
82
+ "id": "main",
83
+ "name": "LayerNormalization.LastAxisRowVec2",
84
+ "shader": "norm-row-stats.wgsl.jinja",
85
+ "derive": {
86
+ "modeSpec": "\"layer\"",
87
+ "vec4": true,
88
+ "hasBias": "present.b",
89
+ "writeStats": false,
90
+ "hidden": "dim(shapes.x, -1)",
91
+ "wg": "min(normWorkgroupCap, pow2ceil(dim(shapes.x, -1) / 2))",
92
+ "epsilon": "attrs.epsilon",
93
+ "hiddenVec": "dim(shapes.x, -1) / 2",
94
+ "vecType": "\"vec2<f32>\"",
95
+ "combineSubgroups": "hasSubgroupId",
96
+ "packedWidth": 2
97
+ },
98
+ "bindings": ["x", "scale", "y", "params"],
99
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
100
+ }
101
+ ],
102
+ "demoteWhen": ["not lastAxisScalarFastOk"]
103
+ },
104
+ {
105
+ "id": "last_axis_bias_row_vec2",
106
+ "priority": 106,
107
+ "when": ["dtypes.T == \"f32\"", "lastAxisExactScaleOk", "dim(shapes.x, -1) % 4 == 2", "noStatsOutputs", "present.b and biasExactOk"],
108
+ "derive": { "scalar": "dtypes.T", "vectorScalar": "\"vec2<f32>\"" },
109
+ "passes": [
110
+ {
111
+ "id": "main",
112
+ "name": "LayerNormalization.LastAxisRowVec2",
113
+ "shader": "norm-row-stats.wgsl.jinja",
114
+ "derive": {
115
+ "modeSpec": "\"layer\"",
116
+ "vec4": true,
117
+ "hasBias": "present.b",
118
+ "writeStats": false,
119
+ "hidden": "dim(shapes.x, -1)",
120
+ "wg": "min(normWorkgroupCap, pow2ceil(dim(shapes.x, -1) / 2))",
121
+ "epsilon": "attrs.epsilon",
122
+ "hiddenVec": "dim(shapes.x, -1) / 2",
123
+ "vecType": "\"vec2<f32>\"",
124
+ "combineSubgroups": "hasSubgroupId",
125
+ "packedWidth": 2
126
+ },
127
+ "bindings": ["x", "scale", "bias", "y", "params"],
128
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
129
+ }
130
+ ],
131
+ "demoteWhen": ["not lastAxisScalarFastOk"]
132
+ },
133
  {
134
  "id": "last_axis_row_vec4",
135
  "priority": 110,
 
142
  "shader": "norm-row-stats.wgsl.jinja",
143
  "derive": {
144
  "modeSpec": "\"layer\"",
145
+ "compensateHalfStats": true,
146
  "vec4": true,
147
+ "hasBias": "present.b",
148
+ "writeStats": "fullStatsOutputs",
 
 
149
  "hidden": "dim(shapes.x, -1)",
150
  "wg": "lastAxisWgVec4",
151
  "epsilon": "attrs.epsilon",
 
154
  "combineSubgroups": "hasSubgroupId"
155
  },
156
  "bindings": ["x", "scale", "y", "params"],
157
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
158
  }
159
  ]
160
  },
 
171
  "shader": "norm-row-stats.wgsl.jinja",
172
  "derive": {
173
  "modeSpec": "\"layer\"",
174
+ "compensateHalfStats": true,
175
  "vec4": false,
176
+ "hasBias": "present.b",
177
+ "writeStats": "fullStatsOutputs",
 
 
178
  "hidden": "dim(shapes.x, -1)",
179
  "wg": "lastAxisWg",
180
  "epsilon": "attrs.epsilon",
181
  "combineSubgroups": "hasSubgroupId"
182
  },
183
  "bindings": ["x", "scale", "y", "params"],
184
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
185
  }
186
  ]
187
  },
 
197
  "shader": "norm-row-stats.wgsl.jinja",
198
  "derive": {
199
  "modeSpec": "\"layer\"",
200
+ "compensateHalfStats": true,
201
  "vec4": true,
202
+ "hasBias": "present.b",
203
+ "writeStats": "fullStatsOutputs",
 
 
204
  "hidden": "dim(shapes.x, -1)",
205
  "wg": "lastAxisWgVec4",
206
  "epsilon": "attrs.epsilon",
 
209
  "combineSubgroups": "hasSubgroupId"
210
  },
211
  "bindings": ["x", "scale", "bias", "y", "params"],
212
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
213
  }
214
  ]
215
  },
 
226
  "shader": "norm-row-stats.wgsl.jinja",
227
  "derive": {
228
  "modeSpec": "\"layer\"",
229
+ "compensateHalfStats": true,
230
  "vec4": false,
231
+ "hasBias": "present.b",
232
+ "writeStats": "fullStatsOutputs",
 
 
233
  "hidden": "dim(shapes.x, -1)",
234
  "wg": "lastAxisWg",
235
  "epsilon": "attrs.epsilon",
236
  "combineSubgroups": "hasSubgroupId"
237
  },
238
  "bindings": ["x", "scale", "bias", "y", "params"],
239
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
240
  }
241
  ]
242
  },
 
252
  "shader": "norm-row-stats.wgsl.jinja",
253
  "derive": {
254
  "modeSpec": "\"layer\"",
255
+ "compensateHalfStats": true,
256
  "vec4": true,
257
+ "hasBias": "present.b",
258
+ "writeStats": "fullStatsOutputs",
 
 
259
  "hidden": "dim(shapes.x, -1)",
260
  "wg": "lastAxisWgVec4",
261
  "epsilon": "attrs.epsilon",
 
264
  "combineSubgroups": "hasSubgroupId"
265
  },
266
  "bindings": ["x", "scale", "y", "mean_out", "inv_std_out", "params"],
267
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
268
  }
269
  ]
270
  },
 
281
  "shader": "norm-row-stats.wgsl.jinja",
282
  "derive": {
283
  "modeSpec": "\"layer\"",
284
+ "compensateHalfStats": true,
285
  "vec4": false,
286
+ "hasBias": "present.b",
287
+ "writeStats": "fullStatsOutputs",
 
 
288
  "hidden": "dim(shapes.x, -1)",
289
  "wg": "lastAxisWg",
290
  "epsilon": "attrs.epsilon",
291
  "combineSubgroups": "hasSubgroupId"
292
  },
293
  "bindings": ["x", "scale", "y", "mean_out", "inv_std_out", "params"],
294
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
295
  }
296
  ]
297
  },
 
307
  "shader": "norm-row-stats.wgsl.jinja",
308
  "derive": {
309
  "modeSpec": "\"layer\"",
310
+ "compensateHalfStats": true,
311
  "vec4": true,
312
+ "hasBias": "present.b",
313
+ "writeStats": "fullStatsOutputs",
 
 
314
  "hidden": "dim(shapes.x, -1)",
315
  "wg": "lastAxisWgVec4",
316
  "epsilon": "attrs.epsilon",
 
319
  "combineSubgroups": "hasSubgroupId"
320
  },
321
  "bindings": ["x", "scale", "bias", "y", "mean_out", "inv_std_out", "params"],
322
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
323
  }
324
  ]
325
  },
 
336
  "shader": "norm-row-stats.wgsl.jinja",
337
  "derive": {
338
  "modeSpec": "\"layer\"",
339
+ "compensateHalfStats": true,
340
  "vec4": false,
341
+ "hasBias": "present.b",
342
+ "writeStats": "fullStatsOutputs",
 
 
343
  "hidden": "dim(shapes.x, -1)",
344
  "wg": "lastAxisWg",
345
  "epsilon": "attrs.epsilon",
346
  "combineSubgroups": "hasSubgroupId"
347
  },
348
  "bindings": ["x", "scale", "bias", "y", "mean_out", "inv_std_out", "params"],
349
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
350
  }
351
  ]
352
  },
 
362
  "shader": "norm-row-stats.wgsl.jinja",
363
  "derive": {
364
  "modeSpec": "\"layer\"",
365
+ "compensateHalfStats": true,
366
  "vec4": true,
367
  "hasBias": true,
368
  "writeStats": false,
 
 
369
  "hidden": "suffixAxisSize",
370
  "wg": "suffixAxisWgVec4",
371
  "epsilon": "attrs.epsilon",
 
374
  "combineSubgroups": "hasSubgroupId"
375
  },
376
  "bindings": ["x", "scale", "bias", "y", "params"],
377
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
378
  }
379
  ]
380
  },
 
583
  "priority": 31,
584
  "when": ["not present.b and meanOnlyOutputs and meanRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
585
  "derive": {
586
+ "hasBias": "present.b",
587
+ "writeMean": "present.mean",
588
+ "writeInvStdDev": "present.invStdDev",
589
  "scalar": "dtypes.T",
590
  "hiddenSize": "genericHiddenSize",
591
  "workgroupSize": "genericWorkgroupSize",
 
597
  "name": "LayerNormalization.MeanOnly",
598
  "shader": "layer-normalization.wgsl.jinja",
599
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale" },
600
+ "bindings": ["x_main", "scale_main", "y_main", "mean_out", "params"],
601
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
602
  }
603
  ]
 
607
  "priority": 32,
608
  "when": ["present.b and meanOnlyOutputs and biasBroadcastOk and meanRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
609
  "derive": {
610
+ "hasBias": "present.b",
611
+ "writeMean": "present.mean",
612
+ "writeInvStdDev": "present.invStdDev",
613
  "scalar": "dtypes.T",
614
  "hiddenSize": "genericHiddenSize",
615
  "workgroupSize": "genericWorkgroupSize",
 
621
  "name": "LayerNormalization.BiasMeanOnly",
622
  "shader": "layer-normalization.wgsl.jinja",
623
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale", "biasShape": "shapes.b" },
624
+ "bindings": ["x_main", "scale_main", "bias_b", "y_main", "mean_out", "params"],
625
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
626
  }
627
  ]
 
631
  "priority": 33,
632
  "when": ["not present.b and invStdOnlyOutputs and invStdRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
633
  "derive": {
634
+ "hasBias": "present.b",
635
+ "writeMean": "present.mean",
636
+ "writeInvStdDev": "present.invStdDev",
637
  "scalar": "dtypes.T",
638
  "hiddenSize": "genericHiddenSize",
639
  "workgroupSize": "genericWorkgroupSize",
 
645
  "name": "LayerNormalization.InvStdDevOnly",
646
  "shader": "layer-normalization.wgsl.jinja",
647
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale" },
648
+ "bindings": ["x_main", "scale_main", "y_main", "inv_std_out", "params"],
649
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
650
  }
651
  ]
 
655
  "priority": 34,
656
  "when": ["present.b and invStdOnlyOutputs and biasBroadcastOk and invStdRowsOk", "lastAxisBroadcastScaleOk or suffixAxisBroadcastScaleOk"],
657
  "derive": {
658
+ "hasBias": "present.b",
659
+ "writeMean": "present.mean",
660
+ "writeInvStdDev": "present.invStdDev",
661
  "scalar": "dtypes.T",
662
  "hiddenSize": "genericHiddenSize",
663
  "workgroupSize": "genericWorkgroupSize",
 
669
  "name": "LayerNormalization.BiasInvStdDevOnly",
670
  "shader": "layer-normalization.wgsl.jinja",
671
  "derive": { "xShape": "shapes.x", "scaleShape": "shapes.scale", "biasShape": "shapes.b" },
672
+ "bindings": ["x_main", "scale_main", "bias_b", "y_main", "inv_std_out", "params"],
673
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
674
  }
675
  ]
676
+ },
677
+ {
678
+ "id": "last_axis_broadcast_row_vec4",
679
+ "priority": 90,
680
+ "when": ["dtypes.T == \"f32\"", "lastAxisBroadcastScaleOk", "dim(shapes.x, -1) >= 4", "dim(shapes.x, -1) % 4 == 0", "ranks.scale >= 1", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "noStatsOutputs", "not present.b", "true"],
681
+ "demoteWhen": ["false"],
682
+ "derive": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
683
+ "passes": [
684
+ {
685
+ "id": "main",
686
+ "name": "LayerNormalization.BroadcastRowsVec4",
687
+ "shader": "norm-row-stats.wgsl.jinja",
688
+ "derive": {
689
+ "modeSpec": "\"layer\"",
690
+ "vec4": true,
691
+ "hasBias": "present.b",
692
+ "writeStats": false,
693
+ "hidden": "dim(shapes.x, -1)",
694
+ "wg": "lastAxisWgVec4",
695
+ "epsilon": "attrs.epsilon",
696
+ "hiddenVec": "dim(shapes.x, -1) / 4",
697
+ "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"",
698
+ "combineSubgroups": "hasSubgroupId",
699
+ "affineRowBroadcast": true,
700
+ "xRowShape": "prefix(shapes.x, ranks.x - 1)",
701
+ "scaleRowShape": "prefix(shapes.scale, ranks.scale - 1)",
702
+ "biasRowShape": "prefix(shapes.b, ranks.b - 1) if present.b else []",
703
+ "batchRows": "1",
704
+ "batchLanes": "lastAxisWgVec4"
705
+ },
706
+ "bindings": ["x", "scale", "y", "params"],
707
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
708
+ }
709
+ ]
710
+ },
711
+ {
712
+ "id": "last_axis_broadcast_bias_row_vec4",
713
+ "priority": 91,
714
+ "when": ["dtypes.T == \"f32\"", "lastAxisBroadcastScaleOk", "dim(shapes.x, -1) >= 4", "dim(shapes.x, -1) % 4 == 0", "ranks.scale >= 1", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "noStatsOutputs", "present.b and biasBroadcastOk and ranks.b >= 1 and dim(shapes.b, -1) == dim(shapes.x, -1)", "true"],
715
+ "demoteWhen": ["false"],
716
+ "derive": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
717
+ "passes": [
718
+ {
719
+ "id": "main",
720
+ "name": "LayerNormalization.BroadcastRowsVec4",
721
+ "shader": "norm-row-stats.wgsl.jinja",
722
+ "derive": {
723
+ "modeSpec": "\"layer\"",
724
+ "vec4": true,
725
+ "hasBias": "present.b",
726
+ "writeStats": false,
727
+ "hidden": "dim(shapes.x, -1)",
728
+ "wg": "lastAxisWgVec4",
729
+ "epsilon": "attrs.epsilon",
730
+ "hiddenVec": "dim(shapes.x, -1) / 4",
731
+ "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"",
732
+ "combineSubgroups": "hasSubgroupId",
733
+ "affineRowBroadcast": true,
734
+ "xRowShape": "prefix(shapes.x, ranks.x - 1)",
735
+ "scaleRowShape": "prefix(shapes.scale, ranks.scale - 1)",
736
+ "biasRowShape": "prefix(shapes.b, ranks.b - 1) if present.b else []",
737
+ "batchRows": "1",
738
+ "batchLanes": "lastAxisWgVec4"
739
+ },
740
+ "bindings": ["x", "scale", "bias", "y", "params"],
741
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
742
+ }
743
+ ]
744
+ },
745
+ {
746
+ "id": "last_axis_broadcast_rows_vec4",
747
+ "priority": 95,
748
+ "when": ["dtypes.T == \"f32\"", "lastAxisBroadcastScaleOk", "dim(shapes.x, -1) >= 4", "dim(shapes.x, -1) % 4 == 0", "ranks.scale >= 1", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "noStatsOutputs", "not present.b", "floor(normWorkgroupCap / lastAxisWgVec4) > 1 and normWorkgroupCap * 8 <= device.limits.maxComputeWorkgroupStorageSize"],
749
+ "demoteWhen": ["normRows < floor(normWorkgroupCap / lastAxisWgVec4)"],
750
+ "derive": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
751
+ "passes": [
752
+ {
753
+ "id": "main",
754
+ "name": "LayerNormalization.BroadcastRowsVec4",
755
+ "shader": "norm-row-stats.wgsl.jinja",
756
+ "derive": {
757
+ "modeSpec": "\"layer\"",
758
+ "vec4": true,
759
+ "hasBias": "present.b",
760
+ "writeStats": false,
761
+ "hidden": "dim(shapes.x, -1)",
762
+ "wg": "lastAxisWgVec4 * floor(normWorkgroupCap / lastAxisWgVec4)",
763
+ "epsilon": "attrs.epsilon",
764
+ "hiddenVec": "dim(shapes.x, -1) / 4",
765
+ "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"",
766
+ "combineSubgroups": "false",
767
+ "affineRowBroadcast": true,
768
+ "xRowShape": "prefix(shapes.x, ranks.x - 1)",
769
+ "scaleRowShape": "prefix(shapes.scale, ranks.scale - 1)",
770
+ "biasRowShape": "prefix(shapes.b, ranks.b - 1) if present.b else []",
771
+ "batchRows": "floor(normWorkgroupCap / lastAxisWgVec4)",
772
+ "batchLanes": "lastAxisWgVec4"
773
+ },
774
+ "bindings": ["x", "scale", "y", "params"],
775
+ "dispatch": {
776
+ "x": "min(ceilDiv(normRows, floor(normWorkgroupCap / lastAxisWgVec4)), 65535)",
777
+ "y": "ceilDiv(ceilDiv(normRows, floor(normWorkgroupCap / lastAxisWgVec4)), 65535)",
778
+ "z": 1
779
+ }
780
+ }
781
+ ]
782
+ },
783
+ {
784
+ "id": "last_axis_broadcast_bias_rows_vec4",
785
+ "priority": 96,
786
+ "when": ["dtypes.T == \"f32\"", "lastAxisBroadcastScaleOk", "dim(shapes.x, -1) >= 4", "dim(shapes.x, -1) % 4 == 0", "ranks.scale >= 1", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "noStatsOutputs", "present.b and biasBroadcastOk and ranks.b >= 1 and dim(shapes.b, -1) == dim(shapes.x, -1)", "floor(normWorkgroupCap / lastAxisWgVec4) > 1 and normWorkgroupCap * 8 <= device.limits.maxComputeWorkgroupStorageSize"],
787
+ "demoteWhen": ["normRows < floor(normWorkgroupCap / lastAxisWgVec4)"],
788
+ "derive": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
789
+ "passes": [
790
+ {
791
+ "id": "main",
792
+ "name": "LayerNormalization.BroadcastRowsVec4",
793
+ "shader": "norm-row-stats.wgsl.jinja",
794
+ "derive": {
795
+ "modeSpec": "\"layer\"",
796
+ "vec4": true,
797
+ "hasBias": "present.b",
798
+ "writeStats": false,
799
+ "hidden": "dim(shapes.x, -1)",
800
+ "wg": "lastAxisWgVec4 * floor(normWorkgroupCap / lastAxisWgVec4)",
801
+ "epsilon": "attrs.epsilon",
802
+ "hiddenVec": "dim(shapes.x, -1) / 4",
803
+ "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"",
804
+ "combineSubgroups": "false",
805
+ "affineRowBroadcast": true,
806
+ "xRowShape": "prefix(shapes.x, ranks.x - 1)",
807
+ "scaleRowShape": "prefix(shapes.scale, ranks.scale - 1)",
808
+ "biasRowShape": "prefix(shapes.b, ranks.b - 1) if present.b else []",
809
+ "batchRows": "floor(normWorkgroupCap / lastAxisWgVec4)",
810
+ "batchLanes": "lastAxisWgVec4"
811
+ },
812
+ "bindings": ["x", "scale", "bias", "y", "params"],
813
+ "dispatch": {
814
+ "x": "min(ceilDiv(normRows, floor(normWorkgroupCap / lastAxisWgVec4)), 65535)",
815
+ "y": "ceilDiv(ceilDiv(normRows, floor(normWorkgroupCap / lastAxisWgVec4)), 65535)",
816
+ "z": 1
817
+ }
818
+ }
819
+ ]
820
  }
821
  ]
822
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.LayerNormalization",
3
- "id": "_ai_onnx_layernormalization_webgpu_7b13eb1",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,16 +8,18 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "IPZzLq64+ycfl0fAzym0hDLorrSrZ5YDwLnpVEGgGHc=",
11
- "layer-normalization.wgsl.jinja": "NJ1/CeeYHnToR+Ki5VHv4U4gxuHKa9zYg9PG2QOndME=",
12
- "manifest.json": "aZkl41hMG0XVB5WthNHRcxLOw39oCNULp/kF2RczI8Y=",
13
- "norm-row-stats.wgsl.jinja": "HUUqntKH7vqbffRpSu33tnmSU7PudcOCtl5QhV1xlGk=",
14
- "test.json": "qJxDvb9POxT3rI+vLcahsq2mqIwAfen9b53MgOSG5Go="
15
  }
16
  },
17
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
18
  "webgpu": {
19
- "manifestSpec": "2.0",
20
  "variants": {
 
 
21
  "last_axis_row_vec4": ["norm-row-stats.wgsl.jinja"],
22
  "last_axis_row": ["norm-row-stats.wgsl.jinja"],
23
  "last_axis_bias_row_vec4": ["norm-row-stats.wgsl.jinja"],
@@ -38,7 +40,11 @@
38
  "mean_only": ["layer-normalization.wgsl.jinja"],
39
  "bias_mean_only": ["layer-normalization.wgsl.jinja"],
40
  "inv_std_dev_only": ["layer-normalization.wgsl.jinja"],
41
- "bias_inv_std_dev_only": ["layer-normalization.wgsl.jinja"]
 
 
 
 
42
  }
43
  }
44
  }
 
1
  {
2
  "name": "ai.onnx.LayerNormalization",
3
+ "id": "_ai_onnx_layernormalization_webgpu_cb0cade",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "IPZzLq64+ycfl0fAzym0hDLorrSrZ5YDwLnpVEGgGHc=",
11
+ "layer-normalization.wgsl.jinja": "7h2F0b9zKAOqO4f5hhaJiX/TKGURziAlfQWUpO4IrWk=",
12
+ "manifest.json": "GWp3ACAaJzLWehyonJd7wD1r+5FDN4j8MJMiFV55SV8=",
13
+ "norm-row-stats.wgsl.jinja": "66BZfZ7q6X6xnqSuDyWvUov8wSdBYvxkmmfuqUl9uSk=",
14
+ "test.json": "qdUKVCiVNSQXTPzMr4JSIXCceUNzOoV4a5bLM2LNSo0="
15
  }
16
  },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
  "webgpu": {
19
+ "manifestSpec": "2.1",
20
  "variants": {
21
+ "last_axis_row_vec2": ["norm-row-stats.wgsl.jinja"],
22
+ "last_axis_bias_row_vec2": ["norm-row-stats.wgsl.jinja"],
23
  "last_axis_row_vec4": ["norm-row-stats.wgsl.jinja"],
24
  "last_axis_row": ["norm-row-stats.wgsl.jinja"],
25
  "last_axis_bias_row_vec4": ["norm-row-stats.wgsl.jinja"],
 
40
  "mean_only": ["layer-normalization.wgsl.jinja"],
41
  "bias_mean_only": ["layer-normalization.wgsl.jinja"],
42
  "inv_std_dev_only": ["layer-normalization.wgsl.jinja"],
43
+ "bias_inv_std_dev_only": ["layer-normalization.wgsl.jinja"],
44
+ "last_axis_broadcast_row_vec4": ["norm-row-stats.wgsl.jinja"],
45
+ "last_axis_broadcast_bias_row_vec4": ["norm-row-stats.wgsl.jinja"],
46
+ "last_axis_broadcast_rows_vec4": ["norm-row-stats.wgsl.jinja"],
47
+ "last_axis_broadcast_bias_rows_vec4": ["norm-row-stats.wgsl.jinja"]
48
  }
49
  }
50
  }
build/webgpu/norm-row-stats.wgsl.jinja CHANGED
@@ -1,16 +1,6 @@
1
- {% if usesF16Spec %}
2
- enable f16;
3
- {% endif %}
4
- {% set combineSubgroups = combineSubgroups %}
5
- {% set scalarIo = scalarIo if scalarIo is defined else false %}
6
- {% set packedBf16Embedding = packedBf16Embedding if packedBf16Embedding is defined else false %}
7
- {% set writeStats = writeStats if writeStats is defined else false %}
8
- {% set rmsChainNorm = rmsChainNorm if rmsChainNorm is defined else false %}
9
- {% set hiddenPairs = hiddenPairs | default(0) %}
10
- {% set numRows = numRows | default(0) %}
11
- {% set epsilon = epsilon | default("0.0") %}
12
- {% set epsilon2 = epsilon2 | default("0.0") %}
13
- {% set hasBias = hasBias is defined and hasBias %}
14
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
15
  if combineSubgroups else ", tid: u32" %}
16
  {% set reduceThreadArguments = ", sg_lane, sg_id, num_sg"
@@ -19,68 +9,359 @@ enable f16;
19
  enable subgroups;
20
  {% endif %}
21
  {{ env.wgsl.resourceDeclarations }}
22
-
23
- // Workgroup-parallel single-pass row statistics + fused normalize/affine.
24
- //
25
- // One workgroup owns one contiguous normalization span ("row": a last-axis
26
- // row, an instance plane, or a channel group). Threads stride the row once,
27
- // accumulating (sum, sum_sq) simultaneously. Partials are reduced either with
28
- // subgroupAdd plus a shared-memory combine or with a portable shared-memory
29
- // tree, then every thread applies the fused normalize + affine write.
30
- //
31
- // Shifted moments avoid cancellation from a large common offset; scaling uses
32
- // inverseSqrt(variance + EPSILON).
33
- const HIDDEN: u32 = {{ hidden }}u;
34
- {% if vec4 %}
35
- const HIDDEN_V: u32 = {{ hiddenVec }}u;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36
  {% endif %}
37
- {% if packedBf16Embedding %}
38
- const HIDDEN_PAIRS: u32 = {{ hiddenPairs }}u;
39
- const NUM_ROWS: u32 = {{ numRows }}u;
 
40
  {% endif %}
41
- const WG: u32 = {{ wg }}u;
42
- const EPSILON: f32 = {{ epsilon }};
43
- {% if rmsChainNorm %}
44
- const EPSILON2: f32 = {{ epsilon2 }};
45
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
 
47
- {% if packedBf16Embedding %}
48
- {% if vec4 %}
49
- fn unpack_bf16_pair(word: u32) -> vec2<f32> {
50
- let bits = vec2<u32>(word & 0xffffu, word >> 16u);
51
- return bitcast<vec2<f32>>(bits << vec2<u32>(16u));
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
  }
53
  {% endif %}
54
 
55
- {% if not vec4 %}
56
- fn embedding_scalar(source_row: u32, hidden: u32) -> f32 {
57
- if (source_row >= NUM_ROWS) {
58
- return 0.0;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
59
  }
60
- let word = x[source_row * HIDDEN_PAIRS + (hidden >> 1u)];
61
- let bits = select(word & 0xffffu, word >> 16u, (hidden & 1u) != 0u);
62
- return bitcast<f32>(bits << 16u);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
63
  }
64
  {% endif %}
 
 
65
 
 
 
 
 
 
 
 
 
 
 
 
 
 
66
  {% if vec4 %}
67
- fn embedding_vec4(source_row: u32, hidden_vec: u32) -> vec4<f32> {
68
- if (source_row >= NUM_ROWS) {
69
- return vec4<f32>(0.0);
 
 
 
 
 
 
70
  }
71
- let base = source_row * HIDDEN_PAIRS + hidden_vec * 2u;
72
- let low = unpack_bf16_pair(x[base]);
73
- let high = unpack_bf16_pair(x[base + 1u]);
74
- return vec4<f32>(low, high);
75
- }
76
  {% endif %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
  {% endif %}
78
-
79
- {% if vec4 and scalarIo %}
80
- fn load_vec4(index: u32) -> vec4<f32> {
81
- return vec4<f32>(x[index], x[index + 1u], x[index + 2u], x[index + 3u]);
82
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
83
  {% endif %}
 
 
 
 
 
 
84
 
85
  {% if combineSubgroups %}
86
  var<workgroup> sg_partials: array<vec2<f32>, WG>;
@@ -110,17 +391,22 @@ fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
110
  tr0[tid] = value.x;
111
  tr1[tid] = value.y;
112
  workgroupBarrier();
113
- var stride: u32 = WG / 2u;
114
  loop {
115
  if (stride == 0u) { break; }
116
- if (tid < stride) {
117
  tr0[tid] = tr0[tid] + tr0[tid + stride];
118
  tr1[tid] = tr1[tid] + tr1[tid + stride];
119
  }
120
  stride = stride / 2u;
121
  workgroupBarrier();
122
  }
 
 
 
 
123
  let reduced = vec2<f32>(tr0[0], tr1[0]);
 
124
  workgroupBarrier();
125
  return reduced;
126
  }
@@ -134,27 +420,33 @@ fn main(
134
  @builtin(subgroup_id) sg_id: u32,
135
  @builtin(num_subgroups) num_sg: u32{% endif %}
136
  ) {
 
 
 
 
137
  let row = wg_id.x + wg_id.y * params.rowStride;
138
  if (row >= params.rows) {
139
  return;
140
  }
 
141
  let tid = lid.x;
142
- {% if packedBf16Embedding %}
143
- let source_row = indices[row];
144
- {% if vec4 %}
145
  let base = row * HIDDEN_V;
146
  {% else %}
147
  let base = row * HIDDEN;
148
  {% endif %}
149
- {% elif vec4 and not scalarIo %}
150
- let base = row * HIDDEN_V;
151
- {% else %}
152
- let base = row * HIDDEN;
 
153
  {% endif %}
154
 
 
155
  {% if vec4 %}
156
- {% if scalarIo %}
157
- let shift = f32(x[base]);
 
158
  {% else %}
159
  let shift = f32(x[base].x);
160
  {% endif %}
@@ -164,27 +456,15 @@ fn main(
164
 
165
  var acc = vec2<f32>(0.0, 0.0);
166
  {% if vec4 %}
167
- for (var i = tid; i < HIDDEN_V; i = i + WG) {
168
- {% if packedBf16Embedding %}
169
- let v = embedding_vec4(source_row, i);
170
- embedding_out[base + i] = v;
171
- {% elif scalarIo %}
172
- let v = load_vec4(base + i * 4u);
173
- {% else %}
174
- let v = vec4<f32>(x[base + i]);
175
- {% endif %}
176
- let d = v - vec4<f32>(shift);
177
- acc.x = acc.x + d.x + d.y + d.z + d.w;
178
  acc.y = acc.y + dot(d, d);
179
  }
180
  {% else %}
181
  for (var i = tid; i < HIDDEN; i = i + WG) {
182
- {% if packedBf16Embedding %}
183
- let v = embedding_scalar(source_row, i);
184
- embedding_out[base + i] = v;
185
- {% else %}
186
  let v = f32(x[base + i]);
187
- {% endif %}
188
  let d = v - shift;
189
  acc.x = acc.x + d;
190
  acc.y = acc.y + d * d;
@@ -204,50 +484,20 @@ fn main(
204
  }
205
  {% endif %}
206
 
207
- {% if rmsChainNorm %}
208
- var acc2 = 0.0;
209
- {% endif %}
210
  {% if vec4 %}
211
- for (var i = tid; i < HIDDEN_V; i = i + WG) {
212
- {% if packedBf16Embedding %}
213
- let idx = base + i;
214
- let v = embedding_vec4(source_row, i);
215
- {% elif scalarIo %}
216
- let idx = base + i * 4u;
217
- let v = load_vec4(idx);
218
- {% else %}
219
  let idx = base + i;
220
- let v = vec4<f32>(x[idx]);
221
- {% endif %}
222
- var value = (v - vec4<f32>(row_mean)) * inv * vec4<f32>(scale[i]);
223
  {% if hasBias %}
224
- value = value + vec4<f32>(bias[i]);
225
  {% endif %}
226
  y[idx] = {{ vecType }}(value);
227
  }
228
- {% if rmsChainNorm %}
229
-
230
- // The chained second norm reads the residual row this loop just stored. This
231
- // barrier completes those stores and any preceding shared-scratch use before
232
- // the next reduction reuses its scratch; each lane then re-reads only the
233
- // elements it wrote itself.
234
- workgroupBarrier();
235
- let total2 = reduce_scalar(acc2{{ reduceThreadArguments }});
236
- let inv2 = inverseSqrt(total2 / f32(HIDDEN) + EPSILON2);
237
- for (var i = tid; i < HIDDEN_V; i = i + WG) {
238
- let idx = base + i;
239
- let hv = vec4<f32>(y[idx]);
240
- normed2[idx] = {{ vecType }}(hv * inv2 * vec4<f32>(scale2[i]));
241
- }
242
- {% endif %}
243
  {% else %}
244
  for (var i = tid; i < HIDDEN; i = i + WG) {
245
  let idx = base + i;
246
- {% if packedBf16Embedding %}
247
- let v = embedding_scalar(source_row, i);
248
- {% else %}
249
  let v = f32(x[idx]);
250
- {% endif %}
251
  var value = (v - row_mean) * inv * f32(scale[i]);
252
  {% if hasBias %}
253
  value = value + f32(bias[i]);
@@ -256,3 +506,4 @@ fn main(
256
  }
257
  {% endif %}
258
  }
 
 
1
+ {% set scalarIo = false %}
2
+ {% set packedWidth = packedWidth | default(4) %}
3
+ {% set packedF32 = "vec" ~ packedWidth ~ "<f32>" %}
 
 
 
 
 
 
 
 
 
 
4
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
5
  if combineSubgroups else ", tid: u32" %}
6
  {% set reduceThreadArguments = ", sg_lane, sg_id, num_sg"
 
9
  enable subgroups;
10
  {% endif %}
11
  {{ env.wgsl.resourceDeclarations }}
12
+ {% set affineRowBroadcast = affineRowBroadcast | default(false) %}
13
+ {% set batchRows = batchRows | default(1) %}
14
+ {% set batchLanes = batchLanes | default(0) %}
15
+ {% set xRowShape = xRowShape | default([]) %}
16
+ {% set scaleRowShape = scaleRowShape | default([]) %}
17
+ {% set biasRowShape = biasRowShape | default([]) %}
18
+ {% if affineRowBroadcast %}
19
+ {% macro offset_fn(fn_name, opShape, opRank, op_same, op_numel, outShape, outRank, out_numel) %}
20
+ fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif %}) -> u32 {
21
+ {% if out_numel == 0 %}
22
+ return 0u;
23
+ {% elif op_numel == 1 %}
24
+ return 0u;
25
+ {% elif op_same %}
26
+ return out_index;
27
+ {% else %}
28
+ var offset = 0u;
29
+ {% for axis in range(outRank) %}
30
+ {% set op_axis = axis - (outRank - opRank) %}
31
+ {% if op_axis >= 0 and opShape[op_axis] != 1 %}
32
+ {% set c_stride = namespace(value=1) %}
33
+ {% for j in range(axis + 1, outRank) %}
34
+ {% set c_stride.value = c_stride.value * outShape[j] %}
35
+ {% endfor %}
36
+ {% set op_stride = namespace(value=1) %}
37
+ {% for j in range(op_axis + 1, opRank) %}
38
+ {% set op_stride.value = op_stride.value * opShape[j] %}
39
+ {% endfor %}
40
+ {% if c_stride.value == 1 %}
41
+ let coord{{ axis }} = out_index % {{ outShape[axis] }}u;
42
+ {% else %}
43
+ let coord{{ axis }} = (out_index / {{ c_stride.value }}u) % {{ outShape[axis] }}u;
44
  {% endif %}
45
+ {% if op_stride.value == 1 %}
46
+ offset = offset + coord{{ axis }};
47
+ {% else %}
48
+ offset = offset + coord{{ axis }} * {{ op_stride.value }}u;
49
  {% endif %}
 
 
 
 
50
  {% endif %}
51
+ {% endfor %}
52
+ return offset;
53
+ {% endif %}
54
+ }{% endmacro %}
55
+ {% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
56
+ {% set op_numel = namespace(value=1) %}
57
+ {% for d in opShape %}
58
+ {% set op_numel.value = op_numel.value * d %}
59
+ {% endfor %}
60
+ {% set out_numel = namespace(value=1) %}
61
+ {% for d in outShape %}
62
+ {% set out_numel.value = out_numel.value * d %}
63
+ {% endfor %}
64
+ {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %}){% endmacro %}
65
+ {% macro broadcast_offset_fn(fn_name, opShape, opRank, outShape, outRank) %}
66
+ {% set op_numel = namespace(value=1) %}
67
+ {% for d in opShape %}
68
+ {% set op_numel.value = op_numel.value * d %}
69
+ {% endfor %}
70
+ {% set out_numel = namespace(value=1) %}
71
+ {% for d in outShape %}
72
+ {% set out_numel.value = out_numel.value * d %}
73
+ {% endfor %}
74
+ {% set op_same = namespace(value=(opRank == outRank)) %}
75
+ {% if op_same.value %}
76
+ {% for axis in range(outRank) %}
77
+ {% if opShape[axis] != outShape[axis] %}
78
+ {% set op_same.value = false %}
79
+ {% endif %}
80
+ {% endfor %}
81
+ {% endif %}
82
+ {{ offset_fn(fn_name, opShape, opRank, op_same.value, op_numel.value, outShape, outRank, out_numel.value) }}{% endmacro %}
83
+ {{ broadcast_offset_fn("scale_row_offset", scaleRowShape, scaleRowShape | length, xRowShape, xRowShape | length) }}
84
+ {% if hasBias %}
85
+ {{ broadcast_offset_fn("bias_row_offset", biasRowShape, biasRowShape | length, xRowShape, xRowShape | length) }}
86
+ {% endif %}
87
+ {% endif %}
88
+ {% if scalar == "f16" and compensateHalfStats is defined and compensateHalfStats %}
89
+ {% set halfOutputScalar = "f16" %}
90
+ {% set halfStageVector = vec4 %}
91
+ fn round_f16_bits_rte(value: f32) -> u32 {
92
+ let bits = bitcast<u32>(value);
93
+ let sign = (bits >> 16u) & 0x8000u;
94
+ let exponent_f32 = (bits >> 23u) & 0xffu;
95
+ let mantissa_f32 = bits & 0x7fffffu;
96
 
97
+ if (exponent_f32 == 0xffu) {
98
+ if (mantissa_f32 != 0u) {
99
+ return 0x7e00u;
100
+ }
101
+ return sign | 0x7c00u;
102
+ }
103
+
104
+ var exponent_f16 = i32(exponent_f32) - 127 + 15;
105
+ if (exponent_f16 >= 0x1f) {
106
+ return sign | 0x7c00u;
107
+ }
108
+
109
+ if (exponent_f16 <= 0) {
110
+ if (exponent_f16 < -10) {
111
+ return sign;
112
+ }
113
+ let significand = mantissa_f32 | 0x800000u;
114
+ let shift = u32(14 - exponent_f16);
115
+ let halfway = 1u << (shift - 1u);
116
+ let discarded = significand & ((1u << shift) - 1u);
117
+ var fraction = significand >> shift;
118
+ if (discarded > halfway || (discarded == halfway && (fraction & 1u) == 1u)) {
119
+ fraction = fraction + 1u;
120
+ }
121
+ return sign | fraction;
122
+ }
123
+
124
+ let halfway = 1u << 12u;
125
+ let discarded = mantissa_f32 & 0x1fffu;
126
+ var mantissa_f16 = mantissa_f32 >> 13u;
127
+ if (discarded > halfway || (discarded == halfway && (mantissa_f16 & 1u) == 1u)) {
128
+ mantissa_f16 = mantissa_f16 + 1u;
129
+ if (mantissa_f16 == 0x400u) {
130
+ mantissa_f16 = 0u;
131
+ exponent_f16 = exponent_f16 + 1;
132
+ }
133
+ }
134
+ if (exponent_f16 >= 0x1f) {
135
+ return sign | 0x7c00u;
136
+ }
137
+ return sign | (u32(exponent_f16) << 10u) | mantissa_f16;
138
+ }
139
+
140
+ fn widen_f16_bits(value: u32) -> f32 {
141
+ return unpack2x16float(value & 0xffffu).x;
142
+ }
143
+
144
+ // Typed ONNX edges must survive arithmetic fusion and narrow/wide casts.
145
+ // Integer rounding also fixes the ties-to-even rule independently of the
146
+ // implementation's floating-point conversion rounding mode.
147
+ fn half_stage(value: f32) -> f32 {
148
+ return widen_f16_bits(round_f16_bits_rte(value));
149
+ }
150
+ {% if halfStageVector | default(false) %}
151
+
152
+ fn half_stage4(value: vec4<f32>) -> vec4<f32> {
153
+ return vec4<f32>(half_stage(value.x), half_stage(value.y),
154
+ half_stage(value.z), half_stage(value.w));
155
  }
156
  {% endif %}
157
 
158
+ // Half output magnifies statistics errors at rounding midpoints. Keep a low
159
+ // residual through the reduction and normalization, then round the typed
160
+ // float32-normalized/half-scale/half-bias edges explicitly. This is the standard ONNX half path;
161
+ // other normalization contracts retain their existing arithmetic.
162
+ const HIDDEN: u32 = {{ hidden }}u;
163
+ const WG: u32 = {{ wg }}u;
164
+ const EPSILON: f32 = {{ epsilon }};
165
+ {% set halfWriteMean = halfWriteMean if halfWriteMean is defined else (writeStats and modeSpec == "layer") %}
166
+ {% set halfWriteInv = halfWriteInv if halfWriteInv is defined else writeStats %}
167
+
168
+ fn pair_add(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> {
169
+ let s = fma(a.x, 1.0, b.x);
170
+ // Materialize each rounded subtraction in the error-free transform. Plain
171
+ // cancellation expressions do not preserve the intended evaluation tree on
172
+ // every shader backend.
173
+ let bv = fma(-1.0, a.x, s);
174
+ let av = fma(-1.0, bv, s);
175
+ let a_error = fma(-1.0, av, a.x);
176
+ let b_error = fma(-1.0, bv, b.x);
177
+ let error = fma(a_error, 1.0, b_error);
178
+ let e = fma(fma(error, 1.0, a.y), 1.0, b.y);
179
+ let hi = fma(s, 1.0, e);
180
+ return vec2<f32>(hi, fma(-1.0, fma(-1.0, s, hi), e));
181
+ }
182
+
183
+ fn pair_mul(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> {
184
+ let p = fma(a.x, b.x, 0.0);
185
+ let error = fma(a.x, b.x, -p);
186
+ return pair_add(vec2<f32>(p, 0.0), vec2<f32>(error + (a.x * b.y + b.x * a.y), 0.0));
187
+ }
188
+
189
+ fn pair_div(a: vec2<f32>, b: f32) -> vec2<f32> {
190
+ let q = a.x / b;
191
+ let residual = pair_add(a, -pair_mul(vec2<f32>(q, 0.0), vec2<f32>(b, 0.0)));
192
+ return pair_add(vec2<f32>(q, 0.0), vec2<f32>((residual.x + residual.y) / b, 0.0));
193
+ }
194
+
195
+ fn pair_inverse_sqrt(a: vec2<f32>) -> vec2<f32> {
196
+ let r = inverseSqrt(a.x);
197
+ let rr = pair_mul(vec2<f32>(r, 0.0), vec2<f32>(r, 0.0));
198
+ let residual = pair_add(vec2<f32>(1.0, 0.0), -pair_mul(a, rr));
199
+ return pair_add(vec2<f32>(r, 0.0), vec2<f32>((0.5 * r) * (residual.x + residual.y), 0.0));
200
+ }
201
+
202
+ fn half_normalized(value: vec2<f32>) -> f32 {
203
+ // stash_type=1 materializes Normalized as float32 before its cast to half.
204
+ // Collapse the compensated residual at that typed edge; rounding the pair
205
+ // directly to half can choose a different result at a float32 midpoint.
206
+ return widen_f16_bits(round_f16_bits_rte(fma(value.x, 1.0, value.y)));
207
+ }
208
+
209
+ {% if combineSubgroups %}
210
+ var<workgroup> sg_partials: array<vec2<f32>, WG>;
211
+ fn reduce_pair(value: vec2<f32>, sg_lane: u32, sg_id: u32, num_sg: u32, sg_size: u32) -> vec2<f32> {
212
+ var s = value;
213
+ let active_lanes = subgroupBallot(true);
214
+ let counts = countOneBits(active_lanes);
215
+ if (counts.x + counts.y + counts.z + counts.w == sg_size) {
216
+ for (var step = sg_size / 2u; step > 0u; step /= 2u) {
217
+ let other = subgroupShuffleDown(s, step);
218
+ if (sg_lane + step < sg_size) { s = pair_add(s, other); }
219
+ }
220
+ s = subgroupBroadcastFirst(s);
221
+ } else {
222
+ // Small workgroups need not fill a subgroup. Enumerate its actual active
223
+ // lanes instead of adding indeterminate shuffle results from inactive lanes.
224
+ s = vec2<f32>(0.0);
225
+ for (var word = 0u; word < 4u; word++) {
226
+ var mask = active_lanes[word];
227
+ while (mask != 0u) {
228
+ let lane = word * 32u + firstTrailingBit(mask);
229
+ s = pair_add(s, subgroupShuffle(value, lane));
230
+ mask &= mask - 1u;
231
+ }
232
+ }
233
  }
234
+ if (num_sg == 1u) { return s; }
235
+ if (subgroupElect()) { sg_partials[sg_id] = s; }
236
+ workgroupBarrier();
237
+ var total = vec2<f32>(0.0);
238
+ for (var i = 0u; i < num_sg; i++) { total = pair_add(total, sg_partials[i]); }
239
+ // Another reduction may immediately reuse the same storage.
240
+ workgroupBarrier();
241
+ return total;
242
+ }
243
+ {% else %}
244
+ var<workgroup> partial: array<vec2<f32>, WG>;
245
+ fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
246
+ partial[tid] = value;
247
+ workgroupBarrier();
248
+ for (var step = WG / 2u; step > 0u; step /= 2u) {
249
+ if (tid < step) { partial[tid] = pair_add(partial[tid], partial[tid + step]); }
250
+ workgroupBarrier();
251
+ }
252
+ let total = partial[0];
253
+ workgroupBarrier();
254
+ return total;
255
  }
256
  {% endif %}
257
+ {% set reduceArgs = "sg_lane, sg_id, num_sg, sg_size" if combineSubgroups else "tid" %}
258
+ {% if not vec4 %}
259
 
260
+ fn load_value(index: u32) -> f32 {
261
+ return f32(x[index]);
262
+ }
263
+ {% endif %}
264
+
265
+ fn normalize_half_row(row: u32, tid: u32
266
+ {% if combineSubgroups %}
267
+ , sg_lane: u32, sg_id: u32, num_sg: u32, sg_size: u32
268
+ {% endif %}
269
+ ) {
270
+ if (row >= params.rows) { return; }
271
+ let base = row * HIDDEN;
272
+ var local_sum = vec2<f32>(0.0);
273
  {% if vec4 %}
274
+ for (var i = tid; i < HIDDEN / 4u; i += WG) {
275
+ let v = vec4<f32>(x[base / 4u + i]);
276
+ {% for component in ["x", "y", "z", "w"] %}
277
+ local_sum = pair_add(local_sum, vec2<f32>(v.{{ component }}, 0.0));
278
+ {% endfor %}
279
+ }
280
+ {% else %}
281
+ for (var i = tid; i < HIDDEN; i += WG) {
282
+ local_sum = pair_add(local_sum, vec2<f32>(load_value(base + i), 0.0));
283
  }
 
 
 
 
 
284
  {% endif %}
285
+ let mean = pair_div(reduce_pair(local_sum, {{ reduceArgs }}), f32(HIDDEN));
286
+ var local_square = vec2<f32>(0.0);
287
+ {% if vec4 %}
288
+ for (var i = tid; i < HIDDEN / 4u; i += WG) {
289
+ let v = vec4<f32>(x[base / 4u + i]);
290
+ {% for component in ["x", "y", "z", "w"] %}
291
+ let centered_{{ component }} = pair_add(vec2<f32>(v.{{ component }}, 0.0), -mean);
292
+ local_square = pair_add(local_square, pair_mul(centered_{{ component }}, centered_{{ component }}));
293
+ {% endfor %}
294
+ }
295
+ {% else %}
296
+ for (var i = tid; i < HIDDEN; i += WG) {
297
+ let centered = pair_add(vec2<f32>(load_value(base + i), 0.0), -mean);
298
+ local_square = pair_add(local_square, pair_mul(centered, centered));
299
+ }
300
+ {% endif %}
301
+ let variance = pair_div(reduce_pair(local_square, {{ reduceArgs }}), f32(HIDDEN));
302
+ let inv = pair_inverse_sqrt(pair_add(variance, vec2<f32>(EPSILON, 0.0)));
303
+ {% if halfWriteMean %}
304
+ if (tid == 0u) { mean_out[row] = mean.x; }
305
+ {% endif %}
306
+ {% if halfWriteInv %}
307
+ if (tid == 0u) { inv_std_out[row] = inv.x; }
308
+ {% endif %}
309
+ {% if vec4 %}
310
+ for (var i = tid; i < HIDDEN / 4u; i += WG) {
311
+ let v = vec4<f32>(x[base / 4u + i]);
312
+ var normalized: vec4<f32>;
313
+ for (var component = 0u; component < 4u; component++) {
314
+ normalized[component] = half_normalized(pair_mul(pair_add(vec2<f32>(v[component], 0.0), -mean), inv));
315
+ }
316
+ var value = half_stage4(normalized * vec4<f32>(scale[i]));
317
+ {% if modeSpec == "layer" and hasBias %}
318
+ value = half_stage4(value + vec4<f32>(bias[i]));
319
+ {% endif %}
320
+ y[base / 4u + i] = vec4<f16>(value);
321
+ }
322
+ {% else %}
323
+ for (var i = tid; i < HIDDEN; i += WG) {
324
+ let normalized = half_normalized(pair_mul(pair_add(vec2<f32>(load_value(base + i), 0.0), -mean), inv));
325
+ var value = half_stage(normalized * f32(scale[{{ halfScaleOffset | default("i") }}]));
326
+ {% if modeSpec == "layer" and hasBias %}
327
+ value = half_stage(value + f32(bias[{{ halfBiasOffset | default("i") }}]));
328
+ {% endif %}
329
+ y[base + i] = {{ halfOutputScalar }}(value);
330
+ }
331
  {% endif %}
 
 
 
 
332
  }
333
+ @compute @workgroup_size(WG, 1, 1)
334
+ fn main(@builtin(workgroup_id) wg_id: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>
335
+ {% if combineSubgroups %}
336
+ , @builtin(subgroup_invocation_id) sg_lane: u32, @builtin(subgroup_id) sg_id: u32,
337
+ @builtin(num_subgroups) num_sg: u32, @builtin(subgroup_size) sg_size: u32
338
+ {% endif %}
339
+ ) {
340
+ let row = wg_id.x + wg_id.y * params.rowStride;
341
+ normalize_half_row(row, lid.x{% if combineSubgroups %}, sg_lane, sg_id, num_sg, sg_size{% endif %});
342
+ }
343
+ {% else %}
344
+
345
+ // Workgroup-parallel single-pass row statistics + fused normalize/affine.
346
+ //
347
+ // One workgroup owns one contiguous normalization span ("row": a last-axis
348
+ // row, an instance plane, or a channel group). Threads stride the row once,
349
+ // accumulating (sum, sum_sq) simultaneously. Partials are reduced either with
350
+ // subgroupAdd plus a shared-memory combine or with a portable shared-memory
351
+ // tree, then every thread applies the fused normalize + affine write.
352
+ //
353
+ // Shifted moments avoid cancellation from a large common offset; scaling uses
354
+ // inverseSqrt(variance + EPSILON).
355
+ const HIDDEN: u32 = {{ hidden }}u;
356
+ {% if vec4 %}
357
+ const HIDDEN_V: u32 = {{ hiddenVec }}u;
358
  {% endif %}
359
+ const WG: u32 = {{ wg }}u;
360
+ {% if batchRows > 1 %}
361
+ const ROW_LANES: u32 = {{ batchLanes }}u;
362
+ const ROW_BATCH: u32 = {{ batchRows }}u;
363
+ {% endif %}
364
+ const EPSILON: f32 = {{ epsilon }};
365
 
366
  {% if combineSubgroups %}
367
  var<workgroup> sg_partials: array<vec2<f32>, WG>;
 
391
  tr0[tid] = value.x;
392
  tr1[tid] = value.y;
393
  workgroupBarrier();
394
+ var stride: u32 = {% if batchRows > 1 %}ROW_LANES{% else %}WG{% endif %} / 2u;
395
  loop {
396
  if (stride == 0u) { break; }
397
+ if (tid{% if batchRows > 1 %} % ROW_LANES{% endif %} < stride) {
398
  tr0[tid] = tr0[tid] + tr0[tid + stride];
399
  tr1[tid] = tr1[tid] + tr1[tid + stride];
400
  }
401
  stride = stride / 2u;
402
  workgroupBarrier();
403
  }
404
+ {% if batchRows > 1 %}
405
+ let row_base = tid - tid % ROW_LANES;
406
+ let reduced = vec2<f32>(tr0[row_base], tr1[row_base]);
407
+ {% else %}
408
  let reduced = vec2<f32>(tr0[0], tr1[0]);
409
+ {% endif %}
410
  workgroupBarrier();
411
  return reduced;
412
  }
 
420
  @builtin(subgroup_id) sg_id: u32,
421
  @builtin(num_subgroups) num_sg: u32{% endif %}
422
  ) {
423
+ {% if batchRows > 1 %}
424
+ let row = (wg_id.x + wg_id.y * params.rowStride) * ROW_BATCH + lid.x / ROW_LANES;
425
+ let row_active = row < params.rows;
426
+ {% else %}
427
  let row = wg_id.x + wg_id.y * params.rowStride;
428
  if (row >= params.rows) {
429
  return;
430
  }
431
+ {% endif %}
432
  let tid = lid.x;
433
+ {% if vec4 and not scalarIo %}
 
 
434
  let base = row * HIDDEN_V;
435
  {% else %}
436
  let base = row * HIDDEN;
437
  {% endif %}
438
+
439
+ {% if affineRowBroadcast %}
440
+ let scale_base = {{ broadcast_offset_call("scale_row_offset", scaleRowShape, xRowShape, "row") }} * HIDDEN_V;
441
+ {% if hasBias %}
442
+ let bias_base = {{ broadcast_offset_call("bias_row_offset", biasRowShape, xRowShape, "row") }} * HIDDEN_V;
443
  {% endif %}
444
 
445
+ {% endif %}
446
  {% if vec4 %}
447
+ {% if batchRows > 1 %}
448
+ var shift = 0.0;
449
+ if (row_active) { shift = f32(x[base].x); }
450
  {% else %}
451
  let shift = f32(x[base].x);
452
  {% endif %}
 
456
 
457
  var acc = vec2<f32>(0.0, 0.0);
458
  {% if vec4 %}
459
+ for (var i = tid{% if batchRows > 1 %} % ROW_LANES{% endif %}; {% if batchRows > 1 %}row_active && {% endif %}i < HIDDEN_V; i = i + {% if batchRows > 1 %}ROW_LANES{% else %}WG{% endif %}) {
460
+ let v = {{ packedF32 }}(x[base + i]);
461
+ let d = v - {{ packedF32 }}(shift);
462
+ acc.x = acc.x + d.x + d.y{% if packedWidth == 4 %} + d.z + d.w{% endif %};
 
 
 
 
 
 
 
463
  acc.y = acc.y + dot(d, d);
464
  }
465
  {% else %}
466
  for (var i = tid; i < HIDDEN; i = i + WG) {
 
 
 
 
467
  let v = f32(x[base + i]);
 
468
  let d = v - shift;
469
  acc.x = acc.x + d;
470
  acc.y = acc.y + d * d;
 
484
  }
485
  {% endif %}
486
 
 
 
 
487
  {% if vec4 %}
488
+ for (var i = tid{% if batchRows > 1 %} % ROW_LANES{% endif %}; {% if batchRows > 1 %}row_active && {% endif %}i < HIDDEN_V; i = i + {% if batchRows > 1 %}ROW_LANES{% else %}WG{% endif %}) {
 
 
 
 
 
 
 
489
  let idx = base + i;
490
+ let v = {{ packedF32 }}(x[idx]);
491
+ var value = (v - {{ packedF32 }}(row_mean)) * inv * {{ packedF32 }}(scale[{% if affineRowBroadcast %}scale_base + {% endif %}i]);
 
492
  {% if hasBias %}
493
+ value = value + {{ packedF32 }}(bias[{% if affineRowBroadcast %}bias_base + {% endif %}i]);
494
  {% endif %}
495
  y[idx] = {{ vecType }}(value);
496
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
497
  {% else %}
498
  for (var i = tid; i < HIDDEN; i = i + WG) {
499
  let idx = base + i;
 
 
 
500
  let v = f32(x[idx]);
 
501
  var value = (v - row_mean) * inv * f32(scale[i]);
502
  {% if hasBias %}
503
  value = value + f32(bias[i]);
 
506
  }
507
  {% endif %}
508
  }
509
+ {% endif %}
build/webgpu/test.json CHANGED
@@ -4,6 +4,157 @@
4
  "onnx_backend_layer_normalization_input_x": [1.764052391052246, 0.40015721321105957, 0.978738009929657, 2.2408931255340576, 1.8675580024719238, -0.9772778749465942, 0.9500884413719177, -0.15135720372200012, -0.10321885347366333, 0.4105985164642334, 0.14404356479644775, 1.4542734622955322, 0.7610377073287964, 0.12167501449584961, 0.44386324286460876, 0.3336743414402008, 1.4940791130065918, -0.2051582634449005, 0.3130677044391632, -0.8540957570075989, -2.5529897212982178, 0.653618574142456, 0.8644362092018127, -0.7421650290489197, 2.269754648208618, -1.4543657302856445, 0.04575851559638977, -0.18718385696411133, 1.5327792167663574, 1.4693588018417358, 0.154947429895401, 0.37816253304481506, -0.8877857327461243, -1.980796456336975, -0.34791216254234314, 0.15634897351264954, 1.2302906513214111, 1.202379822731018, -0.38732680678367615, -0.302302747964859, -1.0485529899597168, -1.420017957687378, -1.7062702178955078, 1.950775384902954, -0.5096521973609924, -0.4380742907524109, -1.2527953386306763, 0.7774903774261475, -1.6138978004455566, -0.21274028718471527, -0.8954665660858154, 0.38690251111984253, -0.5108051300048828, -1.18063223361969, -0.02818222902715206, 0.4283318817615509, 0.06651721894741058, 0.30247190594673157, -0.6343221068382263, -0.3627411723136902, -0.6724604368209839, -0.35955315828323364, -0.8131462931632996, -1.7262825965881348, 0.17742614448070526, -0.4017809331417084, -1.630198359489441, 0.46278226375579834, -0.9072983860969543, 0.05194539576768875, 0.7290905714035034, 0.12898291647434235, 1.1394007205963135, -1.234825849533081, 0.4023416340351105, -0.6848101019859314, -0.8707971572875977, -0.5788496732711792, -0.3115525245666504, 0.056165341287851334, -1.1651498079299927, 0.9008265137672424, 0.4656624495983124, -1.5362436771392822, 1.4882521629333496, 1.895889163017273, 1.1787796020507812, -0.1799248307943344, -1.0707526206970215, 1.0544517040252686, -0.4031769335269928, 1.222445011138916, 0.2082749754190445, 0.9766390323638916, 0.3563663959503174, 0.7065731883049011, 0.01050002034753561, 1.7858705520629883, 0.12691208720207214, 0.4019893705844879, 1.8831506967544556, -1.3477590084075928, -1.2704850435256958, 0.969396710395813, -1.1731233596801758, 1.9436211585998535, -0.4136189818382263, -0.747454822063446, 1.922942042350769, 1.4805147647857666, 1.8675589561462402, 0.9060446619987488, -0.8612256646156311, 1.910064935684204, -0.26800337433815, 0.8024563789367676, 0.9472519755363464, -0.15501008927822113, 0.6140793561935425, 0.922206699848175]
5
  },
6
  "cases": [
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7
  {
8
  "name": "subgroup_vec4_stats_no_bias_2x512",
9
  "attrs": { "epsilon": 0.00001, "axis": -1 },
@@ -1783,6 +1934,373 @@
1783
  "mean": { "dtype": "float32", "shape": [2, 1], "tolerance": 0.000001 },
1784
  "invStdDev": { "dtype": "float32", "shape": [2, 1], "tolerance": 0.00001 }
1785
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1786
  }
1787
  ]
1788
  }
 
4
  "onnx_backend_layer_normalization_input_x": [1.764052391052246, 0.40015721321105957, 0.978738009929657, 2.2408931255340576, 1.8675580024719238, -0.9772778749465942, 0.9500884413719177, -0.15135720372200012, -0.10321885347366333, 0.4105985164642334, 0.14404356479644775, 1.4542734622955322, 0.7610377073287964, 0.12167501449584961, 0.44386324286460876, 0.3336743414402008, 1.4940791130065918, -0.2051582634449005, 0.3130677044391632, -0.8540957570075989, -2.5529897212982178, 0.653618574142456, 0.8644362092018127, -0.7421650290489197, 2.269754648208618, -1.4543657302856445, 0.04575851559638977, -0.18718385696411133, 1.5327792167663574, 1.4693588018417358, 0.154947429895401, 0.37816253304481506, -0.8877857327461243, -1.980796456336975, -0.34791216254234314, 0.15634897351264954, 1.2302906513214111, 1.202379822731018, -0.38732680678367615, -0.302302747964859, -1.0485529899597168, -1.420017957687378, -1.7062702178955078, 1.950775384902954, -0.5096521973609924, -0.4380742907524109, -1.2527953386306763, 0.7774903774261475, -1.6138978004455566, -0.21274028718471527, -0.8954665660858154, 0.38690251111984253, -0.5108051300048828, -1.18063223361969, -0.02818222902715206, 0.4283318817615509, 0.06651721894741058, 0.30247190594673157, -0.6343221068382263, -0.3627411723136902, -0.6724604368209839, -0.35955315828323364, -0.8131462931632996, -1.7262825965881348, 0.17742614448070526, -0.4017809331417084, -1.630198359489441, 0.46278226375579834, -0.9072983860969543, 0.05194539576768875, 0.7290905714035034, 0.12898291647434235, 1.1394007205963135, -1.234825849533081, 0.4023416340351105, -0.6848101019859314, -0.8707971572875977, -0.5788496732711792, -0.3115525245666504, 0.056165341287851334, -1.1651498079299927, 0.9008265137672424, 0.4656624495983124, -1.5362436771392822, 1.4882521629333496, 1.895889163017273, 1.1787796020507812, -0.1799248307943344, -1.0707526206970215, 1.0544517040252686, -0.4031769335269928, 1.222445011138916, 0.2082749754190445, 0.9766390323638916, 0.3563663959503174, 0.7065731883049011, 0.01050002034753561, 1.7858705520629883, 0.12691208720207214, 0.4019893705844879, 1.8831506967544556, -1.3477590084075928, -1.2704850435256958, 0.969396710395813, -1.1731233596801758, 1.9436211585998535, -0.4136189818382263, -0.747454822063446, 1.922942042350769, 1.4805147647857666, 1.8675589561462402, 0.9060446619987488, -0.8612256646156311, 1.910064935684204, -0.26800337433815, 0.8024563789367676, 0.9472519755363464, -0.15501008927822113, 0.6140793561935425, 0.922206699848175]
5
  },
6
  "cases": [
7
+ {
8
+ "name": "half_normalized_midpoint_vec4",
9
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
10
+ "inputs": {
11
+ "x": {
12
+ "dtype": "float16",
13
+ "shape": [1, 12],
14
+ "data": {
15
+ "kind": "values",
16
+ "values": [0.103759765625, -0.128662109375, 0.271484375, 0.06829833984375, -0.00689697265625, 0.10546875, 0.568359375, -0.1162109375, 1.0185546875, 0.52392578125, -0.31591796875, 0.36279296875]
17
+ }
18
+ },
19
+ "scale": {
20
+ "dtype": "float16",
21
+ "shape": [12],
22
+ "data": {
23
+ "kind": "values",
24
+ "values": [0.9931640625, 0.82421875, 1.1064453125, 0.751953125, 1.2080078125, 1.0224609375, 0.9033203125, 1.228515625, 1.1533203125, 0.77001953125, 1.04296875, 1.1240234375]
25
+ }
26
+ },
27
+ "b": {
28
+ "dtype": "float16",
29
+ "shape": [12],
30
+ "data": {
31
+ "kind": "values",
32
+ "values": [0.115234375, -0.0238189697265625, -0.06988525390625, -0.04443359375, -0.1641845703125, 0.0662841796875, 0.00609588623046875, 0.0999755859375, 0.2476806640625, -0.0048675537109375, 0.2154541015625, -0.18359375]
33
+ }
34
+ }
35
+ },
36
+ "outputs": {
37
+ "y": {
38
+ "dtype": "float16",
39
+ "shape": [1, 12],
40
+ "data": {
41
+ "kind": "values",
42
+ "values": [-0.16845703125, -0.80224609375, 0.139892578125, -0.3349609375, -0.8876953125, -0.2208251953125, 0.9365234375, -1.017578125, 2.908203125, 0.69189453125, -1.322265625, 0.3203125]
43
+ },
44
+ "tolerance": 0.000002
45
+ }
46
+ },
47
+ "provenance": {
48
+ "source": "ONNX LayerNormalization typed normalization and scale stages",
49
+ "notes": "Half-normalized midpoint followed by affine rounding. Pinned with independent float64 statistics and explicit half stages; a float32-only reduction can cross the midpoint."
50
+ }
51
+ },
52
+ {
53
+ "name": "onnx17_half_affine_stages_scalar_broadcast",
54
+ "attrs": { "axis": -1, "epsilon": 0, "stash_type": 1 },
55
+ "inputs": {
56
+ "x": { "dtype": "float16", "shape": [1, 3], "data": { "kind": "values", "values": [0.3, 1.7, -0.9] } },
57
+ "scale": { "dtype": "float16", "shape": [], "data": { "kind": "values", "values": [1.3] } },
58
+ "b": { "dtype": "float16", "shape": [], "data": { "kind": "values", "values": [0.07] } }
59
+ },
60
+ "outputs": {
61
+ "y": {
62
+ "dtype": "float16",
63
+ "shape": [1, 3],
64
+ "tolerance": 0.000002,
65
+ "data": { "kind": "values", "values": [-0.0115966796875, 1.701171875, -1.4794921875] }
66
+ }
67
+ },
68
+ "provenance": {
69
+ "source": "ONNX LayerNormalization-17 typed affine stages",
70
+ "notes": "Scalar scale and bias exercise both rank-zero broadcast offsets in the generic half shader. Exact half values independently match explicit primitives on ONNX Runtime's CPU provider and typed PyTorch CPU stages."
71
+ }
72
+ },
73
+ {
74
+ "name": "onnx17_half_affine_stages_scalar",
75
+ "attrs": { "axis": -1, "epsilon": 0, "stash_type": 1 },
76
+ "inputs": {
77
+ "x": { "dtype": "float16", "shape": [1, 3], "data": { "kind": "values", "values": [0.3, 1.7, -0.9] } },
78
+ "scale": { "dtype": "float16", "shape": [3], "data": { "kind": "values", "values": [1.3, 2.1, -0.7] } },
79
+ "b": { "dtype": "float16", "shape": [3], "data": { "kind": "values", "values": [0.07, -0.3, 0.9] } }
80
+ },
81
+ "outputs": {
82
+ "y": {
83
+ "dtype": "float16",
84
+ "shape": [1, 3],
85
+ "tolerance": 0.000002,
86
+ "data": { "kind": "values", "values": [-0.0115966796875, 2.333984375, 1.734375] }
87
+ }
88
+ },
89
+ "provenance": {
90
+ "source": "ONNX LayerNormalization-17 function: Cast Normalized to T, then Mul and Add in T",
91
+ "notes": "Exact half stages independently checked with ONNX Runtime's CPU primitive graph and PyTorch's CPU explicit stages. Native fused LayerNorm differs at the first output; the spec is authoritative."
92
+ }
93
+ },
94
+ {
95
+ "name": "onnx17_half_affine_stages_vec4",
96
+ "attrs": { "axis": -1, "epsilon": 0, "stash_type": 1 },
97
+ "inputs": {
98
+ "x": {
99
+ "dtype": "float16",
100
+ "shape": [1, 12],
101
+ "data": { "kind": "values", "values": [0.3, 1.7, -0.9, 0.3, 1.7, -0.9, 0.3, 1.7, -0.9, 0.3, 1.7, -0.9] }
102
+ },
103
+ "scale": {
104
+ "dtype": "float16",
105
+ "shape": [12],
106
+ "data": { "kind": "values", "values": [1.3, 2.1, -0.7, 1.3, 2.1, -0.7, 1.3, 2.1, -0.7, 1.3, 2.1, -0.7] }
107
+ },
108
+ "b": {
109
+ "dtype": "float16",
110
+ "shape": [12],
111
+ "data": { "kind": "values", "values": [0.07, -0.3, 0.9, 0.07, -0.3, 0.9, 0.07, -0.3, 0.9, 0.07, -0.3, 0.9] }
112
+ }
113
+ },
114
+ "outputs": {
115
+ "y": {
116
+ "dtype": "float16",
117
+ "shape": [1, 12],
118
+ "tolerance": 0.000002,
119
+ "data": {
120
+ "kind": "values",
121
+ "values": [-0.0115966796875, 2.333984375, 1.734375, -0.0115966796875, 2.333984375, 1.734375, -0.0115966796875, 2.333984375, 1.734375, -0.0115966796875, 2.333984375, 1.734375]
122
+ }
123
+ }
124
+ },
125
+ "provenance": {
126
+ "source": "ONNX LayerNormalization-17 typed affine stages",
127
+ "notes": "Four repetitions of the 3-element pattern from this file's onnx17_half_affine_stages_scalar case (12 elements total) load in vec4 groups; float16 scale and bias check half-precision product rounding and addition, with exact expected output values."
128
+ }
129
+ },
130
+ {
131
+ "name": "onnx17_half_affine_stages_suffix_broadcast",
132
+ "attrs": { "axis": 1, "epsilon": 0, "stash_type": 1 },
133
+ "inputs": {
134
+ "x": {
135
+ "dtype": "float16",
136
+ "shape": [2, 2, 3],
137
+ "data": { "kind": "values", "values": [0.3, 1.7, -0.9, 0.3, 1.7, -0.9, 0.3, 1.7, -0.9, 0.3, 1.7, -0.9] }
138
+ },
139
+ "scale": { "dtype": "float16", "shape": [3], "data": { "kind": "values", "values": [1.3, 2.1, -0.7] } },
140
+ "b": { "dtype": "float16", "shape": [3], "data": { "kind": "values", "values": [0.07, -0.3, 0.9] } }
141
+ },
142
+ "outputs": {
143
+ "y": {
144
+ "dtype": "float16",
145
+ "shape": [2, 2, 3],
146
+ "tolerance": 0.000002,
147
+ "data": {
148
+ "kind": "values",
149
+ "values": [-0.0115966796875, 2.333984375, 1.734375, -0.0115966796875, 2.333984375, 1.734375, -0.0115966796875, 2.333984375, 1.734375, -0.0115966796875, 2.333984375, 1.734375]
150
+ }
151
+ }
152
+ },
153
+ "provenance": {
154
+ "source": "ONNX LayerNormalization-17 typed affine stages",
155
+ "notes": "Suffix-axis normalization with broadcast scale/bias exercises the generic shader on two independent rows."
156
+ }
157
+ },
158
  {
159
  "name": "subgroup_vec4_stats_no_bias_2x512",
160
  "attrs": { "epsilon": 0.00001, "axis": -1 },
 
1934
  "mean": { "dtype": "float32", "shape": [2, 1], "tolerance": 0.000001 },
1935
  "invStdDev": { "dtype": "float32", "shape": [2, 1], "tolerance": 0.00001 }
1936
  }
1937
+ },
1938
+ {
1939
+ "name": "broadcast_affine_vec4_full",
1940
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
1941
+ "inputs": {
1942
+ "x": {
1943
+ "dtype": "float32",
1944
+ "shape": [3, 64],
1945
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
1946
+ },
1947
+ "scale": {
1948
+ "dtype": "float32",
1949
+ "shape": [3, 64],
1950
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
1951
+ },
1952
+ "b": {
1953
+ "dtype": "float32",
1954
+ "shape": [3, 64],
1955
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
1956
+ }
1957
+ },
1958
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 64], "tolerance": 0.000002, "relTolerance": 0.00001 } },
1959
+ "provenance": {
1960
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
1961
+ }
1962
+ },
1963
+ {
1964
+ "name": "broadcast_affine_vec4_shifted",
1965
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
1966
+ "inputs": {
1967
+ "x": {
1968
+ "dtype": "float32",
1969
+ "shape": [3, 64],
1970
+ "data": {
1971
+ "kind": "cycle",
1972
+ "values": [8191.75, 8192.125, 8192.5, 8191.9375, 8193.125, 8191.25, 8194.0, 8190.75, 8192.25]
1973
+ }
1974
+ },
1975
+ "scale": {
1976
+ "dtype": "float32",
1977
+ "shape": [3, 64],
1978
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
1979
+ },
1980
+ "b": {
1981
+ "dtype": "float32",
1982
+ "shape": [3, 64],
1983
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
1984
+ }
1985
+ },
1986
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 64], "tolerance": 0.000002, "relTolerance": 0.00001 } },
1987
+ "provenance": {
1988
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
1989
+ }
1990
+ },
1991
+ {
1992
+ "name": "broadcast_affine_vec4_outer_mixed",
1993
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
1994
+ "inputs": {
1995
+ "x": {
1996
+ "dtype": "float32",
1997
+ "shape": [2, 3, 4, 64],
1998
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
1999
+ },
2000
+ "scale": {
2001
+ "dtype": "float32",
2002
+ "shape": [2, 1, 4, 64],
2003
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2004
+ },
2005
+ "b": {
2006
+ "dtype": "float32",
2007
+ "shape": [1, 3, 1, 64],
2008
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
2009
+ }
2010
+ },
2011
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 4, 64], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2012
+ "provenance": {
2013
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2014
+ }
2015
+ },
2016
+ {
2017
+ "name": "broadcast_affine_vec4_shared_scale_full_bias",
2018
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2019
+ "inputs": {
2020
+ "x": {
2021
+ "dtype": "float32",
2022
+ "shape": [3, 12],
2023
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2024
+ },
2025
+ "scale": {
2026
+ "dtype": "float32",
2027
+ "shape": [12],
2028
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2029
+ },
2030
+ "b": {
2031
+ "dtype": "float32",
2032
+ "shape": [3, 12],
2033
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
2034
+ }
2035
+ },
2036
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 12], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2037
+ "provenance": {
2038
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2039
+ }
2040
+ },
2041
+ {
2042
+ "name": "broadcast_affine_vec4_full_scale_no_bias",
2043
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2044
+ "inputs": {
2045
+ "x": {
2046
+ "dtype": "float32",
2047
+ "shape": [3, 1024],
2048
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2049
+ },
2050
+ "scale": {
2051
+ "dtype": "float32",
2052
+ "shape": [3, 1024],
2053
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2054
+ }
2055
+ },
2056
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 1024], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2057
+ "provenance": {
2058
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2059
+ }
2060
+ },
2061
+ {
2062
+ "name": "broadcast_affine_vec4_right_aligned",
2063
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2064
+ "inputs": {
2065
+ "x": {
2066
+ "dtype": "float32",
2067
+ "shape": [2, 3, 8],
2068
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2069
+ },
2070
+ "scale": {
2071
+ "dtype": "float32",
2072
+ "shape": [3, 8],
2073
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2074
+ }
2075
+ },
2076
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 8], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2077
+ "provenance": {
2078
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2079
+ }
2080
+ },
2081
+ {
2082
+ "name": "broadcast_affine_vec4_mixed_shifted",
2083
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2084
+ "inputs": {
2085
+ "x": {
2086
+ "dtype": "float32",
2087
+ "shape": [2, 3, 8],
2088
+ "data": {
2089
+ "kind": "cycle",
2090
+ "values": [39999.75, 40000.125, 40000.5, 39999.9375, 40001.125, 39999.25, 40002.0, 39998.75, 40000.25]
2091
+ }
2092
+ },
2093
+ "scale": {
2094
+ "dtype": "float32",
2095
+ "shape": [2, 1, 8],
2096
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2097
+ },
2098
+ "b": {
2099
+ "dtype": "float32",
2100
+ "shape": [1, 3, 8],
2101
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
2102
+ }
2103
+ },
2104
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 8], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2105
+ "provenance": {
2106
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2107
+ }
2108
+ },
2109
+ {
2110
+ "name": "broadcast_affine_vec4_full_scale_shared_bias",
2111
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2112
+ "inputs": {
2113
+ "x": {
2114
+ "dtype": "float32",
2115
+ "shape": [2, 3, 8],
2116
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2117
+ },
2118
+ "scale": {
2119
+ "dtype": "float32",
2120
+ "shape": [2, 3, 8],
2121
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2122
+ },
2123
+ "b": {
2124
+ "dtype": "float32",
2125
+ "shape": [8],
2126
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
2127
+ }
2128
+ },
2129
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 8], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2130
+ "provenance": {
2131
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2132
+ }
2133
+ },
2134
+ {
2135
+ "name": "broadcast_affine_vec4_zero_variance",
2136
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2137
+ "inputs": {
2138
+ "x": { "dtype": "float32", "shape": [2, 3, 8], "data": { "kind": "constant", "value": 40000.0 } },
2139
+ "scale": {
2140
+ "dtype": "float32",
2141
+ "shape": [2, 1, 8],
2142
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2143
+ },
2144
+ "b": {
2145
+ "dtype": "float32",
2146
+ "shape": [1, 3, 8],
2147
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
2148
+ }
2149
+ },
2150
+ "outputs": { "y": { "dtype": "float32", "shape": [2, 3, 8], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2151
+ "provenance": {
2152
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2153
+ }
2154
+ },
2155
+ {
2156
+ "name": "broadcast_affine_vec4_batch_tail",
2157
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2158
+ "inputs": {
2159
+ "x": {
2160
+ "dtype": "float32",
2161
+ "shape": [17, 64],
2162
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2163
+ },
2164
+ "scale": {
2165
+ "dtype": "float32",
2166
+ "shape": [17, 64],
2167
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2168
+ },
2169
+ "b": {
2170
+ "dtype": "float32",
2171
+ "shape": [17, 64],
2172
+ "data": { "kind": "cycle", "values": [-0.125, 0.0, 0.25, 1.5, -0.75, 0.5, 0.125] }
2173
+ }
2174
+ },
2175
+ "outputs": { "y": { "dtype": "float32", "shape": [17, 64], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2176
+ "provenance": {
2177
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2178
+ }
2179
+ },
2180
+ {
2181
+ "name": "broadcast_affine_vec4_batch_tail_no_bias",
2182
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2183
+ "inputs": {
2184
+ "x": {
2185
+ "dtype": "float32",
2186
+ "shape": [17, 64],
2187
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2188
+ },
2189
+ "scale": {
2190
+ "dtype": "float32",
2191
+ "shape": [17, 64],
2192
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2193
+ }
2194
+ },
2195
+ "outputs": { "y": { "dtype": "float32", "shape": [17, 64], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2196
+ "provenance": {
2197
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2198
+ }
2199
+ },
2200
+ {
2201
+ "name": "broadcast_affine_vec4_wide_full_scale",
2202
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2203
+ "inputs": {
2204
+ "x": {
2205
+ "dtype": "float32",
2206
+ "shape": [3, 4096],
2207
+ "data": { "kind": "cycle", "values": [-0.25, 0.125, 0.5, -0.0625, 1.125, -0.75, 2.0, -1.25, 0.25] }
2208
+ },
2209
+ "scale": {
2210
+ "dtype": "float32",
2211
+ "shape": [3, 4096],
2212
+ "data": { "kind": "cycle", "values": [0.25, 0.5, 1.0, 1.5, -0.5] }
2213
+ }
2214
+ },
2215
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 4096], "tolerance": 0.000002, "relTolerance": 0.00001 } },
2216
+ "provenance": {
2217
+ "notes": "Packed broadcast-affine rows: independent outer indexing, shifted moments, and partial row groups are checked against the CPU reference without changing tolerance."
2218
+ }
2219
+ },
2220
+ {
2221
+ "name": "f32_last_axis_vec2_scale_only_h126",
2222
+ "attrs": { "axis": 1, "epsilon": 0.00001 },
2223
+ "inputs": {
2224
+ "x": {
2225
+ "dtype": "float32",
2226
+ "shape": [3, 126],
2227
+ "data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.13 }
2228
+ },
2229
+ "scale": {
2230
+ "dtype": "float32",
2231
+ "shape": [126],
2232
+ "data": { "kind": "cycle", "values": [0.5, -1.0, 1.25, 2.0, -0.25] }
2233
+ }
2234
+ },
2235
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 126], "tolerance": 0.00001 } }
2236
+ },
2237
+ {
2238
+ "name": "f32_last_axis_vec2_scale_only_h514",
2239
+ "attrs": { "axis": 1, "epsilon": 0.00001 },
2240
+ "inputs": {
2241
+ "x": {
2242
+ "dtype": "float32",
2243
+ "shape": [3, 514],
2244
+ "data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.13 }
2245
+ },
2246
+ "scale": {
2247
+ "dtype": "float32",
2248
+ "shape": [514],
2249
+ "data": { "kind": "cycle", "values": [0.5, -1.0, 1.25, 2.0, -0.25] }
2250
+ }
2251
+ },
2252
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 514], "tolerance": 0.00001 } }
2253
+ },
2254
+ {
2255
+ "name": "f32_last_axis_vec2_scale_only_h1022",
2256
+ "attrs": { "axis": 1, "epsilon": 0.00001 },
2257
+ "inputs": {
2258
+ "x": {
2259
+ "dtype": "float32",
2260
+ "shape": [3, 1022],
2261
+ "data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.13 }
2262
+ },
2263
+ "scale": {
2264
+ "dtype": "float32",
2265
+ "shape": [1022],
2266
+ "data": { "kind": "cycle", "values": [0.5, -1.0, 1.25, 2.0, -0.25] }
2267
+ }
2268
+ },
2269
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 1022], "tolerance": 0.00001 } }
2270
+ },
2271
+ {
2272
+ "name": "f32_last_axis_vec2_scale_only_h4098",
2273
+ "attrs": { "axis": 1, "epsilon": 0.00001 },
2274
+ "inputs": {
2275
+ "x": {
2276
+ "dtype": "float32",
2277
+ "shape": [3, 4098],
2278
+ "data": { "kind": "fillFloat32", "sinStep": 0.19, "cosStep": 0.13 }
2279
+ },
2280
+ "scale": {
2281
+ "dtype": "float32",
2282
+ "shape": [4098],
2283
+ "data": { "kind": "cycle", "values": [0.5, -1.0, 1.25, 2.0, -0.25] }
2284
+ }
2285
+ },
2286
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 4098], "tolerance": 0.00001 } }
2287
+ },
2288
+ {
2289
+ "name": "f32_last_axis_vec2_shifted_bias_h6",
2290
+ "attrs": { "axis": -1, "epsilon": 0.00001 },
2291
+ "inputs": {
2292
+ "x": {
2293
+ "dtype": "float32",
2294
+ "shape": [3, 6],
2295
+ "data": {
2296
+ "kind": "values",
2297
+ "values": [7.0, 7.0, 7.0, 7.0, 7.0, 7.0, 10000.0, 10000.125, 9999.875, 10000.25, 9999.75, 10000.0, -8.0, 2.0, 0.0, 7.0, -4.0, 3.0]
2298
+ }
2299
+ },
2300
+ "scale": { "dtype": "float32", "shape": [6], "data": { "kind": "cycle", "values": [0.5, -1.0, 1.25] } },
2301
+ "b": { "dtype": "float32", "shape": [6], "data": { "kind": "cycle", "values": [-0.25, 1.0, 0.75] } }
2302
+ },
2303
+ "outputs": { "y": { "dtype": "float32", "shape": [3, 6], "tolerance": 0.00001 } }
2304
  }
2305
  ]
2306
  }