{% 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(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, b: vec2) -> vec2 { 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(hi, fma(-1.0, fma(-1.0, s, hi), e)); } fn pair_mul(a: vec2, b: vec2) -> vec2 { let p = fma(a.x, b.x, 0.0); let error = fma(a.x, b.x, -p); return pair_add(vec2(p, 0.0), vec2(error + (a.x * b.y + b.x * a.y), 0.0)); } fn pair_div(a: vec2, b: f32) -> vec2 { let q = a.x / b; let residual = pair_add(a, -pair_mul(vec2(q, 0.0), vec2(b, 0.0))); return pair_add(vec2(q, 0.0), vec2((residual.x + residual.y) / b, 0.0)); } fn pair_inverse_sqrt(a: vec2) -> vec2 { let r = inverseSqrt(a.x); let rr = pair_mul(vec2(r, 0.0), vec2(r, 0.0)); let residual = pair_add(vec2(1.0, 0.0), -pair_mul(a, rr)); return pair_add(vec2(r, 0.0), vec2((0.5 * r) * (residual.x + residual.y), 0.0)); } fn half_normalized(value: vec2) -> 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 partial: array, WG>; fn reduce_pair(value: vec2, tid: u32) -> vec2 { 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(0.0); for (var i = tid; i < HIDDEN; i += WG) { local_sum = pair_add(local_sum, vec2(load_value(base + i), 0.0)); } let mean = pair_div(reduce_pair(local_sum, {{ reduceArgs }}), f32(HIDDEN)); var local_square = vec2(0.0); for (var i = tid; i < HIDDEN; i += WG) { let centered = pair_add(vec2(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(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(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, @builtin(local_invocation_id) lid: vec3) { let row = wg.x + wg.y * params.rowStride; normalize_half_row(row, lid.x); } {% else %} var partial: array; var row_mean: f32; var 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, @builtin(local_invocation_id) lid: vec3) { 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 %}