struct Scalar { size: vec4, data: array, } struct Matrix { size: vec4, data: array>, } struct MatrixBatch { size: vec4, data: array>, } @group(0) @binding(0) var ebs: MatrixBatch; @group(0) @binding(1) var ows: Matrix; @group(0) @binding(2) var tmp: Scalar; @group(0) @binding(3) var dts: MatrixBatch; @compute @workgroup_size(1) fn main(@builtin(global_invocation_id) global_id: vec3) { dts.size.x = ows.size.y; dts.size.y = ebs.size.y; dts.size.z = ebs.size.z; let vc = ebs.size.x / 4; let eo = (global_id.z * ebs.size.y + global_id.y) * vc; let oo = 4 * global_id.x * vc; var s = vec4(0f); for (var i = 0u; i < vc; i++) { let eb = ebs.data[eo + i]; s.x += dot(eb, ows.data[oo + 0 * vc + i]); s.y += dot(eb, ows.data[oo + 1 * vc + i]); s.z += dot(eb, ows.data[oo + 2 * vc + i]); s.w += dot(eb, ows.data[oo + 3 * vc + i]); } dts.data[(global_id.z * dts.size.y + global_id.y) * dts.size.x / 4 + global_id.x] = tmp.data[0] * s; }