{% if usesF16 %} enable f16; {% endif %} {% 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; var partial: array; var row_mean: f32; var row_inv: f32; {% set xNumel = namespace(value=1) %} {% for dim in source.xShape %} {% set xNumel.value = xNumel.value * dim %} {% endfor %} {% set scaleNumel = namespace(value=1) %} {% for dim in source.scaleShape %} {% set scaleNumel.value = scaleNumel.value * dim %} {% endfor %} {% if scaleNumel.value != 1 %} {{ offset_fn("scale_offset", source.scaleShape, source.scaleShape | length, source.scaleShape == source.xShape, scaleNumel.value, source.xShape, source.xShape | length, xNumel.value) }} {% endif %} {% if hasBias %} {% set biasNumel = namespace(value=1) %} {% for dim in source.biasShape %} {% set biasNumel.value = biasNumel.value * dim %} {% endfor %} {% if biasNumel.value != 1 %} {{ offset_fn("bias_offset", source.biasShape, source.biasShape | length, source.biasShape == source.xShape, biasNumel.value, source.xShape, source.xShape | length, xNumel.value) }} {% endif %} {% endif %} {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %} {% if op == "max" %} {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]); {%- else %} {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ 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) %} var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u; loop { {% if form == "head" %} {% if breakInline %} if ({{ svar }} == 0u) { break; } {% else %} if ({{ svar }} == 0u) { break; } {% endif %} {% endif %} {% if bodyInline %} if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} } {% else %} if ({{ idx }} < {{ svar }}) { {% for a in arrays %} {{ wgsl_tree_fold_stmt(a, op, idx, svar) }} {% endfor %} } {% endif %} {% if form == "head" %} {% if barrierFirst %} workgroupBarrier(); {{ svar }} = {{ svar }} / 2u; {% else %} {{ svar }} = {{ svar }} / 2u; workgroupBarrier(); {% endif %} {% else %} workgroupBarrier(); if ({{ svar }} == 1u) { break; } {{ svar }} = {{ svar }} / 2u; {% endif %} } {%- 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. {% set trailingBarrier = trailingBarrier is defined and trailingBarrier %} fn reduce_sum(value: f32, tid: u32) -> f32 { partial[tid] = value; workgroupBarrier(); {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }} {% if trailingBarrier %} let total = partial[0]; workgroupBarrier(); return total; {% else %} return partial[0]; {% endif %} } @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", source.scaleShape, source.xShape, "index") }}{% endif %}]); {% if hasBias %} value = value + f32(bias[{% if biasNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("bias_offset", source.biasShape, source.xShape, "index") }}{% endif %}]); {% endif %} y[index] = {{ scalar }}(value); } }