Download build/webgpu/layer-normalization.wgsl.jinja from webgpu-kernels/ai.onnx.LayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 11.7 kB
-
https://huggingface.co/kernels/webgpu-kernels/ai.onnx.LayerNormalization/resolve/v1/build/webgpu/layer-normalization.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/ai.onnx.LayerNormalization@v1/build/webgpu/layer-normalization.wgsl.jinja
-
curl -L -o layer-normalization.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/ai.onnx.LayerNormalization/resolve/v1/build/webgpu/layer-normalization.wgsl.jinja
11.7 kB
| {% macro offset_fn(fn_name, opShape, opRank, op_same, op_numel, outShape, outRank, out_numel) %} | |
| fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif %}) -> u32 { | |
| {% if out_numel == 0 %} | |
| return 0u; | |
| {% elif op_numel == 1 %} | |
| return 0u; | |
| {% elif op_same %} | |
| return out_index; | |
| {% else %} | |
| var offset = 0u; | |
| {% for axis in range(outRank) %} | |
| {% set op_axis = axis - (outRank - opRank) %} | |
| {% if op_axis >= 0 and opShape[op_axis] != 1 %} | |
| {% set c_stride = namespace(value=1) %} | |
| {% for j in range(axis + 1, outRank) %} | |
| {% set c_stride.value = c_stride.value * outShape[j] %} | |
| {% endfor %} | |
| {% set op_stride = namespace(value=1) %} | |
| {% for j in range(op_axis + 1, opRank) %} | |
| {% set op_stride.value = op_stride.value * opShape[j] %} | |
| {% endfor %} | |
| {% if c_stride.value == 1 %} | |
| let coord{{ axis }} = out_index % {{ outShape[axis] }}u; | |
| {% else %} | |
| let coord{{ axis }} = (out_index / {{ c_stride.value }}u) % {{ outShape[axis] }}u; | |
| {% endif %} | |
| {% if op_stride.value == 1 %} | |
| offset = offset + coord{{ axis }}; | |
| {% else %} | |
| offset = offset + coord{{ axis }} * {{ op_stride.value }}u; | |
| {% endif %} | |
| {% endif %} | |
| {% endfor %} | |
| return offset; | |
| {% endif %} | |
| }{% endmacro %} | |
| {% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %} | |
| {% set op_numel = namespace(value=1) %} | |
| {% for d in opShape %} | |
| {% set op_numel.value = op_numel.value * d %} | |
| {% endfor %} | |
| {% set out_numel = namespace(value=1) %} | |
| {% for d in outShape %} | |
| {% set out_numel.value = out_numel.value * d %} | |
| {% endfor %} | |
| {{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %}){% endmacro %} | |
| {{ env.wgsl.resourceDeclarations }} | |
| const HIDDEN: u32 = {{ hiddenSize }}u; | |
| const EPSILON: f32 = {{ epsilon }}; | |
| const WG: u32 = {{ workgroupSize }}u; | |
| {% set xNumel = namespace(value=1) %} | |
| {% for dim in xShape %} | |
| {% set xNumel.value = xNumel.value * dim %} | |
| {% endfor %} | |
| {% set scaleNumel = namespace(value=1) %} | |
| {% for dim in scaleShape %} | |
| {% set scaleNumel.value = scaleNumel.value * dim %} | |
| {% endfor %} | |
| {% if scaleNumel.value != 1 %} | |
| {{ offset_fn("scale_offset", scaleShape, scaleShape | length, scaleShape == xShape, scaleNumel.value, xShape, xShape | length, xNumel.value) }} | |
| {% endif %} | |
| {% if hasBias %} | |
| {% set biasNumel = namespace(value=1) %} | |
| {% for dim in biasShape %} | |
| {% set biasNumel.value = biasNumel.value * dim %} | |
| {% endfor %} | |
| {% if biasNumel.value != 1 %} | |
| {{ offset_fn("bias_offset", biasShape, biasShape | length, biasShape == xShape, biasNumel.value, xShape, xShape | length, xNumel.value) }} | |
| {% endif %} | |
| {% endif %} | |
| {% if scalar == "f16" %} | |
| {% set modeSpec = "layer" %} | |
| {% set halfWriteMean = writeMean %} | |
| {% set halfWriteInv = writeInvStdDev %} | |
| {% set halfScaleOffset = "0u" if scaleNumel.value == 1 else broadcast_offset_call("scale_offset", scaleShape, xShape, "base + i") %} | |
| {% if hasBias %} | |
| {% set halfBiasOffset = "0u" if biasNumel.value == 1 else broadcast_offset_call("bias_offset", biasShape, xShape, "base + i") %} | |
| {% endif %} | |
| {% set halfOutputScalar = "f16" %} | |
| fn round_f16_bits_rte(value: f32) -> u32 { | |
| let bits = bitcast<u32>(value); | |
| let sign = (bits >> 16u) & 0x8000u; | |
| let exponent_f32 = (bits >> 23u) & 0xffu; | |
| let mantissa_f32 = bits & 0x7fffffu; | |
| if (exponent_f32 == 0xffu) { | |
| if (mantissa_f32 != 0u) { | |
| return 0x7e00u; | |
| } | |
| return sign | 0x7c00u; | |
| } | |
| var exponent_f16 = i32(exponent_f32) - 127 + 15; | |
| if (exponent_f16 >= 0x1f) { | |
| return sign | 0x7c00u; | |
| } | |
| if (exponent_f16 <= 0) { | |
| if (exponent_f16 < -10) { | |
| return sign; | |
| } | |
| let significand = mantissa_f32 | 0x800000u; | |
| let shift = u32(14 - exponent_f16); | |
| let halfway = 1u << (shift - 1u); | |
| let discarded = significand & ((1u << shift) - 1u); | |
| var fraction = significand >> shift; | |
| if (discarded > halfway || (discarded == halfway && (fraction & 1u) == 1u)) { | |
| fraction = fraction + 1u; | |
| } | |
| return sign | fraction; | |
| } | |
| let halfway = 1u << 12u; | |
| let discarded = mantissa_f32 & 0x1fffu; | |
| var mantissa_f16 = mantissa_f32 >> 13u; | |
| if (discarded > halfway || (discarded == halfway && (mantissa_f16 & 1u) == 1u)) { | |
| mantissa_f16 = mantissa_f16 + 1u; | |
| if (mantissa_f16 == 0x400u) { | |
| mantissa_f16 = 0u; | |
| exponent_f16 = exponent_f16 + 1; | |
| } | |
| } | |
| if (exponent_f16 >= 0x1f) { | |
| return sign | 0x7c00u; | |
| } | |
| return sign | (u32(exponent_f16) << 10u) | mantissa_f16; | |
| } | |
| fn widen_f16_bits(value: u32) -> f32 { | |
| return unpack2x16float(value & 0xffffu).x; | |
| } | |
| // Typed ONNX edges must survive arithmetic fusion and narrow/wide casts. | |
| // Integer rounding also fixes the ties-to-even rule independently of the | |
| // implementation's floating-point conversion rounding mode. | |
| fn half_stage(value: f32) -> f32 { | |
| return widen_f16_bits(round_f16_bits_rte(value)); | |
| } | |
| // Half output magnifies statistics errors at rounding midpoints. Keep a low | |
| // residual through the reduction and normalization, then round the typed | |
| // float32-normalized/half-scale/half-bias edges explicitly. This is the standard ONNX half path; | |
| // other normalization contracts retain their existing arithmetic. | |
| {% set halfWriteMean = halfWriteMean if halfWriteMean is defined else (writeStats and modeSpec == "layer") %} | |
| {% set halfWriteInv = halfWriteInv if halfWriteInv is defined else writeStats %} | |
| fn pair_add(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> { | |
| let s = fma(a.x, 1.0, b.x); | |
| // Materialize each rounded subtraction in the error-free transform. Plain | |
| // cancellation expressions do not preserve the intended evaluation tree on | |
| // every shader backend. | |
| let bv = fma(-1.0, a.x, s); | |
| let av = fma(-1.0, bv, s); | |
| let a_error = fma(-1.0, av, a.x); | |
| let b_error = fma(-1.0, bv, b.x); | |
| let error = fma(a_error, 1.0, b_error); | |
| let e = fma(fma(error, 1.0, a.y), 1.0, b.y); | |
| let hi = fma(s, 1.0, e); | |
| return vec2<f32>(hi, fma(-1.0, fma(-1.0, s, hi), e)); | |
| } | |
| fn pair_mul(a: vec2<f32>, b: vec2<f32>) -> vec2<f32> { | |
| let p = fma(a.x, b.x, 0.0); | |
| let error = fma(a.x, b.x, -p); | |
| return pair_add(vec2<f32>(p, 0.0), vec2<f32>(error + (a.x * b.y + b.x * a.y), 0.0)); | |
| } | |
| fn pair_div(a: vec2<f32>, b: f32) -> vec2<f32> { | |
| let q = a.x / b; | |
| let residual = pair_add(a, -pair_mul(vec2<f32>(q, 0.0), vec2<f32>(b, 0.0))); | |
| return pair_add(vec2<f32>(q, 0.0), vec2<f32>((residual.x + residual.y) / b, 0.0)); | |
| } | |
| fn pair_inverse_sqrt(a: vec2<f32>) -> vec2<f32> { | |
| let r = inverseSqrt(a.x); | |
| let rr = pair_mul(vec2<f32>(r, 0.0), vec2<f32>(r, 0.0)); | |
| let residual = pair_add(vec2<f32>(1.0, 0.0), -pair_mul(a, rr)); | |
| return pair_add(vec2<f32>(r, 0.0), vec2<f32>((0.5 * r) * (residual.x + residual.y), 0.0)); | |
| } | |
| fn half_normalized(value: vec2<f32>) -> f32 { | |
| // stash_type=1 materializes Normalized as float32 before its cast to half. | |
| // Collapse the compensated residual at that typed edge; rounding the pair | |
| // directly to half can choose a different result at a float32 midpoint. | |
| return widen_f16_bits(round_f16_bits_rte(fma(value.x, 1.0, value.y))); | |
| } | |
| var<workgroup> partial: array<vec2<f32>, WG>; | |
| fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> { | |
| partial[tid] = value; | |
| workgroupBarrier(); | |
| for (var step = WG / 2u; step > 0u; step /= 2u) { | |
| if (tid < step) { partial[tid] = pair_add(partial[tid], partial[tid + step]); } | |
| workgroupBarrier(); | |
| } | |
| let total = partial[0]; | |
| workgroupBarrier(); | |
| return total; | |
| } | |
| {% set reduceArgs = "tid" %} | |
| fn load_value(index: u32) -> f32 { | |
| return f32(x[index]); | |
| } | |
| fn normalize_half_row(row: u32, tid: u32 | |
| ) { | |
| if (row >= params.rows) { return; } | |
| let base = row * HIDDEN; | |
| var local_sum = vec2<f32>(0.0); | |
| for (var i = tid; i < HIDDEN; i += WG) { | |
| local_sum = pair_add(local_sum, vec2<f32>(load_value(base + i), 0.0)); | |
| } | |
| let mean = pair_div(reduce_pair(local_sum, {{ reduceArgs }}), f32(HIDDEN)); | |
| var local_square = vec2<f32>(0.0); | |
| for (var i = tid; i < HIDDEN; i += WG) { | |
| let centered = pair_add(vec2<f32>(load_value(base + i), 0.0), -mean); | |
| local_square = pair_add(local_square, pair_mul(centered, centered)); | |
| } | |
| let variance = pair_div(reduce_pair(local_square, {{ reduceArgs }}), f32(HIDDEN)); | |
| let inv = pair_inverse_sqrt(pair_add(variance, vec2<f32>(EPSILON, 0.0))); | |
| {% if halfWriteMean %} | |
| if (tid == 0u) { mean_out[row] = mean.x; } | |
| {% endif %} | |
| {% if halfWriteInv %} | |
| if (tid == 0u) { inv_std_out[row] = inv.x; } | |
| {% endif %} | |
| for (var i = tid; i < HIDDEN; i += WG) { | |
| let normalized = half_normalized(pair_mul(pair_add(vec2<f32>(load_value(base + i), 0.0), -mean), inv)); | |
| var value = half_stage(normalized * f32(scale[{{ halfScaleOffset | default("i") }}])); | |
| {% if modeSpec == "layer" and hasBias %} | |
| value = half_stage(value + f32(bias[{{ halfBiasOffset | default("i") }}])); | |
| {% endif %} | |
| y[base + i] = {{ halfOutputScalar }}(value); | |
| } | |
| } | |
| @compute @workgroup_size(WG, 1, 1) | |
| fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) { | |
| let row = wg.x + wg.y * params.rowStride; | |
| normalize_half_row(row, lid.x); | |
| } | |
| {% else %} | |
| var<workgroup> partial: array<f32, WG>; | |
| var<workgroup> row_mean: f32; | |
| var<workgroup> row_inv: f32; | |
| {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %} | |
| {% if op == "max" or op == "min" %} | |
| {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %} | |
| {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %} | |
| {% 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) %} | |
| var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u; | |
| loop { | |
| if ({{ svar }} == 0u) { | |
| break; | |
| } | |
| if ({{ idx }} < {{ svar }}) { | |
| {% for a in arrays %} | |
| {{ wgsl_tree_fold_stmt(a, op, idx, svar) }} | |
| {% endfor %} | |
| } | |
| {{ svar }} = {{ svar }} / 2u; | |
| workgroupBarrier(); | |
| }{% endmacro %} | |
| // Reusing partial after this reduction requires a barrier between the read of | |
| // partial[0] and the next write, or the next round can race the prior readers. | |
| fn reduce_sum(value: f32, tid: u32) -> f32 { | |
| partial[tid] = value; | |
| workgroupBarrier(); | |
| {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }} | |
| return partial[0]; | |
| } | |
| @compute @workgroup_size(WG, 1, 1) | |
| fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) { | |
| let row = wg.x + wg.y * params.rowStride; | |
| if (row >= params.rows) { | |
| return; | |
| } | |
| let tid = lid.x; | |
| let base = row * HIDDEN; | |
| var local_sum = 0.0; | |
| for (var d = tid; d < HIDDEN; d = d + WG) { | |
| let value = f32(x[base + d]); | |
| local_sum = local_sum + value; | |
| } | |
| let sum = reduce_sum(local_sum, tid); | |
| if (tid == 0u) { | |
| row_mean = sum / f32(HIDDEN); | |
| } | |
| workgroupBarrier(); | |
| var local_var_sum = 0.0; | |
| for (var d = tid; d < HIDDEN; d = d + WG) { | |
| let diff = f32(x[base + d]) - row_mean; | |
| local_var_sum = local_var_sum + diff * diff; | |
| } | |
| let var_sum = reduce_sum(local_var_sum, tid); | |
| if (tid == 0u) { | |
| let variance = var_sum / f32(HIDDEN); | |
| row_inv = inverseSqrt(variance + EPSILON); | |
| {% if writeMean %} | |
| mean_out[row] = row_mean; | |
| {% endif %} | |
| {% if writeInvStdDev %} | |
| inv_std_out[row] = row_inv; | |
| {% endif %} | |
| } | |
| workgroupBarrier(); | |
| for (var d = tid; d < HIDDEN; d = d + WG) { | |
| let index = base + d; | |
| let normalized = (f32(x[index]) - row_mean) * row_inv; | |
| var value = normalized * f32(scale[{% if scaleNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("scale_offset", scaleShape, xShape, "index") }}{% endif %}]); | |
| {% if hasBias %} | |
| value = value + f32(bias[{% if biasNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("bias_offset", biasShape, xShape, "index") }}{% endif %}]); | |
| {% endif %} | |
| y[index] = {{ scalar }}(value); | |
| } | |
| } | |
| {% endif %} | |