From 5f084cd9ddb1cd4b39d410891e884ad97f2693ba Mon Sep 17 00:00:00 2001 From: Freywar Ulvnaudgari Date: Tue, 21 Jul 2026 14:27:31 +0300 Subject: [PATCH] Add residuals --- shader/center.wgsl | 24 +++ shader/dlds.wgsl | 48 ------ shader/gradinp.wgsl | 47 ------ shader/gradtattw.wgsl | 28 ---- shader/gradtout.wgsl | 28 ---- shader/layernormgrad.wgsl | 51 ++++++ shader/normalize.wgsl | 24 +++ shader/{hotones.wgsl => onehots.wgsl} | 0 shader/otdldo.wgsl | 29 ---- shader/softmaxgrad.wgsl | 26 +++ shader/tmul.wgsl | 34 ++++ shader/totdldowa.wgsl | 29 ---- shader/{attw.wgsl => tweigh.wgsl} | 0 src/data/gpu/gpu-struct.ts | 2 + src/data/gpu/matrix-batch.ts | 20 +++ src/data/gpu/matrix.ts | 19 +++ src/model-gpu.ts | 230 +++++++++++++++++++------- 17 files changed, 367 insertions(+), 272 deletions(-) create mode 100644 shader/center.wgsl delete mode 100644 shader/dlds.wgsl delete mode 100644 shader/gradinp.wgsl delete mode 100644 shader/gradtattw.wgsl delete mode 100644 shader/gradtout.wgsl create mode 100644 shader/layernormgrad.wgsl create mode 100644 shader/normalize.wgsl rename shader/{hotones.wgsl => onehots.wgsl} (100%) delete mode 100644 shader/otdldo.wgsl create mode 100644 shader/softmaxgrad.wgsl create mode 100644 shader/tmul.wgsl delete mode 100644 shader/totdldowa.wgsl rename shader/{attw.wgsl => tweigh.wgsl} (100%) diff --git a/shader/center.wgsl b/shader/center.wgsl new file mode 100644 index 0000000..cc659c2 --- /dev/null +++ b/shader/center.wgsl @@ -0,0 +1,24 @@ +struct MatrixBatch { + size: vec4, + data: array>, +} + +@group(0) @binding(0) +var m: MatrixBatch; +@group(0) @binding(1) +var centered: MatrixBatch; + +@compute @workgroup_size(1) +fn main(@builtin(global_invocation_id) global_id: vec3) { + 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); +} diff --git a/shader/dlds.wgsl b/shader/dlds.wgsl deleted file mode 100644 index c8907d1..0000000 --- a/shader/dlds.wgsl +++ /dev/null @@ -1,48 +0,0 @@ -struct Scalar { - size: vec4, - data: array, -} - -struct MatrixBatch { - size: vec4, - data: array, -} - -@group(0) @binding(0) -var ds: Scalar; -@group(0) @binding(1) -var wa: MatrixBatch; -@group(0) @binding(2) -var v: MatrixBatch; -@group(0) @binding(3) -var otdldo: MatrixBatch; -@group(0) @binding(4) -var dlds: MatrixBatch; - -@compute @workgroup_size(1) -fn main(@builtin(global_invocation_id) global_id: vec3) { - 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; -} diff --git a/shader/gradinp.wgsl b/shader/gradinp.wgsl deleted file mode 100644 index 8f5e6dd..0000000 --- a/shader/gradinp.wgsl +++ /dev/null @@ -1,47 +0,0 @@ -struct MatrixBatch { - size: vec4, - data: array, -} - -@group(0) @binding(0) -var ilogits: MatrixBatch; -@group(0) @binding(0) -var tilogits: MatrixBatch; -@group(0) @binding(1) -var dldsq: MatrixBatch; -@group(0) @binding(2) -var dldsk: MatrixBatch; -@group(0) @binding(3) -var totdldowa: MatrixBatch; -@group(0) @binding(4) -var tqry: MatrixBatch; -@group(0) @binding(5) -var tkey: MatrixBatch; -@group(0) @binding(6) -var tval: MatrixBatch; -@group(0) @binding(7) -var ginp: MatrixBatch; - -@compute @workgroup_size(1) -fn main(@builtin(global_invocation_id) global_id: vec3) { - 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; -} diff --git a/shader/gradtattw.wgsl b/shader/gradtattw.wgsl deleted file mode 100644 index 6af53fa..0000000 --- a/shader/gradtattw.wgsl +++ /dev/null @@ -1,28 +0,0 @@ -struct MatrixBatch { - size: vec4, - data: array, -} - -@group(0) @binding(0) -var e: MatrixBatch; -@group(0) @binding(1) -var dlds: MatrixBatch; -@group(0) @binding(2) -var gtattw: MatrixBatch; - -@compute @workgroup_size(1) -fn main(@builtin(global_invocation_id) global_id: vec3) { - 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; -} diff --git a/shader/gradtout.wgsl b/shader/gradtout.wgsl deleted file mode 100644 index 4b68a76..0000000 --- a/shader/gradtout.wgsl +++ /dev/null @@ -1,28 +0,0 @@ -struct MatrixBatch { - size: vec4, - data: array, -} - -@group(0) @binding(0) -var a: MatrixBatch; -@group(0) @binding(1) -var dldo: MatrixBatch; -@group(0) @binding(2) -var gtout: MatrixBatch; - -@compute @workgroup_size(1) -fn main(@builtin(global_invocation_id) global_id: vec3) { - 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; -} diff --git a/shader/layernormgrad.wgsl b/shader/layernormgrad.wgsl new file mode 100644 index 0000000..faf7396 --- /dev/null +++ b/shader/layernormgrad.wgsl @@ -0,0 +1,51 @@ +struct MatrixBatch { + size: vec4, + data: array>, +} + +@group(0) @binding(0) var m: MatrixBatch; +@group(0) @binding(1) var dldm: MatrixBatch; +@group(0) @binding(2) var grad: MatrixBatch; + +@compute @workgroup_size(1) +fn main(@builtin(global_invocation_id) global_id: vec3) { + 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; + } +} diff --git a/shader/normalize.wgsl b/shader/normalize.wgsl new file mode 100644 index 0000000..97adb0e --- /dev/null +++ b/shader/normalize.wgsl @@ -0,0 +1,24 @@ +struct MatrixBatch { + size: vec4, + data: array>, +} + +@group(0) @binding(0) +var c: MatrixBatch; +@group(0) @binding(1) +var result: MatrixBatch; + +@compute @workgroup_size(1) +fn main(@builtin(global_invocation_id) global_id: vec3) { + 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)); +} diff --git a/shader/hotones.wgsl b/shader/onehots.wgsl similarity index 100% rename from shader/hotones.wgsl rename to shader/onehots.wgsl diff --git a/shader/otdldo.wgsl b/shader/otdldo.wgsl deleted file mode 100644 index a9ae905..0000000 --- a/shader/otdldo.wgsl +++ /dev/null @@ -1,29 +0,0 @@ -struct MatrixBatch { - size: vec4, - data: array, -} - -@group(0) @binding(0) -var dldo: MatrixBatch; -@group(0) @binding(1) -var tout: MatrixBatch; -@group(0) @binding(2) -var otdldo: MatrixBatch; - -@compute @workgroup_size(1) -fn main(@builtin(global_invocation_id) global_id: vec3) { - 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; -} diff --git a/shader/softmaxgrad.wgsl b/shader/softmaxgrad.wgsl new file mode 100644 index 0000000..c3b7156 --- /dev/null +++ b/shader/softmaxgrad.wgsl @@ -0,0 +1,26 @@ +struct MatrixBatch { + size: vec4, + data: array>, +} + +@group(0) @binding(0) +var m: MatrixBatch; +@group(0) @binding(1) +var dldm: MatrixBatch; +@group(0) @binding(2) +var grad: MatrixBatch; + +@compute @workgroup_size(1) +fn main(@builtin(global_invocation_id) global_id: vec3) { + 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)); +} diff --git a/shader/tmul.wgsl b/shader/tmul.wgsl new file mode 100644 index 0000000..2e91f1e --- /dev/null +++ b/shader/tmul.wgsl @@ -0,0 +1,34 @@ +alias number = f32; + +struct PackedMatrixBatch { + size: vec4, + data: array>, +} + +struct MatrixBatch { + size: vec4, + data: array, +} + +@group(0) @binding(0) +var l: PackedMatrixBatch; +@group(0) @binding(1) +var tr: PackedMatrixBatch; +@group(0) @binding(2) +var result: MatrixBatch; + +@compute @workgroup_size(1) +fn main(@builtin(global_invocation_id) global_id: vec3) { + 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; +} diff --git a/shader/totdldowa.wgsl b/shader/totdldowa.wgsl deleted file mode 100644 index e92b10a..0000000 --- a/shader/totdldowa.wgsl +++ /dev/null @@ -1,29 +0,0 @@ -struct MatrixBatch { - size: vec4, - data: array, -} - -@group(0) @binding(0) -var wa: MatrixBatch; -@group(0) @binding(1) -var otdldo: MatrixBatch; -@group(0) @binding(2) -var totdldowa: MatrixBatch; - -@compute @workgroup_size(1) -fn main(@builtin(global_invocation_id) global_id: vec3) { - 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; -} diff --git a/shader/attw.wgsl b/shader/tweigh.wgsl similarity index 100% rename from shader/attw.wgsl rename to shader/tweigh.wgsl diff --git a/src/data/gpu/gpu-struct.ts b/src/data/gpu/gpu-struct.ts index be872ee..2116ef6 100644 --- a/src/data/gpu/gpu-struct.ts +++ b/src/data/gpu/gpu-struct.ts @@ -151,6 +151,8 @@ export abstract class GPUStruct; + public buffer(): GPUBuffer { if (this._dirty === 'front') { this.write(); diff --git a/src/data/gpu/matrix-batch.ts b/src/data/gpu/matrix-batch.ts index 297ec15..a787bef 100644 --- a/src/data/gpu/matrix-batch.ts +++ b/src/data/gpu/matrix-batch.ts @@ -9,6 +9,7 @@ const SHADER_SCALE = fs.readFileSync('shader/bscale.wgsl').toString(); const SHADER_ISUM = fs.readFileSync('shader/isum.wgsl').toString(); const SHADER_MUL = fs.readFileSync('shader/mul.wgsl').toString(); const SHADER_TRANS = fs.readFileSync('shader/transpose.wgsl').toString(); +const SHADER_TMUL = fs.readFileSync('shader/tmul.wgsl').toString(); const SHADER_AVG = fs.readFileSync('shader/avg.wgsl').toString(); const SHADER_TOT = fs.readFileSync('shader/tot.wgsl').toString(); @@ -253,6 +254,25 @@ export class MatrixBatch< return result.calculate('transpose', SHADER_TRANS, [m]); } + public static tmul< + L extends number, + H extends number, + S extends number, + W extends number, + A extends Uint32Array | Int32Array | Float32Array, + T extends NonNullable>, + >( + l: MatrixBatch, + tr: MatrixBatch, + result: MatrixBatch = new MatrixBatch(l._device, l._array, l.layers, l.height, tr.height), + ): MatrixBatch { + this.assertUnique(l, tr, result); + this.assertColocated(l, tr, result); + this.assertPackable(1, l, tr, result); + + return result.calculate('tmul', SHADER_TMUL, [l, tr]); + } + public static avg< L extends number, H extends number, diff --git a/src/data/gpu/matrix.ts b/src/data/gpu/matrix.ts index 1ece278..149983f 100644 --- a/src/data/gpu/matrix.ts +++ b/src/data/gpu/matrix.ts @@ -8,6 +8,7 @@ import { Vector } from './vector'; const SHADER_ISUM = fs.readFileSync('shader/isum.wgsl').toString(); const SHADER_MUL = fs.readFileSync('shader/mul.wgsl').toString(); const SHADER_TRANS = fs.readFileSync('shader/transpose.wgsl').toString(); +const SHADER_TMUL = fs.readFileSync('shader/tmul.wgsl').toString(); const SHADER_UNAVG = fs.readFileSync('shader/unavg.wgsl').toString(); export class Matrix< @@ -197,6 +198,24 @@ export class Matrix< return result.calculate('mul', SHADER_MUL, [l, r]); } + public static tmul< + H extends number, + S extends number, + W extends number, + A extends Uint32Array | Int32Array | Float32Array, + T extends NonNullable>, + >( + l: Matrix, + tr: Matrix, + result: Matrix = new Matrix(l._device, l._array, l.height, tr.height), + ): Matrix { + this.assertUnique(l, tr, result); + this.assertColocated(l, tr, result); + this.assertPackable(1, l, tr, result); + + return result.calculate('tmul', SHADER_TMUL, [l, tr]); + } + public static transpose< H extends number, W extends number, diff --git a/src/model-gpu.ts b/src/model-gpu.ts index 01774b6..e3084d4 100644 --- a/src/model-gpu.ts +++ b/src/model-gpu.ts @@ -1,26 +1,27 @@ import fs from 'node:fs'; import { Writable } from 'node:stream'; +import { GPUStruct } from './data/gpu/gpu-struct'; import { Matrix } from './data/gpu/matrix'; import { Scalar } from './data/gpu/scalar'; import { Vector } from './data/gpu/vector'; import { Matrix as CPUMatrix, Vector as CPUVector, drm, mat, randseq, slope, stddev, vec } from './data/math'; import { TID, Tokenizer } from './tokenizer'; -const SHADER_HOTONES = fs.readFileSync('shader/hotones.wgsl').toString(); +const SHADER_ONE_HOTS = fs.readFileSync('shader/onehots.wgsl').toString(); const SHADER_EMBED = fs.readFileSync('shader/embed.wgsl').toString(); -const SHADER_ATTW = fs.readFileSync('shader/attw.wgsl').toString(); +const SHADER_CENTER = fs.readFileSync('shader/center.wgsl').toString(); +const SHADER_NORMALIZE = fs.readFileSync('shader/normalize.wgsl').toString(); +const SHADER_TWEIGH = fs.readFileSync('shader/tweigh.wgsl').toString(); const SHADER_SOFTMAX_EXP = fs.readFileSync('shader/softmaxexp.wgsl').toString(); const SHADER_SOFTMAX_NORM = fs.readFileSync('shader/softmaxnorm.wgsl').toString(); const SHADER_DLDO = fs.readFileSync('shader/dldo.wgsl').toString(); -const SHADER_OTDLDO = fs.readFileSync('shader/otdldo.wgsl').toString(); -const SHADER_TOTDLDOWA = fs.readFileSync('shader/totdldowa.wgsl').toString(); -const SHADER_DLDS = fs.readFileSync('shader/dlds.wgsl').toString(); -const SHADER_GRAD_TOUT = fs.readFileSync('shader/gradtout.wgsl').toString(); -const SHADER_GRAD_TATTW = fs.readFileSync('shader/gradtattw.wgsl').toString(); -const SHADER_GRAD_INP = fs.readFileSync('shader/gradinp.wgsl').toString(); +const SHADER_SOFTMAX_GRAD = fs.readFileSync('shader/softmaxgrad.wgsl').toString(); +const SHADER_LAYERNORM_GRAD = fs.readFileSync('shader/layernormgrad.wgsl').toString(); const SHADER_ARGMAX = fs.readFileSync('shader/argmax.wgsl').toString(); const SHADER_SAMPLES = fs.readFileSync('shader/samples.wgsl').toString(); +const LOGGING = false; + export interface ModelInitData< V extends number, W extends number, @@ -126,14 +127,16 @@ export class Model< readonly itids: Vector; readonly ilogits: Matrix; readonly e: Matrix; + readonly ce: Matrix; + readonly ne: Matrix; readonly q: Matrix; readonly k: Matrix; - readonly tk: Matrix; readonly qtk: Matrix; readonly v: Matrix; readonly s: Matrix; readonly se: Matrix; readonly wa: Matrix; + readonly ar: Matrix; readonly a: Matrix; readonly o: Matrix; readonly temp: Scalar; @@ -156,6 +159,12 @@ export class Model< readonly ttids: Vector; }; + protected async log(name: string, s: GPUStruct, sample: number = 4): Promise { + if (LOGGING) { + console.log(name, (await s.get() as number[]).slice(0, sample)); + } + } + protected softmax( m: Matrix, e: Matrix = new Matrix(this.device, Float32Array, m.height, m.width), @@ -166,28 +175,75 @@ export class Model< return r; } + protected softmax_grad( + m: Matrix, + dldm: Matrix, + result: Matrix = new Matrix(this.device, Float32Array, m.height, m.width), + ): Matrix { + result.calculate('softmaxgrad', SHADER_SOFTMAX_GRAD, [m, dldm]); + return result; + } + + protected async layernorm( + m: Matrix, + c: Matrix = new Matrix(this.device, Float32Array, m.height, m.width), + r: Matrix = new Matrix(this.device, Float32Array, m.height, m.width), + ): Promise> { + c.calculate('center', SHADER_CENTER, [m], false, 4); + r.calculate('normalize', SHADER_NORMALIZE, [m], false, 4); + await this.log('layernorm c', c); + await this.log('layernorm r', r); + return r; + } + + protected layernorm_grad( + m: Matrix, + dldm: Matrix, + result: Matrix = new Matrix(this.device, Float32Array, m.height, m.width), + ): Matrix { + result.calculate('layernormgrad', SHADER_LAYERNORM_GRAD, [m, dldm]); + return result; + } + protected async ilogits(itids: Vector = this.processing.itids): Promise> { - this.processing.ilogits.calculate('hotones', SHADER_HOTONES, [itids]); + this.processing.ilogits.calculate('onehots', SHADER_ONE_HOTS, [itids]); + await this.log('itids', itids); + await this.log('ilogits', this.processing.ilogits); return this.processing.ilogits; } protected async e(ilogits: Matrix = this.processing.ilogits): Promise> { this.processing.e.calculate('embed', SHADER_EMBED, [ilogits, this.weights.inp, this.weights.pos]); + await this.log('inp', this.weights.inp); + await this.log('pos', this.weights.pos); + await this.log('e', this.processing.e); return this.processing.e; } - protected q(e: Matrix = this.processing.e): Matrix { - this.processing.q.calculate('attw', SHADER_ATTW, [e, this.weights.tqry]); + protected async ne(e: Matrix = this.processing.e): Promise> { + this.layernorm(e, this.processing.ce, this.processing.ne); + await this.log('ne', this.processing.ne); + return this.processing.ne; + } + + protected async q(ne: Matrix = this.processing.ne): Promise> { + this.processing.q.calculate('tqry', SHADER_TWEIGH, [ne, this.weights.tqry]); + await this.log('tqry', this.weights.tqry); + await this.log('q', this.processing.q); return this.processing.q; } - protected k(e: Matrix = this.processing.e): Matrix { - this.processing.k.calculate('attw', SHADER_ATTW, [e, this.weights.tkey]); + protected async k(ne: Matrix = this.processing.ne): Promise> { + this.processing.k.calculate('tkey', SHADER_TWEIGH, [ne, this.weights.tkey]); + await this.log('tkey', this.weights.tkey); + await this.log('k', this.processing.k); return this.processing.k; } - protected v(e: Matrix = this.processing.e): Matrix { - this.processing.v.calculate('attw', SHADER_ATTW, [e, this.weights.tval]); + protected async v(ne: Matrix = this.processing.ne): Promise> { + this.processing.v.calculate('tval', SHADER_TWEIGH, [ne, this.weights.tval]); + await this.log('tval', this.weights.tval); + await this.log('v', this.processing.v); return this.processing.v; } @@ -195,8 +251,7 @@ export class Model< q: Matrix = this.processing.q, k: Matrix = this.processing.k, ): Matrix { - Matrix.transpose(k, this.processing.tk); - Matrix.mul(q, this.processing.tk, this.processing.qtk); + Matrix.tmul(q, k, this.processing.qtk); Matrix.scale(this.processing.ds, this.processing.qtk, this.processing.s); return this.processing.s; } @@ -207,15 +262,17 @@ export class Model< } protected a( + e: Matrix = this.processing.e, wa: Matrix = this.processing.wa, v: Matrix = this.processing.v, ): Matrix { - Matrix.mul(wa, v, this.processing.a); + Matrix.mul(wa, v, this.processing.ar); + Matrix.sum(e, this.processing.ar, this.processing.a); return this.processing.a; } protected o(a: Matrix = this.processing.a): Matrix { - this.processing.o.calculate('attw', SHADER_ATTW, [a, this.weights.tout]); + this.processing.o.calculate('tout', SHADER_TWEIGH, [a, this.weights.tout]); return this.processing.o; } @@ -252,7 +309,7 @@ export class Model< protected async upd( ilogits: Matrix = this.processing.ilogits, - e: Matrix = this.processing.e, + ne: Matrix = this.processing.ne, q: Matrix = this.processing.q, k: Matrix = this.processing.k, v: Matrix = this.processing.v, @@ -270,39 +327,80 @@ export class Model< this.processing.lr, ]); - this.processing.otdldo.calculate('otdldo', SHADER_OTDLDO, [this.processing.dldo, this.weights.tout]); - - this.processing.dlds.calculate('dlds', SHADER_DLDS, [ - this.processing.ds, + this.softmax_grad( wa, - v, - this.processing.otdldo, - ]); + Matrix.mul(this.processing.dldo, Matrix.mul(this.weights.tout, Matrix.transpose(v))), + this.processing.dlds, + ); + + Matrix.mul( + Matrix.transpose(ilogits), + Matrix.sum( + Matrix.mul(this.processing.dldo, this.weights.tout), + this.layernorm_grad( + ne, + Matrix.sum( + Matrix.mul( + Matrix.mul( + this.processing.dlds, + Matrix.scale(this.processing.ds, k), + ), + this.weights.tqry, + ) + , + Matrix.sum( + Matrix.mul( + Matrix.mul( + Matrix.transpose(this.processing.dlds), + Matrix.scale(this.processing.ds, q), + ), + this.weights.tkey, + ), + Matrix.mul( + Matrix.mul(Matrix.transpose(wa), this.processing.dldo), + Matrix.mul(this.weights.tout, this.weights.tval), + ), + ), + ), + ), + ), + this.processing.ginp, + ); + + Matrix.transpose( + Matrix.scale(this.processing.ds, Matrix.mul( + Matrix.transpose(ne), + Matrix.mul(this.processing.dlds, k), + )), + this.processing.gtqry, + ); + + Matrix.transpose( + Matrix.scale(this.processing.ds, Matrix.mul( + Matrix.transpose(ne), + Matrix.mul(Matrix.transpose(this.processing.dlds), q), + )), + this.processing.gtkey, + ); + + Matrix.transpose( + Matrix.mul( + Matrix.mul(Matrix.transpose(ne), Matrix.transpose(wa)), + Matrix.mul(this.processing.dldo, this.weights.tout), + ), + this.processing.gtval, + ); + + Matrix.transpose( + Matrix.mul(Matrix.transpose(a), this.processing.dldo), + this.processing.gtout, + ); - Matrix.mul(this.processing.dlds, k, this.processing.dldsk); - Matrix.mul(Matrix.transpose(this.processing.dlds), q, this.processing.dldsq); - - this.processing.ginp.calculate('gradinp', SHADER_GRAD_INP, [ - ilogits, - this.processing.dldsq, - this.processing.dldsk, - this.processing.totdldowa, - this.weights.tqry, - this.weights.tkey, - this.weights.tval, - ]); Matrix.sum(this.processing.ginp, this.weights.inp, this.weights.inp); - - this.processing.gtout.calculate('gradtout', SHADER_GRAD_TOUT, [a, this.processing.dldo]); - Matrix.sum(this.processing.gtout, this.weights.tout, this.weights.tout); - this.processing.gtqry.calculate('gradtqry', SHADER_GRAD_TATTW, [e, this.processing.dldsk]); Matrix.sum(this.processing.gtqry, this.weights.tqry, this.weights.tqry); - this.processing.gtkey.calculate('gradtkey', SHADER_GRAD_TATTW, [e, this.processing.dldsq]); Matrix.sum(this.processing.gtkey, this.weights.tkey, this.weights.tkey); - - this.processing.totdldowa.calculate('totdldowa', SHADER_TOTDLDOWA, [wa, this.processing.otdldo]); - this.processing.gtval.calculate('gradtval', SHADER_GRAD_TATTW, [e, this.processing.totdldowa]); Matrix.sum(this.processing.gtval, this.weights.tval, this.weights.tval); + Matrix.sum(this.processing.gtout, this.weights.tout, this.weights.tout); } public constructor( @@ -349,15 +447,17 @@ export class Model< itids: new Vector(this.device, Int32Array, this.ws), ilogits: new Matrix(this.device, Float32Array, this.ws, this.vs), e: new Matrix(this.device, Float32Array, this.ws, this.ds), + ce: new Matrix(this.device, Float32Array, this.ws, this.ds), + ne: new Matrix(this.device, Float32Array, this.ws, this.ds), q: new Matrix(this.device, Float32Array, this.ws, this.ds), k: new Matrix(this.device, Float32Array, this.ws, this.ds), - tk: new Matrix(this.device, Float32Array, this.ds, this.ws), qtk: new Matrix(this.device, Float32Array, this.ws, this.ws), v: new Matrix(this.device, Float32Array, this.ws, this.ds), s: new Matrix(this.device, Float32Array, this.ws, this.ws), se: new Matrix(this.device, Float32Array, this.ws, this.ws), wa: new Matrix(this.device, Float32Array, this.ws, this.ws), a: new Matrix(this.device, Float32Array, this.ws, this.ds), + ar: new Matrix(this.device, Float32Array, this.ws, this.ds), o: new Matrix(this.device, Float32Array, this.ws, this.vs), temp: new Scalar(this.device, Float32Array, 1), tempo: new Matrix(this.device, Float32Array, this.ws, this.vs), @@ -458,6 +558,8 @@ export class Model< const tids = this.tokenizer.tokenize(data); const batches = tids.length - this.ws - 1; + const seed = tids.slice(tids.length / 2, tids.length / 2 + 100); + let lr = 0.003; for (let epoch = 0; epoch < epochs; epoch++) { @@ -471,15 +573,16 @@ export class Model< const ilogits = await this.ilogits(this.processing.itids); const e = await this.e(ilogits); - const q = this.q(e); - const k = this.k(e); - const v = this.v(e); + const ne = await this.ne(e); + const q = await this.q(ne); + const k = await this.k(ne); + const v = await this.v(ne); const s = this.s(q, k); const wa = await this.wa(s); - const a = this.a(wa, v); + const a = this.a(e, wa, v); const o = this.o(a); const ologits = await this.ologits(o); - await this.upd(ilogits, e, q, k, v, wa, a, ologits, this.processing.ttids, lr); + await this.upd(ilogits, ne, q, k, v, wa, a, ologits, this.processing.ttids, lr); const queued: number = Date.now(); @@ -567,19 +670,19 @@ export class Model< samples?.write(`Epoch ${epoch + 1}:\n`); samples?.write('Sample (temp=0.3):\n'); - samples?.write(this.tokenizer.detokenize(await this.run(this.tokenizer.tokenize('Hello?'), { temp: 0.3 })).join('')); + samples?.write(this.tokenizer.detokenize(await this.run(seed, { temp: 0.3 })).join('')); samples?.write('\nEnd of sample.\n'); samples?.write('Sample (temp=0.7):\n'); - samples?.write(this.tokenizer.detokenize(await this.run(this.tokenizer.tokenize('Hello?'), { temp: 0.7 })).join('')); + samples?.write(this.tokenizer.detokenize(await this.run(seed, { temp: 0.7 })).join('')); samples?.write('\nEnd of sample.\n'); samples?.write('Sample (temp=1):\n'); - samples?.write(this.tokenizer.detokenize(await this.run(this.tokenizer.tokenize('Hello?'), { temp: 1 })).join('')); + samples?.write(this.tokenizer.detokenize(await this.run(seed, { temp: 1 })).join('')); samples?.write('\nEnd of sample.\n'); samples?.write('Sample (temp=1.3):\n'); - samples?.write(this.tokenizer.detokenize(await this.run(this.tokenizer.tokenize('Hello?'), { temp: 1.3 })).join('')); + samples?.write(this.tokenizer.detokenize(await this.run(seed, { temp: 1.3 })).join('')); samples?.write('\nEnd of sample.\n\n'); samples?.write('Sample (argmax):\n'); - samples?.write(this.tokenizer.detokenize(await this.run(this.tokenizer.tokenize('Hello?'), { temp: 'max' })).join('')); + samples?.write(this.tokenizer.detokenize(await this.run(seed, { temp: 'max' })).join('')); samples?.write('\nEnd of sample.\n\n'); snapshotted = end; @@ -606,12 +709,13 @@ export class Model< await this.processing.ttids.set([...itids.slice(1), 0]); const ilogits = await this.ilogits(this.processing.itids); const e = await this.e(ilogits); - const q = this.q(e); - const k = this.k(e); - const v = this.v(e); + const ne = await this.ne(e); + const q = await this.q(ne); + const k = await this.k(ne); + const v = await this.v(ne); const s = this.s(q, k); const wa = await this.wa(s); - const a = this.a(wa, v); + const a = this.a(e, wa, v); const o = this.o(a); const ologits = await this.ologits(o, temp === 'max' ? 1 : temp); const out = temp === 'max' ? this.otids(ologits) : this.otids(await this.rnd(), ologits);