struct MatrixBatch { size: vec4, data: array>, } @group(0) @binding(0) var m: MatrixBatch; @group(0) @binding(1) var dldm: MatrixBatch; @group(0) @binding(2) var grad: MatrixBatch; @compute @workgroup_size(1) fn main(@builtin(global_invocation_id) global_id: vec3) { grad.size = m.size; let vc = m.size.x / 4; let ro = (global_id.z * m.size.y + global_id.y) * vc; var sum = vec4(0f); var sumsq = vec4(0f); for (var i = 0u; i < vc; i++) { let v = m.data[ro + i]; sum += v; sumsq += v * v; } let mean = (sum.x + sum.y + sum.z + sum.w) / f32(m.size.x); let variance = (sumsq.x + sumsq.y + sumsq.z + sumsq.w) / f32(m.size.x) - mean * mean; let inv_std = inverseSqrt((sumsq.x + sumsq.y + sumsq.z + sumsq.w) / f32(m.size.x) - mean * mean + 1e-5); // Recompute x̂ var sum_grad = vec4(0.0); var sum_grad_xhat = vec4(0.0); for (var i = 0u; i < vc; i++) { let v = m.data[ro + i]; let g = dldm.data[ro + i]; let xhat = (v - vec4(mean)) * vec4(inv_std); sum_grad += g; sum_grad_xhat += g * xhat; } let mg = (sum_grad.x + sum_grad.y + sum_grad.z + sum_grad.w) / f32(m.size.x); let mgxhat = (sum_grad_xhat.x + sum_grad_xhat.y + sum_grad_xhat.z + sum_grad_xhat.w) / f32(m.size.x); for (var i = 0u; i < vc; i++) { let v = m.data[ro + i]; let g = dldm.data[ro + i]; let xhat = (v - vec4(mean)) * vec4(inv_std); let dx = vec4(inv_std) * (g - vec4(mg) - xhat * vec4(mgxhat)); grad.data[ro + i] = dx; } }