struct Scalar { size: vec4, data: array, } struct MatrixBatch { size: vec4, data: array, } @group(0) @binding(0) var ds: Scalar; @group(0) @binding(1) var wa: MatrixBatch; @group(0) @binding(2) var v: MatrixBatch; @group(0) @binding(3) var otdldo: MatrixBatch; @group(0) @binding(4) var dlds: MatrixBatch; @compute @workgroup_size(1) fn main(@builtin(global_invocation_id) global_id: vec3) { dlds.size = wa.size; var s = 0f; var li = (global_id.z * v.size.y + global_id.x) * v.size.x; var ri = global_id.z * otdldo.size.x * otdldo.size.y + global_id.y; for (var i = 0u; i < v.size.x; i++) { var cs = 0f; var lj = global_id.z * wa.size.y * wa.size.x + global_id.y * wa.size.x; var mj = global_id.z * v.size.x * v.size.y + i; for (var j = 0u; j < wa.size.x; j++) { cs += wa.data[lj] * v.data[mj]; lj++; mj += v.size.x; } s += (v.data[li] - cs) * otdldo.data[ri]; li++; ri += otdldo.size.x; } let resi = (global_id.z * dlds.size.y + global_id.y) * dlds.size.x + global_id.x; dlds.data[resi] = ds.data[0] * wa.data[resi] * s; }