| {% 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<workgroup> partial: array<f32, WG>; |
| var<workgroup> row_mean: f32; |
| var<workgroup> 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<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", 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); |
| } |
| } |
| |