Xenova HF Staff commited on
Commit
b635b3b
·
verified ·
1 Parent(s): 28e88e5

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -51,7 +51,7 @@ Attributes and default values (overridable per request):
51
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
52
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
53
  - [`test.json`](build/webgpu/test.json) — correctness cases
54
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
55
  - [`group-normalization-splitk-apply.wgsl.jinja`](build/webgpu/group-normalization-splitk-apply.wgsl.jinja)
56
  - [`group-normalization-splitk-partials.wgsl.jinja`](build/webgpu/group-normalization-splitk-partials.wgsl.jinja)
57
  - [`group-normalization-stash-f16-serial.wgsl.jinja`](build/webgpu/group-normalization-stash-f16-serial.wgsl.jinja)
@@ -60,7 +60,7 @@ Attributes and default values (overridable per request):
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.
 
51
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
52
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
53
  - [`test.json`](build/webgpu/test.json) — correctness cases
54
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
55
  - [`group-normalization-splitk-apply.wgsl.jinja`](build/webgpu/group-normalization-splitk-apply.wgsl.jinja)
56
  - [`group-normalization-splitk-partials.wgsl.jinja`](build/webgpu/group-normalization-splitk-partials.wgsl.jinja)
57
  - [`group-normalization-stash-f16-serial.wgsl.jinja`](build/webgpu/group-normalization-stash-f16-serial.wgsl.jinja)
 
60
  ## Use with `@huggingface/kernels`
61
 
62
  ```sh
63
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
64
  ```
65
 
66
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/group-normalization-splitk-apply.wgsl.jinja CHANGED
@@ -10,7 +10,8 @@ const EPSILON: f32 = {{ epsilon }};
10
 
11
  @compute @workgroup_size(WG, 1, 1)
12
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
13
- let row = wg.x;
 
14
  let part = wg.z;
15
  if (row >= params.rows) { return; }
16
  var pair = vec2<f32>(0.0);
 
10
 
11
  @compute @workgroup_size(WG, 1, 1)
12
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
13
+ // wg.y carries the row index past the dispatch fold width; wg.z is the part.
14
+ let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
15
  let part = wg.z;
16
  if (row >= params.rows) { return; }
17
  var pair = vec2<f32>(0.0);
build/webgpu/group-normalization-splitk-partials.wgsl.jinja CHANGED
@@ -7,7 +7,8 @@ var<workgroup> reduction: array<vec2<f32>, WG>;
7
 
8
  @compute @workgroup_size(WG, 1, 1)
9
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
10
- let row = wg.x;
 
11
  let part = wg.z;
12
  if (row >= params.rows) { return; }
13
  let tid = lid.x;
 
7
 
8
  @compute @workgroup_size(WG, 1, 1)
9
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
10
+ // wg.y carries the row index past the dispatch fold width; wg.z is the part.
11
+ let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
12
  let part = wg.z;
13
  if (row >= params.rows) { return; }
14
  let tid = lid.x;
build/webgpu/group-normalization-stash-f16-serial.wgsl.jinja CHANGED
@@ -1,5 +1,3 @@
1
- {% if usesF16 %}enable f16;
2
- {% endif %}
3
  {{ env.wgsl.resourceDeclarations }}
4
 
5
  const HIDDEN: u32 = {{ hiddenSize }}u;
@@ -61,7 +59,6 @@ fn widen_f16_bits(value: u32) -> f32 {
61
  return unpack2x16float(value & 0xffffu).x;
62
  }
63
 
64
-
65
  // ONNX GroupNormalization-21 expresses the stash_type=FLOAT16 stage as a
66
  // graph of f16 tensor operators. One invocation owns a complete group so each
67
  // intermediate addition and arithmetic stage remains f16.
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  const HIDDEN: u32 = {{ hiddenSize }}u;
 
59
  return unpack2x16float(value & 0xffffu).x;
60
  }
61
 
 
62
  // ONNX GroupNormalization-21 expresses the stash_type=FLOAT16 stage as a
63
  // graph of f16 tensor operators. One invocation owns a complete group so each
64
  // intermediate addition and arithmetic stage remains f16.
build/webgpu/manifest.json CHANGED
@@ -40,12 +40,11 @@
40
  "groupSplitCovered": "groupRowCovered and groupRows <= tunables.SPLIT_STATS_MAX_ROWS and groupRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and groupHidden >= tunables.SPLIT_STATS_MIN_HIDDEN and groupPartialBytes <= device.limits.maxStorageBufferBindingSize and groupPartialBytes <= device.limits.maxBufferSize"
41
  },
42
  "bindings": {
43
- "x": { "buffer": "read-only-storage", "elementType": "$ioElement" },
44
- "scale": { "buffer": "read-only-storage", "elementType": "$scalar" },
45
- "bias": { "buffer": "read-only-storage", "elementType": "$scalar" },
46
- "y": { "buffer": "storage", "elementType": "$ioElement" },
47
  "params": {
48
- "buffer": "uniform",
49
  "struct": [
50
  { "name": "rows", "type": "u32", "value": "groupRows" },
51
  {
@@ -55,12 +54,8 @@
55
  }
56
  ]
57
  },
58
- "x_2": { "name": "x", "buffer": "read-only-storage", "elementType": "$scalar" },
59
- "params_2": {
60
- "name": "params",
61
- "buffer": "uniform",
62
- "struct": [{ "name": "rows", "type": "u32", "value": "groupRows" }]
63
- }
64
  },
65
  "variants": [
66
  {
@@ -70,7 +65,6 @@
70
  "derive": {
71
  "scalar": "dtypes.T",
72
  "ioElement": "dtypes.T",
73
- "usesF16": "dtypes.T == \"f16\"",
74
  "hiddenSize": "groupHidden",
75
  "spatial": "groupSpatial",
76
  "channelsPerGroup": "groupChannelsPerGroup",
@@ -93,7 +87,6 @@
93
  "when": ["groupSplitCovered"],
94
  "derive": {
95
  "scalar": "dtypes.T",
96
- "usesF16": "dtypes.T == \"f16\"",
97
  "hiddenSize": "groupHidden",
98
  "spatial": "groupSpatial",
99
  "channelsPerGroup": "groupChannelsPerGroup",
@@ -108,7 +101,7 @@
108
  "id": "partials",
109
  "name": "GroupNormalization.SplitKPartials",
110
  "shader": "group-normalization-splitk-partials.wgsl.jinja",
111
- "bindings": ["x_2", { "name": "partials", "buffer": "storage", "elementType": "vec2<f32>" }, "params_2"],
112
  "dispatch": {
113
  "x": "min(groupRows, DISPATCH_FOLD_WIDTH)",
114
  "y": "ceilDiv(groupRows, DISPATCH_FOLD_WIDTH)",
@@ -120,12 +113,12 @@
120
  "name": "GroupNormalization.SplitKApply",
121
  "shader": "group-normalization-splitk-apply.wgsl.jinja",
122
  "bindings": [
123
- "x_2",
124
  "scale",
125
  "bias",
126
  { "name": "partials", "buffer": "read-only-storage", "elementType": "vec2<f32>" },
127
  { "arg": "y", "elementType": "$scalar" },
128
- "params_2"
129
  ],
130
  "dispatch": {
131
  "x": "min(groupRows, DISPATCH_FOLD_WIDTH)",
@@ -143,13 +136,11 @@
143
  "passes": [
144
  {
145
  "id": "main",
146
- "name": "GroupNormalization.group_subgroup_vec4",
147
  "shader": "norm-row-stats.wgsl.jinja",
148
  "derive": {
149
- "modeSpec": "\"group\"",
150
  "vec4": true,
151
  "scalar": "dtypes.T",
152
- "usesF16Spec": "dtypes.T == \"f16\"",
153
  "hidden": "groupHidden",
154
  "wg": "groupVec4Workgroup",
155
  "epsilon": "attrs.epsilon",
@@ -161,8 +152,7 @@
161
  "combineSubgroups": "hasSubgroupId"
162
  },
163
  "bindings": ["x", "scale", "bias", "y", "params"],
164
- "dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 },
165
- "subgroupCollectivesWidth": "portable"
166
  }
167
  ]
168
  },
@@ -174,13 +164,11 @@
174
  "passes": [
175
  {
176
  "id": "main",
177
- "name": "GroupNormalization.group_subgroup",
178
  "shader": "norm-row-stats.wgsl.jinja",
179
  "derive": {
180
- "modeSpec": "\"group\"",
181
  "vec4": false,
182
  "scalar": "dtypes.T",
183
- "usesF16Spec": "dtypes.T == \"f16\"",
184
  "hidden": "groupHidden",
185
  "wg": "groupScalarWorkgroup",
186
  "epsilon": "attrs.epsilon",
@@ -190,8 +178,7 @@
190
  "combineSubgroups": "hasSubgroupId"
191
  },
192
  "bindings": ["x", "scale", "bias", "y", "params"],
193
- "dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 },
194
- "subgroupCollectivesWidth": "portable"
195
  }
196
  ]
197
  }
 
40
  "groupSplitCovered": "groupRowCovered and groupRows <= tunables.SPLIT_STATS_MAX_ROWS and groupRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and groupHidden >= tunables.SPLIT_STATS_MIN_HIDDEN and groupPartialBytes <= device.limits.maxStorageBufferBindingSize and groupPartialBytes <= device.limits.maxBufferSize"
41
  },
42
  "bindings": {
43
+ "x": { "elementType": "$ioElement" },
44
+ "scale": { "elementType": "$scalar" },
45
+ "bias": { "elementType": "$scalar" },
46
+ "y": { "elementType": "$ioElement" },
47
  "params": {
 
48
  "struct": [
49
  { "name": "rows", "type": "u32", "value": "groupRows" },
50
  {
 
54
  }
55
  ]
56
  },
57
+ "x_apply": { "name": "x", "elementType": "$scalar" },
58
+ "params_rows": { "name": "params", "struct": [{ "name": "rows", "type": "u32", "value": "groupRows" }] }
 
 
 
 
59
  },
60
  "variants": [
61
  {
 
65
  "derive": {
66
  "scalar": "dtypes.T",
67
  "ioElement": "dtypes.T",
 
68
  "hiddenSize": "groupHidden",
69
  "spatial": "groupSpatial",
70
  "channelsPerGroup": "groupChannelsPerGroup",
 
87
  "when": ["groupSplitCovered"],
88
  "derive": {
89
  "scalar": "dtypes.T",
 
90
  "hiddenSize": "groupHidden",
91
  "spatial": "groupSpatial",
92
  "channelsPerGroup": "groupChannelsPerGroup",
 
101
  "id": "partials",
102
  "name": "GroupNormalization.SplitKPartials",
103
  "shader": "group-normalization-splitk-partials.wgsl.jinja",
104
+ "bindings": ["x_apply", { "name": "partials", "elementType": "vec2<f32>" }, "params_rows"],
105
  "dispatch": {
106
  "x": "min(groupRows, DISPATCH_FOLD_WIDTH)",
107
  "y": "ceilDiv(groupRows, DISPATCH_FOLD_WIDTH)",
 
113
  "name": "GroupNormalization.SplitKApply",
114
  "shader": "group-normalization-splitk-apply.wgsl.jinja",
115
  "bindings": [
116
+ "x_apply",
117
  "scale",
118
  "bias",
119
  { "name": "partials", "buffer": "read-only-storage", "elementType": "vec2<f32>" },
120
  { "arg": "y", "elementType": "$scalar" },
121
+ "params_rows"
122
  ],
123
  "dispatch": {
124
  "x": "min(groupRows, DISPATCH_FOLD_WIDTH)",
 
136
  "passes": [
137
  {
138
  "id": "main",
139
+ "name": "GroupNormalization.GroupSubgroupVec4",
140
  "shader": "norm-row-stats.wgsl.jinja",
141
  "derive": {
 
142
  "vec4": true,
143
  "scalar": "dtypes.T",
 
144
  "hidden": "groupHidden",
145
  "wg": "groupVec4Workgroup",
146
  "epsilon": "attrs.epsilon",
 
152
  "combineSubgroups": "hasSubgroupId"
153
  },
154
  "bindings": ["x", "scale", "bias", "y", "params"],
155
+ "dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 }
 
156
  }
157
  ]
158
  },
 
164
  "passes": [
165
  {
166
  "id": "main",
167
+ "name": "GroupNormalization.GroupSubgroup",
168
  "shader": "norm-row-stats.wgsl.jinja",
169
  "derive": {
 
170
  "vec4": false,
171
  "scalar": "dtypes.T",
 
172
  "hidden": "groupHidden",
173
  "wg": "groupScalarWorkgroup",
174
  "epsilon": "attrs.epsilon",
 
178
  "combineSubgroups": "hasSubgroupId"
179
  },
180
  "bindings": ["x", "scale", "bias", "y", "params"],
181
+ "dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 }
 
182
  }
183
  ]
184
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.GroupNormalization",
3
- "id": "_ai_onnx_groupnormalization_webgpu_a85b638",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,17 +8,17 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "gttctaT32ACO8eeDLTBubmz2KmY+JbRp40eLp4BLAiw=",
11
- "group-normalization-splitk-apply.wgsl.jinja": "bJQ3aD6iCak7YjlH5yZZGPo8FIGY8CuCz3BwVnTnUPE=",
12
- "group-normalization-splitk-partials.wgsl.jinja": "DVgljhomjpQ4XizLxIEh8NckAby7lIb6+plUaGh812A=",
13
- "group-normalization-stash-f16-serial.wgsl.jinja": "Wez9kqS+lzZASbm2BpWZSsHuNTT4rqsus0NiyhmGmvY=",
14
- "manifest.json": "A56uOwydeaOGyMd1ulwyeOrEUd7G8uLRbxougyW6s8g=",
15
- "norm-row-stats.wgsl.jinja": "CyRuHHc7bYmXEhtvfCxLRvjhidAMvicRA5nuJJESwxg=",
16
  "test.json": "jfHQT3aDoszhxWXbi2YtXb9bqMfNRQyaM3HFdPdOWA0="
17
  }
18
  },
19
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
20
  "webgpu": {
21
- "manifestSpec": "2.0",
22
  "variants": {
23
  "group_stash_f16_serial": ["group-normalization-stash-f16-serial.wgsl.jinja"],
24
  "group_splitk": ["group-normalization-splitk-apply.wgsl.jinja", "group-normalization-splitk-partials.wgsl.jinja"],
 
1
  {
2
  "name": "ai.onnx.GroupNormalization",
3
+ "id": "_ai_onnx_groupnormalization_webgpu_c63203e",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "gttctaT32ACO8eeDLTBubmz2KmY+JbRp40eLp4BLAiw=",
11
+ "group-normalization-splitk-apply.wgsl.jinja": "teD7D7rpdNBVIoXSGMhWtjYhB5HegbAuP9Uq7u1socc=",
12
+ "group-normalization-splitk-partials.wgsl.jinja": "Tif7IYCZjN0rU9RwOjuDJRuL+nWRA+fP9NurRkyxfGA=",
13
+ "group-normalization-stash-f16-serial.wgsl.jinja": "R+PB+omhL6ghmEal1QbolJcaSJtiayYgBs0lPi3WcQA=",
14
+ "manifest.json": "QmDFj0Qi2AT20Sw051jLbEE+/46hPAiZMuFkp6ydMfE=",
15
+ "norm-row-stats.wgsl.jinja": "5fSl/yJ/eXZXSkHd/WTM3uDq8Wj6W/0geOtl9i36uXM=",
16
  "test.json": "jfHQT3aDoszhxWXbi2YtXb9bqMfNRQyaM3HFdPdOWA0="
17
  }
18
  },
19
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
20
  "webgpu": {
21
+ "manifestSpec": "2.1",
22
  "variants": {
23
  "group_stash_f16_serial": ["group-normalization-stash-f16-serial.wgsl.jinja"],
24
  "group_splitk": ["group-normalization-splitk-apply.wgsl.jinja", "group-normalization-splitk-partials.wgsl.jinja"],
build/webgpu/norm-row-stats.wgsl.jinja CHANGED
@@ -1,16 +1,5 @@
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 rmsChainNorm = rmsChainNorm if rmsChainNorm is defined else false %}
8
- {% set hiddenPairs = hiddenPairs | default(0) %}
9
- {% set numRows = numRows | default(0) %}
10
- {% set epsilon = epsilon | default("0.0") %}
11
- {% set epsilon2 = epsilon2 | default("0.0") %}
12
- {% set numGroupsSpec = numGroupsSpec | default(0) %}
13
- {% set cpg = cpg | default(0) %}
14
  {% set spatialVec = spatialVec | default(0) %}
15
  {% set spatial = spatial | default(0) %}
16
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
@@ -36,15 +25,8 @@ const HIDDEN: u32 = {{ hidden }}u;
36
  {% if vec4 %}
37
  const HIDDEN_V: u32 = {{ hiddenVec }}u;
38
  {% endif %}
39
- {% if packedBf16Embedding %}
40
- const HIDDEN_PAIRS: u32 = {{ hiddenPairs }}u;
41
- const NUM_ROWS: u32 = {{ numRows }}u;
42
- {% endif %}
43
  const WG: u32 = {{ wg }}u;
44
  const EPSILON: f32 = {{ epsilon }};
45
- {% if rmsChainNorm %}
46
- const EPSILON2: f32 = {{ epsilon2 }};
47
- {% endif %}
48
  const NUM_GROUPS: u32 = {{ numGroupsSpec }}u;
49
  const CPG: u32 = {{ cpg }}u;
50
  {% if vec4 %}
@@ -53,44 +35,6 @@ const SPATIAL_V: u32 = {{ spatialVec }}u;
53
  const SPATIAL: u32 = {{ spatial }}u;
54
  {% endif %}
55
 
56
- {% if packedBf16Embedding %}
57
- {% if vec4 %}
58
- fn unpack_bf16_pair(word: u32) -> vec2<f32> {
59
- let bits = vec2<u32>(word & 0xffffu, word >> 16u);
60
- return bitcast<vec2<f32>>(bits << vec2<u32>(16u));
61
- }
62
- {% endif %}
63
-
64
- {% if not vec4 %}
65
- fn embedding_scalar(source_row: u32, hidden: u32) -> f32 {
66
- if (source_row >= NUM_ROWS) {
67
- return 0.0;
68
- }
69
- let word = x[source_row * HIDDEN_PAIRS + (hidden >> 1u)];
70
- let bits = select(word & 0xffffu, word >> 16u, (hidden & 1u) != 0u);
71
- return bitcast<f32>(bits << 16u);
72
- }
73
- {% endif %}
74
-
75
- {% if vec4 %}
76
- fn embedding_vec4(source_row: u32, hidden_vec: u32) -> vec4<f32> {
77
- if (source_row >= NUM_ROWS) {
78
- return vec4<f32>(0.0);
79
- }
80
- let base = source_row * HIDDEN_PAIRS + hidden_vec * 2u;
81
- let low = unpack_bf16_pair(x[base]);
82
- let high = unpack_bf16_pair(x[base + 1u]);
83
- return vec4<f32>(low, high);
84
- }
85
- {% endif %}
86
- {% endif %}
87
-
88
- {% if vec4 and scalarIo %}
89
- fn load_vec4(index: u32) -> vec4<f32> {
90
- return vec4<f32>(x[index], x[index + 1u], x[index + 2u], x[index + 3u]);
91
- }
92
- {% endif %}
93
-
94
  {% if combineSubgroups %}
95
  var<workgroup> sg_partials: array<vec2<f32>, WG>;
96
 
@@ -148,25 +92,14 @@ fn main(
148
  return;
149
  }
150
  let tid = lid.x;
151
- {% if packedBf16Embedding %}
152
- let source_row = indices[row];
153
- {% if vec4 %}
154
- let base = row * HIDDEN_V;
155
- {% else %}
156
- let base = row * HIDDEN;
157
- {% endif %}
158
- {% elif vec4 and not scalarIo %}
159
  let base = row * HIDDEN_V;
160
  {% else %}
161
  let base = row * HIDDEN;
162
  {% endif %}
163
 
164
  {% if vec4 %}
165
- {% if scalarIo %}
166
- let shift = f32(x[base]);
167
- {% else %}
168
  let shift = f32(x[base].x);
169
- {% endif %}
170
  {% else %}
171
  let shift = f32(x[base]);
172
  {% endif %}
@@ -174,26 +107,14 @@ fn main(
174
  var acc = vec2<f32>(0.0, 0.0);
175
  {% if vec4 %}
176
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
177
- {% if packedBf16Embedding %}
178
- let v = embedding_vec4(source_row, i);
179
- embedding_out[base + i] = v;
180
- {% elif scalarIo %}
181
- let v = load_vec4(base + i * 4u);
182
- {% else %}
183
- let v = vec4<f32>(x[base + i]);
184
- {% endif %}
185
- let d = v - vec4<f32>(shift);
186
  acc.x = acc.x + d.x + d.y + d.z + d.w;
187
  acc.y = acc.y + dot(d, d);
188
  }
189
  {% else %}
190
  for (var i = tid; i < HIDDEN; i = i + WG) {
191
- {% if packedBf16Embedding %}
192
- let v = embedding_scalar(source_row, i);
193
- embedding_out[base + i] = v;
194
- {% else %}
195
  let v = f32(x[base + i]);
196
- {% endif %}
197
  let d = v - shift;
198
  acc.x = acc.x + d;
199
  acc.y = acc.y + d * d;
@@ -208,48 +129,18 @@ fn main(
208
  let row_mean = shift + mean_d;
209
  let g_ch_base = (row % NUM_GROUPS) * CPG;
210
 
211
- {% if rmsChainNorm %}
212
- var acc2 = 0.0;
213
- {% endif %}
214
  {% if vec4 %}
215
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
216
- {% if packedBf16Embedding %}
217
  let idx = base + i;
218
- let v = embedding_vec4(source_row, i);
219
- {% elif scalarIo %}
220
- let idx = base + i * 4u;
221
- let v = load_vec4(idx);
222
- {% else %}
223
- let idx = base + i;
224
- let v = vec4<f32>(x[idx]);
225
- {% endif %}
226
  let ch = g_ch_base + i / SPATIAL_V;
227
  let normed = (v - vec4<f32>(row_mean)) / vec4<f32>(denom);
228
  y[idx] = {{ vecType }}(normed * vec4<f32>(f32(scale[ch])) + vec4<f32>(f32(bias[ch])));
229
  }
230
- {% if rmsChainNorm %}
231
-
232
- // The chained second norm reads the residual row this loop just stored. This
233
- // barrier completes those stores and any preceding shared-scratch use before
234
- // the next reduction reuses its scratch; each lane then re-reads only the
235
- // elements it wrote itself.
236
- workgroupBarrier();
237
- let total2 = reduce_scalar(acc2{{ reduceThreadArguments }});
238
- let inv2 = inverseSqrt(total2 / f32(HIDDEN) + EPSILON2);
239
- for (var i = tid; i < HIDDEN_V; i = i + WG) {
240
- let idx = base + i;
241
- let hv = vec4<f32>(y[idx]);
242
- normed2[idx] = {{ vecType }}(hv * inv2 * vec4<f32>(scale2[i]));
243
- }
244
- {% endif %}
245
  {% else %}
246
  for (var i = tid; i < HIDDEN; i = i + WG) {
247
  let idx = base + i;
248
- {% if packedBf16Embedding %}
249
- let v = embedding_scalar(source_row, i);
250
- {% else %}
251
  let v = f32(x[idx]);
252
- {% endif %}
253
  let ch = g_ch_base + i / SPATIAL;
254
  let normed = (v - row_mean) / denom;
255
  y[idx] = {{ scalar }}(normed * f32(scale[ch]) + f32(bias[ch]));
 
1
+ {% set scalarIo = false %}
2
+ {% set packedF32 = "vec4<f32>" %}
 
 
 
 
 
 
 
 
 
 
 
3
  {% set spatialVec = spatialVec | default(0) %}
4
  {% set spatial = spatial | default(0) %}
5
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
 
25
  {% if vec4 %}
26
  const HIDDEN_V: u32 = {{ hiddenVec }}u;
27
  {% endif %}
 
 
 
 
28
  const WG: u32 = {{ wg }}u;
29
  const EPSILON: f32 = {{ epsilon }};
 
 
 
30
  const NUM_GROUPS: u32 = {{ numGroupsSpec }}u;
31
  const CPG: u32 = {{ cpg }}u;
32
  {% if vec4 %}
 
35
  const SPATIAL: u32 = {{ spatial }}u;
36
  {% endif %}
37
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
38
  {% if combineSubgroups %}
39
  var<workgroup> sg_partials: array<vec2<f32>, WG>;
40
 
 
92
  return;
93
  }
94
  let tid = lid.x;
95
+ {% if vec4 and not scalarIo %}
 
 
 
 
 
 
 
96
  let base = row * HIDDEN_V;
97
  {% else %}
98
  let base = row * HIDDEN;
99
  {% endif %}
100
 
101
  {% if vec4 %}
 
 
 
102
  let shift = f32(x[base].x);
 
103
  {% else %}
104
  let shift = f32(x[base]);
105
  {% endif %}
 
107
  var acc = vec2<f32>(0.0, 0.0);
108
  {% if vec4 %}
109
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
110
+ let v = {{ packedF32 }}(x[base + i]);
111
+ let d = v - {{ packedF32 }}(shift);
 
 
 
 
 
 
 
112
  acc.x = acc.x + d.x + d.y + d.z + d.w;
113
  acc.y = acc.y + dot(d, d);
114
  }
115
  {% else %}
116
  for (var i = tid; i < HIDDEN; i = i + WG) {
 
 
 
 
117
  let v = f32(x[base + i]);
 
118
  let d = v - shift;
119
  acc.x = acc.x + d;
120
  acc.y = acc.y + d * d;
 
129
  let row_mean = shift + mean_d;
130
  let g_ch_base = (row % NUM_GROUPS) * CPG;
131
 
 
 
 
132
  {% if vec4 %}
133
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
 
134
  let idx = base + i;
135
+ let v = {{ packedF32 }}(x[idx]);
 
 
 
 
 
 
 
136
  let ch = g_ch_base + i / SPATIAL_V;
137
  let normed = (v - vec4<f32>(row_mean)) / vec4<f32>(denom);
138
  y[idx] = {{ vecType }}(normed * vec4<f32>(f32(scale[ch])) + vec4<f32>(f32(bias[ch])));
139
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
140
  {% else %}
141
  for (var i = tid; i < HIDDEN; i = i + WG) {
142
  let idx = base + i;
 
 
 
143
  let v = f32(x[idx]);
 
144
  let ch = g_ch_base + i / SPATIAL;
145
  let normed = (v - row_mean) / denom;
146
  y[idx] = {{ scalar }}(normed * f32(scale[ch]) + f32(bias[ch]));