Download build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja from webgpu-kernels/ai.onnx.SimplifiedLayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 2.92 kB
-
https://huggingface.co/kernels/webgpu-kernels/ai.onnx.SimplifiedLayerNormalization/resolve/v1/build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/ai.onnx.SimplifiedLayerNormalization@v1/build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja
-
curl -L -o rms-normalization-splitk-normalize.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/ai.onnx.SimplifiedLayerNormalization/resolve/v1/build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja
2.92 kB
| // Split-K normalize pass. One lane folds the per-row partials in their original | |
| // order and shares the inverse RMS with its workgroup. Output tiles can outnumber | |
| // the reduction splits: their independent work does not need additional scratch. | |
| // Scale offsets follow the suffix-axis broadcast contract. | |
| {{ env.wgsl.resourceDeclarations }} | |
| const HIDDEN: u32 = {{ hiddenSize }}u; | |
| const EPSILON: f32 = {{ epsilon }}; | |
| const WG: u32 = {{ workgroupSize }}u; | |
| const SPLIT: u32 = {{ split }}u; | |
| {% 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 %} | |
| } | |
| var<workgroup> shared_inv: f32; | |
| @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; | |
| if (tid == 0u) { | |
| var total = 0.0; | |
| for (var i = 0u; i < SPLIT; i = i + 1u) { | |
| total = total + partials[row * SPLIT + i]; | |
| } | |
| shared_inv = inverseSqrt(total / f32(HIDDEN) + EPSILON); | |
| } | |
| workgroupBarrier(); | |
| let inv = shared_inv; | |
| {% if writeStats %} | |
| // Every split workgroup folds the same partials, so one designated | |
| // workgroup writes the row statistic. | |
| if (k == 0u && tid == 0u) { | |
| inv_std_out[row] = inv; | |
| } | |
| {% endif %} | |
| let chunk = {{ normalizeChunk }}u; | |
| let start = k * chunk; | |
| var end = start + chunk; | |
| if (end > HIDDEN) { end = HIDDEN; } | |
| let base = row * HIDDEN; | |
| var d = start + tid; | |
| loop { | |
| if (d >= end) { break; } | |
| let index = base + d; | |
| let value = f32(x[index]) * inv * f32(scale[scale_offset({% if scaleRank > 0 %}index{% endif %})]); | |
| y[index] = {{ scalar }}(value); | |
| d = d + WG; | |
| } | |
| } | |