| {{ env.wgsl.resourceDeclarations }} |
| |
| const HIDDEN: u32 = {{ hiddenSize }}u; |
| const EPSILON: f32 = {{ epsilon }}; |
| const WG: u32 = {{ workgroupSize }}u; |
| |
| var<workgroup> partial: array<f32, WG>; |
| |
| {% if scaleRank > 0 %} |
| const X_RANK: u32 = {{ xRank }}u; |
| const SCALE_RANK: u32 = {{ scaleRank }}u; |
| const X_SHAPE: array<u32, {{ xRank }}> = array<u32, {{ xRank }}>({% for d in xShape %}{{ d }}u{% if not loop.last %}, {% endif %}{% endfor %}); |
| const SCALE_SHAPE: array<u32, {{ scaleRank }}> = array<u32, {{ scaleRank }}>({% 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<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_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); |
| } |
| } |
| |