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