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.
|
|
struct MatrixBatch {
|
|
|
size: vec4<u32>,
|
|
|
data: array<vec4<f32>>,
|
|
|
}
|
|
|
|
|
|
@group(0) @binding(0) var<storage, read> m: MatrixBatch;
|
|
|
@group(0) @binding(1) var<storage, read> dldm: MatrixBatch;
|
|
|
@group(0) @binding(2) var<storage, read_write> grad: MatrixBatch;
|
|
|
|
|
|
@compute @workgroup_size(1)
|
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
|
grad.size = m.size;
|
|
|
|
|
|
let vc = m.size.x / 4;
|
|
|
let ro = (global_id.z * m.size.y + global_id.y) * vc;
|
|
|
|
|
|
var sum = vec4(0f);
|
|
|
var sumsq = vec4(0f);
|
|
|
for (var i = 0u; i < vc; i++) {
|
|
|
let v = m.data[ro + i];
|
|
|
sum += v;
|
|
|
sumsq += v * v;
|
|
|
}
|
|
|
|
|
|
let mean = (sum.x + sum.y + sum.z + sum.w) / f32(m.size.x);
|
|
|
let variance = (sumsq.x + sumsq.y + sumsq.z + sumsq.w) / f32(m.size.x) - mean * mean;
|
|
|
let inv_std = inverseSqrt((sumsq.x + sumsq.y + sumsq.z + sumsq.w) / f32(m.size.x) - mean * mean + 1e-5);
|
|
|
|
|
|
// Recompute x̂
|
|
|
var sum_grad = vec4(0.0);
|
|
|
var sum_grad_xhat = vec4(0.0);
|
|
|
|
|
|
for (var i = 0u; i < vc; i++) {
|
|
|
let v = m.data[ro + i];
|
|
|
let g = dldm.data[ro + i];
|
|
|
let xhat = (v - vec4(mean)) * vec4(inv_std);
|
|
|
sum_grad += g;
|
|
|
sum_grad_xhat += g * xhat;
|
|
|
}
|
|
|
|
|
|
let mg = (sum_grad.x + sum_grad.y + sum_grad.z + sum_grad.w) / f32(m.size.x);
|
|
|
let mgxhat = (sum_grad_xhat.x + sum_grad_xhat.y + sum_grad_xhat.z + sum_grad_xhat.w) / f32(m.size.x);
|
|
|
|
|
|
for (var i = 0u; i < vc; i++) {
|
|
|
let v = m.data[ro + i];
|
|
|
let g = dldm.data[ro + i];
|
|
|
let xhat = (v - vec4(mean)) * vec4(inv_std);
|
|
|
let dx = vec4(inv_std) * (g - vec4(mg) - xhat * vec4(mgxhat));
|
|
|
grad.data[ro + i] = dx;
|
|
|
}
|
|
|
}
|
|
|
|