ai.onnx.SimplifiedLayerNormalization / build /webgpu /rms-normalization.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 91d990483a17
8e0c6a5 verified
Raw
History Blame
4.07 kB
{{ env.wgsl.resourceDeclarations }}
const HIDDEN: u32 = {{ hiddenSize }}u;
const EPSILON: f32 = {{ epsilon }};
const WG: u32 = {{ workgroupSize }}u;
var<workgroup> partial: array<f32, WG>;
{% if scaleRank > 0 %}
const X_RANK: u32 = {{ xRank }}u;
const SCALE_RANK: u32 = {{ scaleRank }}u;
const X_SHAPE: array<u32, {{ xRank }}> = array<u32, {{ xRank }}>({% for d in xShape %}{{ d }}u{% if not loop.last %}, {% endif %}{% endfor %});
const SCALE_SHAPE: array<u32, {{ scaleRank }}> = array<u32, {{ scaleRank }}>({% for d in scaleShape %}{{ d }}u{% if not loop.last %}, {% endif %}{% endfor %});
fn x_stride(axis: u32) -> u32 {
var stride = 1u;
for (var i = axis + 1u; i < X_RANK; i += 1u) {
stride *= X_SHAPE[i];
}
return stride;
}
fn scale_stride(axis: u32) -> u32 {
var stride = 1u;
for (var i = axis + 1u; i < SCALE_RANK; i += 1u) {
stride *= SCALE_SHAPE[i];
}
return stride;
}
{% endif %}
fn scale_offset({% if scaleRank > 0 %}out_index: u32{% endif %}) -> u32 {
{% if scaleRank == 0 %}
return 0u;
{% else %}
var rem = out_index;
var offset = 0u;
for (var axis = 0u; axis < X_RANK; axis += 1u) {
let stride = x_stride(axis);
let coord = rem / stride;
rem %= stride;
let scale_axis = i32(axis) - i32(X_RANK - SCALE_RANK);
if (scale_axis >= 0) {
let s_axis = u32(scale_axis);
if (SCALE_SHAPE[s_axis] != 1u) {
offset += coord * scale_stride(s_axis);
}
}
}
return offset;
{% endif %}
}
{% 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 %}
// Reusing partial after this reduction requires a barrier between the read of
// partial[0] and the next write, or the next round can race the prior readers.
{% set trailingBarrier = trailingBarrier is defined and trailingBarrier %}
fn reduce_sum(value: f32, tid: u32) -> f32 {
partial[tid] = value;
workgroupBarrier();
{{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
{% if trailingBarrier %}
let total = partial[0];
workgroupBarrier();
return total;
{% else %}
return partial[0];
{% endif %}
}
@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 tid = lid.x;
let base = row * HIDDEN;
var local_sq = 0.0;
for (var d = tid; d < HIDDEN; d = d + WG) {
let value = f32(x[base + d]);
local_sq = local_sq + value * value;
}
let inv = inverseSqrt(reduce_sum(local_sq, tid) / f32(HIDDEN) + EPSILON);
{% if writeStats %}
if (tid == 0u) {
inv_std_out[row] = inv;
}
{% endif %}
for (var d = tid; d < HIDDEN; d = d + WG) {
let index = base + d;
let value = f32(x[index]) * inv * f32(scale[scale_offset({% if scaleRank > 0 %}index{% endif %})]);
y[base + d] = {{ scalar }}(value);
}
}