File size: 4,742 Bytes
2760d09 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | {% 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<vec2<f32>, WG>;
fn reduce_pair(value: vec2<f32>{% if useSubgroups %}, sg_lane: u32, sg_id: u32, num_sg: u32{% else %}, tid: u32{% endif %}) -> vec2<f32> {
{% if useSubgroups %}
let s = vec2<f32>(subgroupAdd(value.x), subgroupAdd(value.y));
if (num_sg == 1u) {
return s;
}
if (sg_lane == 0u) {
sg_partials[sg_id] = s;
}
workgroupBarrier();
var total = vec2<f32>(0.0, 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;
{% if source.broadcastSkip %}
// skip broadcasts across the batch dim: fold row into [0, skipRows) so every
// batch reuses the same skip row (skipRows == params.rows ⇒ identity).
let skip_base = (row % params.skipRows) * HIDDEN_V;
{% else %}
let skip_base = base;
{% endif %}
let shift = residual_value(base, skip_base{% if source.hasBias %}, 0u{% endif %}).x;
var acc = vec2<f32>(0.0, 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 %});
let d = v - vec4<f32>(shift);
acc.x = acc.x + d.x + d.y + d.z + d.w;
acc.y = acc.y + dot(d, d);
}
let totals = reduce_pair(acc{% if useSubgroups %}, sg_lane, sg_id, num_sg{% else %}, tid{% endif %});
let mean_d = totals.x / f32(HIDDEN);
let variance = max(totals.y / f32(HIDDEN) - mean_d * mean_d, 0.0);
let row_inv = inverseSqrt(variance + params.epsilon);
let row_mean = shift + mean_d;
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 %}
var value = (residual - vec4<f32>(row_mean)) * row_inv * vec4<f32>(gamma[i]);
{% if source.hasBeta %}
value = value + vec4<f32>(beta[i]);
{% endif %}
output[idx] = {{ source.vecType }}(value);
}
}
|