struct MatrixBatch { size: vec4, data: array, } @group(0) @binding(0) var ilogits: MatrixBatch; @group(0) @binding(1) var inp: MatrixBatch; @group(0) @binding(2) var pos: MatrixBatch; @group(0) @binding(3) var r: MatrixBatch; @compute @workgroup_size(1) fn main(@builtin(global_invocation_id) global_id: vec3) { r.size = pos.size; var s = 0f; var li = (global_id.z * ilogits.size.y + global_id.y) * ilogits.size.x; var ri = global_id.z * inp.size.x * inp.size.y + global_id.x; for (var i = 0u; i < ilogits.size.x; i++) { s += ilogits.data[li] * inp.data[ri]; li++; ri += inp.size.x; } var resi = (global_id.z * r.size.y + global_id.y) * r.size.x + global_id.x; r.data[resi] = s + pos.data[resi]; }