Add residuals
This commit is contained in:
@@ -0,0 +1,24 @@
|
||||
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_write> centered: MatrixBatch;
|
||||
|
||||
@compute @workgroup_size(1)
|
||||
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
||||
centered.size = m.size;
|
||||
|
||||
let vc = m.size.x / 4;
|
||||
let ro = (global_id.z * m.size.y + global_id.y) * vc;
|
||||
|
||||
var s = vec4(0f);
|
||||
for (var x = 0u; x < vc; x++) {
|
||||
s += m.data[ro + x];
|
||||
}
|
||||
|
||||
centered.data[ro + global_id.x] = m.data[ro + global_id.x] - vec4(s.x + s.y + s.z + s.w) / f32(m.size.x);
|
||||
}
|
||||
@@ -1,48 +0,0 @@
|
||||
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;
|
||||
}
|
||||
@@ -1,47 +0,0 @@
|
||||
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;
|
||||
}
|
||||
@@ -1,28 +0,0 @@
|
||||
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;
|
||||
}
|
||||
@@ -1,28 +0,0 @@
|
||||
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;
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
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;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
struct MatrixBatch {
|
||||
size: vec4<u32>,
|
||||
data: array<vec4<f32>>,
|
||||
}
|
||||
|
||||
@group(0) @binding(0)
|
||||
var<storage, read> c: 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 = c.size;
|
||||
|
||||
let vc = c.size.x / 4;
|
||||
let ro = (global_id.z * c.size.y + global_id.y) * vc;
|
||||
|
||||
var s = vec4(0f);
|
||||
for (var x = 0u; x < vc; x++) {
|
||||
s += c.data[ro + x] * c.data[ro + x];
|
||||
}
|
||||
|
||||
result.data[ro + global_id.x] = c.data[ro + global_id.x] * vec4(inverseSqrt((s.x + s.y + s.z + s.w) / f32(c.size.x) + 1e-5f));
|
||||
}
|
||||
@@ -1,29 +0,0 @@
|
||||
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;
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
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 s = 0f;
|
||||
for (var x = 0u; x < vc; x++) {
|
||||
s += dot(m.data[ro + x], dldm.data[ro + x]);
|
||||
}
|
||||
|
||||
grad.data[ro + global_id.x] = m.data[ro + global_id.x] * (dldm.data[ro + global_id.x] - vec4(s));
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
alias number = f32;
|
||||
|
||||
struct PackedMatrixBatch {
|
||||
size: vec4<u32>,
|
||||
data: array<vec4<number>>,
|
||||
}
|
||||
|
||||
struct MatrixBatch {
|
||||
size: vec4<u32>,
|
||||
data: array<number>,
|
||||
}
|
||||
|
||||
@group(0) @binding(0)
|
||||
var<storage, read> l: PackedMatrixBatch;
|
||||
@group(0) @binding(1)
|
||||
var<storage, read> tr: PackedMatrixBatch;
|
||||
@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;
|
||||
result.size.x = tr.size.y;
|
||||
|
||||
var s = 0f;
|
||||
let vc = l.size.x / 4;
|
||||
let lo = (global_id.z * l.size.y + global_id.y) * vc;
|
||||
let ro = (global_id.z * tr.size.y + global_id.x) * vc;
|
||||
for (var i = 0u; i < vc; i++) {
|
||||
s += dot(l.data[lo + i], tr.data[ro + i]);
|
||||
}
|
||||
|
||||
result.data[global_id.z * result.size.y * result.size.x + global_id.y * result.size.x + global_id.x] = s;
|
||||
}
|
||||
@@ -1,29 +0,0 @@
|
||||
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;
|
||||
}
|
||||
Reference in New Issue
Block a user