{{ env.wgsl.resourceDeclarations }} const HIDDEN: u32 = {{ hiddenSize }}u; const EPSILON: f32 = {{ epsilon }}; const WG: u32 = {{ workgroupSize }}u; var partial: array; {% if scaleRank > 0 %} const X_RANK: u32 = {{ xRank }}u; const SCALE_RANK: u32 = {{ scaleRank }}u; const X_SHAPE: array = array({% for d in xShape %}{{ d }}u{% if not loop.last %}, {% endif %}{% endfor %}); const SCALE_SHAPE: array = array({% for d in scaleShape %}{{ d }}u{% if not loop.last %}, {% endif %}{% endfor %}); fn x_stride(axis: u32) -> u32 { var stride = 1u; for (var i = axis + 1u; i < X_RANK; i += 1u) { stride *= X_SHAPE[i]; } return stride; } fn scale_stride(axis: u32) -> u32 { var stride = 1u; for (var i = axis + 1u; i < SCALE_RANK; i += 1u) { stride *= SCALE_SHAPE[i]; } return stride; } {% endif %} fn scale_offset({% if scaleRank > 0 %}out_index: u32{% endif %}) -> u32 { {% if scaleRank == 0 %} return 0u; {% else %} var rem = out_index; var offset = 0u; for (var axis = 0u; axis < X_RANK; axis += 1u) { let stride = x_stride(axis); let coord = rem / stride; rem %= stride; let scale_axis = i32(axis) - i32(X_RANK - SCALE_RANK); if (scale_axis >= 0) { let s_axis = u32(scale_axis); if (SCALE_SHAPE[s_axis] != 1u) { offset += coord * scale_stride(s_axis); } } } return offset; {% 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_sq = 0.0; for (var d = tid; d < HIDDEN; d = d + WG) { let value = f32(x[base + d]); local_sq = local_sq + value * value; } let inv = inverseSqrt(reduce_sum(local_sq, tid) / f32(HIDDEN) + EPSILON); {% if writeStats %} if (tid == 0u) { inv_std_out[row] = inv; } {% endif %} for (var d = tid; d < HIDDEN; d = d + WG) { let index = base + d; let value = f32(x[index]) * inv * f32(scale[scale_offset({% if scaleRank > 0 %}index{% endif %})]); y[base + d] = {{ scalar }}(value); } }