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/layernormgrad.wgsl

51 lines
1.5 KiB

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