File size: 4,221 Bytes
69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 69e74ac f1a8138 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 131 132 133 134 135 136 137 | {% 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 broadcastSkip = broadcastSkip is defined and broadcastSkip %}
{% set useSubgroups = useSubgroups %}
{% if usesF16Spec %}
enable f16;
{% endif %}
{% if useSubgroups %}
enable subgroups;
{% endif %}
{{ env.wgsl.resourceDeclarations }}
const HIDDEN: u32 = {{ hidden }}u;
const HIDDEN_V: u32 = {{ hiddenVec }}u;
const WG: u32 = {{ 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 hasBias %}, vi: u32{% endif %}) -> vec4<f32> {
var value = vec4<f32>(input[idx]) + vec4<f32>(skip[skip_idx]);
{% if 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 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 %}
var acc = 0.0;
for (var i = tid; i < HIDDEN_V; i = i + WG) {
let v = residual_value(base + i, skip_base + i{% if 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 hasBias %}, i{% endif %});
{% if writeResidualSum %}
input_skip_bias_sum[idx] = {{ vecType }}(residual);
{% endif %}
output[idx] = {{ vecType }}(residual * row_inv * vec4<f32>(gamma[i]));
}
}
|