struct MatrixBatch { size: vec4, data: array>, } @group(0) @binding(0) var b: MatrixBatch; @group(0) @binding(1) var result: MatrixBatch; @compute @workgroup_size(1) fn main(@builtin(global_invocation_id) global_id: vec3) { result.size = b.size; var vc = result.size.x / 4; var resi = (global_id.z * result.size.y + global_id.y) * vc + global_id.x; var mv = b.data[resi]; var bi = (global_id.z * b.size.y + global_id.y) * vc; for (var i = 0u; i < vc; i++) { mv = max(mv, b.data[bi]); bi++; } var ms = max(max(mv.x, mv.y), max(mv.z, mv.w)); result.data[resi] = exp(b.data[resi] - vec4(ms)); }