Initial commit

This commit is contained in:
2024-07-27 14:39:08 +00:00
commit def3a479be
55 changed files with 315576 additions and 0 deletions
+37
View File
@@ -0,0 +1,37 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
struct FVectorBatch {
size: vec4<u32>,
data: array<f32>,
}
struct UVectorBatch {
size: vec4<u32>,
data: array<u32>,
}
@group(0) @binding(0)
var<storage, read> ps: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: UVectorBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.x = ps.size.y;
result.size.z = ps.size.z;
var mi = 0u;
var mp = 0f;
var pi = global_id.z * ps.size.y * ps.size.x + global_id.x * ps.size.x;
for (var i = 0u; i < ps.size.x; i++) {
if (ps.data[pi] > mp) {
mi = i;
mp = ps.data[pi];
}
pi++;
}
result.data[global_id.z * result.size.x + global_id.x] = mi;
}
+33
View File
@@ -0,0 +1,33 @@
struct PackedMatrixBatch {
size: vec4<u32>,
data: array<vec4<f32>>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> e: PackedMatrixBatch;
@group(0) @binding(1)
var<storage, read> tw: PackedMatrixBatch;
@group(0) @binding(2)
var<storage, read_write> aw: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
aw.size.x = tw.size.y;
aw.size.y = e.size.y;
aw.size.z = e.size.z;
var s = 0f;
let vc = e.size.x / 4;
let lo = (global_id.z * e.size.y + global_id.y) * vc;
let ro = (global_id.z * tw.size.y + global_id.x) * vc;
for (var i = 0u; i < vc; i++) {
s += dot(e.data[lo + i], tw.data[ro + i]);
}
aw.data[(global_id.z * aw.size.y + global_id.y) * aw.size.x + global_id.x] = s;
}
+32
View File
@@ -0,0 +1,32 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
struct Matrix {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> m: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: Matrix;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = m.size;
result.size.z = 1;
var vc = result.size.x / 4;
var resi = global_id.y * vc + global_id.x;
result.data[resi] = vec4(0f);
var mi = resi;
for (var z = 0u; z < m.size.z; z++) {
result.data[resi] += m.data[mi] / vec4(f32(m.size.z));
mi += m.size.y * vc;
}
}
+18
View File
@@ -0,0 +1,18 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(3)
var<storage, read_write> m: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
for (var z = 0u; z < m.size.z; z++) {
for (var y = 0u; y < m.size.y; y++) {
for (var x = 0u; x < m.size.x; x++) {
m.data[z * m.size.y * m.size.x + y * m.size.x + x] = 0f;
}
}
}
}
+26
View File
@@ -0,0 +1,26 @@
alias number = f32;
struct ScalarBatch {
size: vec4<u32>,
data: array<number>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: ScalarBatch;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = r_or_result.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[global_id.z] * r_or_result.data[resi];
}
+21
View File
@@ -0,0 +1,21 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = l.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[resi] / r_or_result.data[resi];
}
+21
View File
@@ -0,0 +1,21 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = l.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[resi] * r_or_result.data[resi];
}
+26
View File
@@ -0,0 +1,26 @@
alias number = f32;
struct Scalar {
size: vec4<u32>,
data: array<number>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: Scalar;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = r_or_result.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[0] * r_or_result.data[resi];
}
+21
View File
@@ -0,0 +1,21 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = l.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[resi] - r_or_result.data[resi];
}
+21
View File
@@ -0,0 +1,21 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = l.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[resi] + r_or_result.data[resi];
}
+34
View File
@@ -0,0 +1,34 @@
alias number = f32;
struct Scalar {
size: vec4<u32>,
data: array<number>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<number>,
}
struct UVectorBatch {
size: vec4<u32>,
data: array<u32>,
}
@group(0) @binding(0)
var<storage, read> m: MatrixBatch;
@group(0) @binding(1)
var<storage, read> v: UVectorBatch;
@group(0) @binding(2)
var<storage, read> lr: Scalar;
@group(0) @binding(3)
var<storage, read_write> grs: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
grs.size = m.size;
var vi = global_id.z * v.size.x + global_id.y;
var gi = vi * grs.size.x + global_id.x;
grs.data[gi] = lr.data[0] * (m.data[gi] - select(0f, 1f, v.data[vi] == global_id.x));
}
+48
View File
@@ -0,0 +1,48 @@
struct Scalar {
size: vec4<u32>,
data: array<f32>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> ds: Scalar;
@group(0) @binding(1)
var<storage, read> wa: MatrixBatch;
@group(0) @binding(2)
var<storage, read> v: MatrixBatch;
@group(0) @binding(3)
var<storage, read> otdldo: MatrixBatch;
@group(0) @binding(4)
var<storage, read_write> dlds: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
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;
}
+43
View File
@@ -0,0 +1,43 @@
struct Scalar {
size: vec4<u32>,
data: array<f32>,
}
struct Matrix {
size: vec4<u32>,
data: array<vec4<f32>>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<f32>>,
}
@group(0) @binding(0)
var<storage, read> ebs: MatrixBatch;
@group(0) @binding(1)
var<storage, read> ows: Matrix;
@group(0) @binding(2)
var<storage, read> tmp: Scalar;
@group(0) @binding(3)
var<storage, read_write> dts: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
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;
}
+30
View File
@@ -0,0 +1,30 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> ilogits: MatrixBatch;
@group(0) @binding(1)
var<storage, read> inp: MatrixBatch;
@group(0) @binding(2)
var<storage, read> pos: MatrixBatch;
@group(0) @binding(3)
var<storage, read_write> r: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
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];
}
+47
View File
@@ -0,0 +1,47 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> ilogits: MatrixBatch;
@group(0) @binding(0)
var<storage, read> tilogits: MatrixBatch;
@group(0) @binding(1)
var<storage, read> dldsq: MatrixBatch;
@group(0) @binding(2)
var<storage, read> dldsk: MatrixBatch;
@group(0) @binding(3)
var<storage, read> totdldowa: MatrixBatch;
@group(0) @binding(4)
var<storage, read> tqry: MatrixBatch;
@group(0) @binding(5)
var<storage, read> tkey: MatrixBatch;
@group(0) @binding(6)
var<storage, read> tval: MatrixBatch;
@group(0) @binding(7)
var<storage, read_write> ginp: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
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;
}
+28
View File
@@ -0,0 +1,28 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> e: MatrixBatch;
@group(0) @binding(1)
var<storage, read> dlds: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> gtattw: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
gtattw.size.z = e.size.z;
gtattw.size.y = e.size.x;
gtattw.size.x = e.size.x;
var s = 0f;
var li = global_id.z * e.size.x * e.size.y + global_id.x;
var ri = global_id.z * dlds.size.x * dlds.size.y + global_id.y;
for (var i = 0u; i < e.size.y; i++) {
s += e.data[li] * dlds.data[ri];
li += e.size.x;
ri += dlds.size.x;
}
gtattw.data[global_id.z * gtattw.size.y * gtattw.size.x + global_id.y * gtattw.size.x + global_id.x] = s;
}
+28
View File
@@ -0,0 +1,28 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> a: MatrixBatch;
@group(0) @binding(1)
var<storage, read> dldo: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> gtout: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
gtout.size = a.size;
gtout.size.y = dldo.size.x;
var s = 0f;
var ai = global_id.z * a.size.y * a.size.x + global_id.x;
var di = global_id.z * dldo.size.x * dldo.size.y + global_id.y;
for (var i = 0u; i < a.size.y; i++) {
s += a.data[ai] * dldo.data[di];
ai += a.size.x;
di += dldo.size.x;
}
gtout.data[global_id.z * gtout.size.y * gtout.size.x + global_id.y * gtout.size.x + global_id.x] = s;
}
+19
View File
@@ -0,0 +1,19 @@
struct UVectorBatch {
size: vec4<u32>,
data: array<u32>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> v: UVectorBatch;
@group(0) @binding(1)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.data[global_id.z * result.size.y * result.size.x + global_id.y * result.size.x + global_id.x] = select(0f, 1f, global_id.x == v.data[global_id.z * v.size.x + global_id.y]);
}
+37
View File
@@ -0,0 +1,37 @@
alias number = f32;
struct VectorBatch {
size: vec4<u32>,
data: array<u32>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> v: VectorBatch;
@group(0) @binding(1)
var<storage, read> m: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = m.size;
var vc = result.size.x / 4;
var resi = (global_id.z * result.size.y + global_id.y) * vc + global_id.x;
result.data[resi] = vec4(0f);
var mi = global_id.z * m.size.y * vc + global_id.x;
var vi = global_id.z * v.size.x;
for (var i = 0u; i < v.size.x; i++) {
if (global_id.y == v.data[vi]) {
result.data[resi] += m.data[mi];
}
mi += vc;
vi++;
}
}
+31
View File
@@ -0,0 +1,31 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<number>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.x = r.size.x;
result.size.y = l.size.y;
result.size.z = l.size.z;
var s = 0f;
var li = global_id.z * l.size.x * l.size.y + global_id.y * l.size.x;
var ri = global_id.z * r.size.x * r.size.y + global_id.x;
for (var i = 0u; i < l.size.x; i++) {
s += l.data[li] * r.data[ri];
li++;
ri += r.size.x;
}
result.data[global_id.z * result.size.y * result.size.x + global_id.y * result.size.x + global_id.x] = s;
}
+29
View File
@@ -0,0 +1,29 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> dldo: MatrixBatch;
@group(0) @binding(1)
var<storage, read> tout: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> otdldo: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
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;
}
+39
View File
@@ -0,0 +1,39 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
struct FVectorBatch {
size: vec4<u32>,
data: array<f32>,
}
struct UVectorBatch {
size: vec4<u32>,
data: array<u32>,
}
@group(0) @binding(0)
var<storage, read> m: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r: FVectorBatch;
@group(0) @binding(2)
var<storage, read_write> result: UVectorBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.x = m.size.y;
result.size.z = m.size.z;
var resi = global_id.z * result.size.x + global_id.x;
var tp = r.data[global_id.z * r.size.x + global_id.x];
var p = 0f;
var mi = (global_id.z * m.size.y + global_id.x) * m.size.x;
for (var i = 0u; i < m.size.x; i++) {
p += m.data[mi];
if (p >= tp) {
result.data[resi] = i;
break;
}
mi++;
}
}
+25
View File
@@ -0,0 +1,25 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<f32>>,
}
@group(0) @binding(0)
var<storage, read> b: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = b.size;
var vc = result.size.x / 4;
var resi = (global_id.z * result.size.y + global_id.y) * vc + global_id.x;
var mv = b.data[resi];
var bi = (global_id.z * b.size.y + global_id.y) * vc;
for (var i = 0u; i < vc; i++) {
mv = max(mv, b.data[bi]);
bi++;
}
var ms = max(max(mv.x, mv.y), max(mv.z, mv.w));
result.data[resi] = exp(b.data[resi] - vec4(ms));
}
+25
View File
@@ -0,0 +1,25 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<f32>>,
}
@group(0) @binding(0)
var<storage, read> b: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = b.size;
var vc = result.size.x / 4;
var resi = (global_id.z * result.size.y + global_id.y) * vc + global_id.x;
var sv = vec4(0f);
var bi = (global_id.z * b.size.y + global_id.y) * vc;
for (var i = 0u; i < vc; i++) {
sv += b.data[bi];
bi++;
}
var ss = sv.x + sv.y + sv.z + sv.w;
result.data[resi] = b.data[resi] / vec4(ss);
}
+32
View File
@@ -0,0 +1,32 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
struct Matrix {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> m: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: Matrix;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = m.size;
result.size.z = 1;
var vc = result.size.x / 4;
var resi = global_id.y * vc + global_id.x;
result.data[resi] = vec4(0f);
var mi = resi;
for (var z = 0u; z < m.size.z; z++) {
result.data[resi] += m.data[mi];
mi += m.size.y * vc;
}
}
+29
View File
@@ -0,0 +1,29 @@
struct MatrixBatch {
size: vec4<u32>,
data: array<f32>,
}
@group(0) @binding(0)
var<storage, read> wa: MatrixBatch;
@group(0) @binding(1)
var<storage, read> otdldo: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> totdldowa: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
totdldowa.size.x = otdldo.size.y;
totdldowa.size.y = wa.size.x;
totdldowa.size.z = otdldo.size.z;
var s = 0f;
var li = global_id.z * otdldo.size.x * otdldo.size.y + global_id.x * otdldo.size.x;
var ri = global_id.z * wa.size.x * wa.size.y + global_id.y;
for (var i = 0u; i < otdldo.size.x; i++) {
s += otdldo.data[li] * wa.data[ri];
li++;
ri += wa.size.x;
}
totdldowa.data[(global_id.z * totdldowa.size.y + global_id.y) * totdldowa.size.x + global_id.x] = s;
}
+20
View File
@@ -0,0 +1,20 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<number>,
}
@group(0) @binding(0)
var<storage, read> in: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> out: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
out.size.x = in.size.y;
out.size.y = in.size.x;
out.size.z = in.size.z;
var l = global_id.z * out.size.y * out.size.x;
out.data[l + global_id.y * out.size.x + global_id.x] = in.data[l + global_id.x * in.size.x + global_id.y];
}
+26
View File
@@ -0,0 +1,26 @@
alias number = f32;
struct Matrix {
size: vec4<u32>,
data: array<vec4<number>>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> m: Matrix;
@group(0) @binding(1)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>, @builtin(num_workgroups) num_workgroups: vec3<u32>) {
result.size = m.size;
result.size.z = num_workgroups.z * 2;
var vc = result.size.x / 4;
var mi = global_id.y * vc + global_id.x;
result.data[global_id.z * result.size.y * vc + mi] = m.data[mi];
}