sync 6fdf6301e2bb
Browse files- README.md +13 -2
- build/webgpu/layer-normalization.wgsl.jinja +193 -54
- build/webgpu/manifest.json +263 -80
- build/webgpu/metadata.json +14 -8
- build/webgpu/norm-row-stats.wgsl.jinja +370 -119
- build/webgpu/test.json +518 -0
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
|
| 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.
|
| 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 |
-
{%
|
| 38 |
{% set op_numel = namespace(value=1) %}
|
| 39 |
-
{% for d in opShape %}
|
|
|
|
|
|
|
| 40 |
{% set out_numel = namespace(value=1) %}
|
| 41 |
-
{% for d in outShape %}
|
| 42 |
-
{
|
| 43 |
-
{%
|
| 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 }}] =
|
| 80 |
-
{
|
| 81 |
-
{
|
| 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 |
-
|
| 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":
|
| 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": { "
|
| 58 |
-
"scale": { "
|
| 59 |
-
"y": { "
|
| 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", "
|
| 68 |
-
"mean_out": { "arg": "mean", "
|
| 69 |
-
"inv_std_out": { "arg": "invStdDev", "
|
| 70 |
-
"
|
| 71 |
-
"
|
| 72 |
-
"
|
| 73 |
-
"
|
| 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":
|
| 90 |
-
"writeStats":
|
| 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":
|
| 121 |
-
"writeStats":
|
| 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":
|
| 149 |
-
"writeStats":
|
| 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":
|
| 180 |
-
"writeStats":
|
| 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":
|
| 208 |
-
"writeStats":
|
| 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":
|
| 239 |
-
"writeStats":
|
| 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":
|
| 267 |
-
"writeStats":
|
| 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":
|
| 298 |
-
"writeStats":
|
| 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":
|
| 548 |
-
"writeMean":
|
| 549 |
-
"writeInvStdDev":
|
| 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": ["
|
| 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":
|
| 572 |
-
"writeMean":
|
| 573 |
-
"writeInvStdDev":
|
| 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": ["
|
| 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":
|
| 596 |
-
"writeMean":
|
| 597 |
-
"writeInvStdDev":
|
| 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": ["
|
| 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":
|
| 620 |
-
"writeMean":
|
| 621 |
-
"writeInvStdDev":
|
| 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": ["
|
| 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": "
|
| 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": "
|
| 12 |
-
"manifest.json": "
|
| 13 |
-
"norm-row-stats.wgsl.jinja": "
|
| 14 |
-
"test.json": "
|
| 15 |
}
|
| 16 |
},
|
| 17 |
-
"provenance": { "kernel": { "sha": "
|
| 18 |
"webgpu": {
|
| 19 |
-
"manifestSpec": "2.
|
| 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 |
-
{%
|
| 2 |
-
|
| 3 |
-
{%
|
| 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 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
{% endif %}
|
| 37 |
-
{% if
|
| 38 |
-
|
| 39 |
-
|
|
|
|
| 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 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 52 |
}
|
| 53 |
{% endif %}
|
| 54 |
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 59 |
}
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
}
|
| 64 |
{% endif %}
|
|
|
|
|
|
|
| 65 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 66 |
{% if vec4 %}
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 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 |
-
|
| 150 |
-
|
| 151 |
-
{
|
| 152 |
-
|
|
|
|
| 153 |
{% endif %}
|
| 154 |
|
|
|
|
| 155 |
{% if vec4 %}
|
| 156 |
-
{% if
|
| 157 |
-
|
|
|
|
| 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 |
-
|
| 169 |
-
let
|
| 170 |
-
|
| 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 =
|
| 221 |
-
{% endif %}
|
| 222 |
-
var value = (v - vec4<f32>(row_mean)) * inv * vec4<f32>(scale[i]);
|
| 223 |
{% if hasBias %}
|
| 224 |
-
value = value +
|
| 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 |
}
|