| {% 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 %}{% set useSubgroups = source.useSubgroups %} |
| {% if source.usesF16 %} |
| enable f16; |
| {% endif %} |
| {% if useSubgroups %} |
| enable subgroups; |
| {% endif %} |
| {{ env.wgsl.resourceDeclarations }} |
| |
| const HIDDEN: u32 = {{ source.hidden }}u; |
| const HIDDEN_V: u32 = {{ source.hiddenVec }}u; |
| const WG: u32 = {{ source.wg }}u; |
| |
| var<workgroup> sg_partials: array<f32, WG>; |
| |
| fn reduce_scalar(value: f32{% if useSubgroups %}, sg_lane: u32, sg_id: u32, num_sg: u32{% else %}, tid: u32{% endif %}) -> f32 { |
| {% if useSubgroups %} |
| let s = subgroupAdd(value); |
| if (num_sg == 1u) { |
| return s; |
| } |
| if (sg_lane == 0u) { |
| sg_partials[sg_id] = s; |
| } |
| workgroupBarrier(); |
| var total = 0.0; |
| for (var i = 0u; i < num_sg; i = i + 1u) { |
| total = total + sg_partials[i]; |
| } |
| return total; |
| {% else %} |
| // No-subgroup tier: workgroup barrier tree-reduction (WG is a power of two). |
| sg_partials[tid] = value; |
| workgroupBarrier(); |
| {{ wgsl_tree_fold(["sg_partials"], idx="tid", wg="WG", form="head", breakInline=true) }} |
| return sg_partials[0]; |
| {% endif %} |
| } |
| |
| // 4 contiguous residual elements (input[idx] + skip[skip_idx] [+ bias]) at vec4 |
| // index `vi`. skip_idx == idx for the normal (non-broadcast) path; for a skip |
| // that broadcasts across the leading/batch dim uses a folded index. |
| fn residual_value(idx: u32, skip_idx: u32{% if source.hasBias %}, vi: u32{% endif %}) -> vec4<f32> { |
| var value = vec4<f32>(input[idx]) + vec4<f32>(skip[skip_idx]); |
| {% if source.hasBias %} |
| value = value + vec4<f32>(bias[vi]); |
| {% endif %} |
| return value; |
| } |
| |
| @compute @workgroup_size(WG, 1, 1) |
| fn main( |
| @builtin(workgroup_id) wg_id: vec3<u32>, |
| @builtin(local_invocation_id) lid: vec3<u32>{% if useSubgroups %}, |
| @builtin(subgroup_invocation_id) sg_lane: u32, |
| @builtin(subgroup_id) sg_id: u32, |
| @builtin(num_subgroups) num_sg: u32{% endif %} |
| ) { |
| let row = wg_id.x + wg_id.y * params.rowStride; |
| if (row >= params.rows) { |
| return; |
| } |
| let tid = lid.x; |
| let base = row * HIDDEN_V; |
| let skip_base = base; |
| |
| |
| var acc = 0.0; |
| for (var i = tid; i < HIDDEN_V; i = i + WG) { |
| let v = residual_value(base + i, skip_base + i{% if source.hasBias %}, i{% endif %}); |
| acc = acc + dot(v, v); |
| } |
| |
| let total = reduce_scalar(acc{% if useSubgroups %}, sg_lane, sg_id, num_sg{% else %}, tid{% endif %}); |
| let row_inv = inverseSqrt(total / f32(HIDDEN) + params.epsilon); |
| |
| for (var i = tid; i < HIDDEN_V; i = i + WG) { |
| let idx = base + i; |
| let residual = residual_value(idx, skip_base + i{% if source.hasBias %}, i{% endif %}); |
| {% if source.writeResidualSum %} |
| input_skip_bias_sum[idx] = {{ source.vecType }}(residual); |
| {% endif %} |
| output[idx] = {{ source.vecType }}(residual * row_inv * vec4<f32>(gamma[i])); |
| } |
| } |
| |