ai.onnx.LayerNormalization / build /webgpu /layer-normalization.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 2e7068faf55e
49b0523 verified
Raw
History Blame
6.15 kB
{% if usesF16 %}
enable f16;
{% endif %}
{% macro offset_fn(fn_name, opShape, opRank, op_same, op_numel, outShape, outRank, out_numel) %}
fn {{ fn_name }}({% if out_numel != 0 and op_numel != 1 %}out_index: u32{% endif %}) -> u32 {
{% if out_numel == 0 %}
return 0u;
{% elif op_numel == 1 %}
return 0u;
{% elif op_same %}
return out_index;
{% else %}
var offset = 0u;
{% for axis in range(outRank) %}
{% set op_axis = axis - (outRank - opRank) %}
{% if op_axis >= 0 and opShape[op_axis] != 1 %}
{% set c_stride = namespace(value=1) %}
{% for j in range(axis + 1, outRank) %}
{% set c_stride.value = c_stride.value * outShape[j] %}
{% endfor %}
{% set op_stride = namespace(value=1) %}
{% for j in range(op_axis + 1, opRank) %}
{% set op_stride.value = op_stride.value * opShape[j] %}
{% endfor %}
{% if c_stride.value == 1 %}
let coord{{ axis }} = out_index % {{ outShape[axis] }}u;
{% else %}
let coord{{ axis }} = (out_index / {{ c_stride.value }}u) % {{ outShape[axis] }}u;
{% endif %}
{% if op_stride.value == 1 %}
offset = offset + coord{{ axis }};
{% else %}
offset = offset + coord{{ axis }} * {{ op_stride.value }}u;
{% endif %}
{% endif %}
{% endfor %}
return offset;
{% endif %}
}
{%- endmacro %}{% macro broadcast_offset_call(fn_name, opShape, outShape, out_index) %}
{% set op_numel = namespace(value=1) %}
{% for d in opShape %}{% set op_numel.value = op_numel.value * d %}{% endfor %}
{% set out_numel = namespace(value=1) %}
{% for d in outShape %}{% set out_numel.value = out_numel.value * d %}{% endfor %}
{{ fn_name }}({% if out_numel.value != 0 and op_numel.value != 1 %}{{ out_index }}{% endif %})
{%- endmacro %}
{{ env.wgsl.resourceDeclarations }}
const HIDDEN: u32 = {{ hiddenSize }}u;
const EPSILON: f32 = {{ epsilon }};
const WG: u32 = {{ workgroupSize }}u;
var<workgroup> partial: array<f32, WG>;
var<workgroup> row_mean: f32;
var<workgroup> row_inv: f32;
{% set xNumel = namespace(value=1) %}
{% for dim in source.xShape %}
{% set xNumel.value = xNumel.value * dim %}
{% endfor %}
{% set scaleNumel = namespace(value=1) %}
{% for dim in source.scaleShape %}
{% set scaleNumel.value = scaleNumel.value * dim %}
{% endfor %}
{% if scaleNumel.value != 1 %}
{{ offset_fn("scale_offset", source.scaleShape, source.scaleShape | length, source.scaleShape == source.xShape, scaleNumel.value, source.xShape, source.xShape | length, xNumel.value) }}
{% endif %}
{% if hasBias %}
{% set biasNumel = namespace(value=1) %}
{% for dim in source.biasShape %}
{% set biasNumel.value = biasNumel.value * dim %}
{% endfor %}
{% if biasNumel.value != 1 %}
{{ offset_fn("bias_offset", source.biasShape, source.biasShape | length, source.biasShape == source.xShape, biasNumel.value, source.xShape, source.xShape | length, xNumel.value) }}
{% endif %}
{% 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_sum = 0.0;
for (var d = tid; d < HIDDEN; d = d + WG) {
let value = f32(x[base + d]);
local_sum = local_sum + value;
}
let sum = reduce_sum(local_sum, tid);
if (tid == 0u) {
row_mean = sum / f32(HIDDEN);
}
workgroupBarrier();
var local_var_sum = 0.0;
for (var d = tid; d < HIDDEN; d = d + WG) {
let diff = f32(x[base + d]) - row_mean;
local_var_sum = local_var_sum + diff * diff;
}
let var_sum = reduce_sum(local_var_sum, tid);
if (tid == 0u) {
let variance = var_sum / f32(HIDDEN);
row_inv = inverseSqrt(variance + EPSILON);
{% if writeMean %}
mean_out[row] = row_mean;
{% endif %}
{% if writeInvStdDev %}
inv_std_out[row] = row_inv;
{% endif %}
}
workgroupBarrier();
for (var d = tid; d < HIDDEN; d = d + WG) {
let index = base + d;
let normalized = (f32(x[index]) - row_mean) * row_inv;
var value = normalized * f32(scale[{% if scaleNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("scale_offset", source.scaleShape, source.xShape, "index") }}{% endif %}]);
{% if hasBias %}
value = value + f32(bias[{% if biasNumel.value == 1 %}0u{% else %}{{ broadcast_offset_call("bias_offset", source.biasShape, source.xShape, "index") }}{% endif %}]);
{% endif %}
y[index] = {{ scalar }}(value);
}
}