struct MatrixBatch { size: vec4, data: array, } @group(0) @binding(0) var ilogits: MatrixBatch; @group(0) @binding(0) var tilogits: MatrixBatch; @group(0) @binding(1) var dldsq: MatrixBatch; @group(0) @binding(2) var dldsk: MatrixBatch; @group(0) @binding(3) var totdldowa: MatrixBatch; @group(0) @binding(4) var tqry: MatrixBatch; @group(0) @binding(5) var tkey: MatrixBatch; @group(0) @binding(6) var tval: MatrixBatch; @group(0) @binding(7) var ginp: MatrixBatch; @compute @workgroup_size(1) fn main(@builtin(global_invocation_id) global_id: vec3) { ginp.size.x = tval.size.x; ginp.size.y = ilogits.size.x; ginp.size.z = ilogits.size.z; var ts = 0f; var li = global_id.z * ilogits.size.y * ilogits.size.x + global_id.y; for (var i = 0u; i < ilogits.size.y; i++) { var s = 0f; var lj = global_id.z * totdldowa.size.x * totdldowa.size.y + i * totdldowa.size.x; var rj = global_id.z * tval.size.x * tval.size.y + global_id.x; for (var j = 0u; j < totdldowa.size.x; j++) { s += totdldowa.data[lj] * tval.data[rj] + dldsq.data[lj] * tkey.data[rj] + dldsk.data[lj] * tqry.data[rj]; lj++; rj += tval.size.x; } ts += ilogits.data[li] * s; li += ilogits.size.x; } ginp.data[global_id.z * ginp.size.y * ginp.size.x + global_id.y * ginp.size.x + global_id.x] = ts; }