Transformer architecture implemented in TypeScript with WebGPU.
You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
nn/shader/dots.wgsl

43 lines
1.1 KiB

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;
}