// Split-K normalize pass. Each workgroup (row = wg.x, split index = wg.z) // folds the SPLIT per-row partial sums of squares into the RMS scale, then // normalizes its HIDDEN/SPLIT slice. SPLIT is small (<=64), so the serial fold // avoids a third combine pass. Scale offsets follow the suffix-axis broadcast // contract. {% if usesF16 %} enable f16; {% endif %} {{ env.wgsl.resourceDeclarations }} const HIDDEN: u32 = {{ hiddenSize }}u; const EPSILON: f32 = {{ epsilon }}; const WG: u32 = {{ workgroupSize }}u; const SPLIT: u32 = {{ split }}u; {% if source.scaleRank > 0 %} const X_RANK: u32 = {{ source.xRank }}u; const SCALE_RANK: u32 = {{ source.scaleRank }}u; const X_SHAPE: array = array({% for d in source.xShape %}{{ d }}u{% if not loop.last %}, {% endif %}{% endfor %}); const SCALE_SHAPE: array = array({% for d in source.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 source.scaleRank > 0 %}out_index: u32{% endif %}) -> u32 { {% if source.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 %} } @compute @workgroup_size(WG, 1, 1) fn main(@builtin(workgroup_id) wg: vec3, @builtin(local_invocation_id) lid: vec3) { let row = wg.x + wg.y * params.rowStride; if (row >= params.rows) { return; } let k = wg.z; let tid = lid.x; var total = 0.0; for (var i = 0u; i < SPLIT; i = i + 1u) { total = total + partials[row * SPLIT + i]; } let inv = inverseSqrt(total / f32(HIDDEN) + EPSILON); let chunk = (HIDDEN + SPLIT - 1u) / SPLIT; 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; // Preserve the ONNX stage boundary: round Normalized to X's dtype before // the affine scale is applied. let normalized = {{ xElement }}(f32(x[index]) * inv); let value = f32(normalized) * f32(scale[scale_offset({% if source.scaleRank > 0 %}index{% endif %})]); y[index] = {{ scalar }}(value); d = d + WG; } }