Add residuals

This commit is contained in:
Freywar Ulvnaudgari
2025-05-15 19:38:59 +03:00
parent 3c5f9a04dd
commit 44cd14f1a5
17 changed files with 366 additions and 271 deletions
+24
View File
@@ -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);
}
-48
View File
@@ -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;
}
-47
View File
@@ -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;
}
-28
View File
@@ -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;
}
-28
View File
@@ -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;
}
+51
View File
@@ -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;
}
}
+24
View File
@@ -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));
}
-29
View File
@@ -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;
}
+26
View File
@@ -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));
}
+34
View File
@@ -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;
}
-29
View File
@@ -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;
}