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 s = 0f; for (var x = 0u; x < vc; x++) { s += dot(m.data[ro + x], dldm.data[ro + x]); } grad.data[ro + global_id.x] = m.data[ro + global_id.x] * (dldm.data[ro + global_id.x] - vec4(s)); }