File size: 3,976 Bytes
69e74ac | 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 | {% 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]));
}
}
|