ai.onnx.RMSNormalization / build /webgpu /rms-normalization-splitk-partials.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 91d990483a17
f3f43cb verified
Raw
History Blame
2.77 kB
{% 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 %}
/* Split-K partial sum-of-squares for tensors with few rows and a large hidden
dimension. A workgroup-per-row kernel exposes too little parallelism in this
regime, so this pass splits each row across SPLIT workgroups
(row = wg.x, split index = wg.z). Each workgroup accumulates a partial
sum-of-squares over its HIDDEN/SPLIT slice and writes one partial to scratch.
The normalize pass folds the SPLIT partials per row. Split-K reassociates the
f32 sum, so this route is not bit-identical to the unsplit reduction. */
{{ env.wgsl.resourceDeclarations }}
const HIDDEN: u32 = {{ hiddenSize }}u;
const WG: u32 = {{ workgroupSize }}u;
const SPLIT: u32 = {{ split }}u;
var<workgroup> red: array<f32, WG>;
@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 k = wg.z;
let tid = lid.x;
let chunk = (HIDDEN + SPLIT - 1u) / SPLIT;
let start = k * chunk;
var end = start + chunk;
if (end > HIDDEN) { end = HIDDEN; }
let base = row * HIDDEN;
var acc = 0.0;
var d = start + tid;
loop {
if (d >= end) { break; }
let v = f32(x[base + d]);
acc = acc + v * v;
d = d + WG;
}
red[tid] = acc;
workgroupBarrier();
{{ wgsl_tree_fold(["red"], idx="tid", wg="WG", typed=true, form="head", breakInline=true, bodyInline=true) }}
if (tid == 0u) {
partials[row * SPLIT + k] = red[0];
}
}