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/src/data/gpu/matrix.ts

252 lines
8.4 KiB

import fs from 'node:fs';
import { Matrix as CPUMatrix, Vector as CPUVector } from '../math';
import { GPUStruct } from './gpu-struct';
import { MatrixBatch } from './matrix-batch';
import { Scalar } from './scalar';
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<
H extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>> = NonNullable<ReturnType<A['at']>>,
> extends GPUStruct<A> {
public get width(): W {
return this.size[0] as W;
}
public get height(): H {
return this.size[1] as H;
}
protected override workgroups(pm: number = 1): [number, number, number] {
return [this.width / GPUStruct.dunit / pm, this.height / GPUStruct.dunit, 1];
}
protected override workgroup(): [number, number, number] {
return [GPUStruct.dunit, GPUStruct.dunit, 1];
}
public constructor(device: GPUDevice, array: new (size: number) => A, data: CPUMatrix<H, W, T>);
public constructor(device: GPUDevice, array: new (size: number) => A, h: H, w: W, data?: CPUMatrix<H, W, T>);
public constructor(device: GPUDevice, array: new (size: number) => A, ...args: unknown[]) {
let [h, w]: [number, number] = [1, 1];
if (typeof args[0] === 'number') {
[h, w] = args as [H, W];
} else {
const [data] = args as [CPUMatrix<H, W, T>];
[h, w] = [data.length, data[0].length];
}
super(device, array, `${h}x${w}`, [w, h, 1, 1]);
Matrix.assertPackable(1, this);
const data = (typeof args[0] === 'number' ? args[2] : args[0]) as CPUMatrix<H, W, T> | undefined;
if (data) {
this.set(data);
}
}
public index(y: number, x: number): number {
return y * this.width + x;
}
public async set(y: number, x: number, v: T): Promise<void>;
public async set(data: CPUMatrix<H, W, T>): Promise<void>;
public async set(slice: T[][], oy?: number, ox?: number): Promise<void>;
public async set(m: Matrix<H, W, A, T>): Promise<void>;
public async set(...args: unknown[]): Promise<void> {
if (typeof args[0] === 'number') {
const [y, x, v]: number[] = args as number[];
if (this._dirty === 'back') {
await this.read();
}
this.data[this.index(y, x)] = v;
this._dirty = 'front';
} else if (Array.isArray(args[0])) {
const [data, oy = 0, ox = 0] = args as [T[][], number | undefined, number | undefined];
if (this._dirty === 'back' && (oy || ox || data.length < this.height || data[0].length < this.width)) {
await this.read();
}
let y = oy;
let x = ox;
for (const row of data ?? []) {
if (y >= this.height) {
break;
}
for (const v of row) {
if (x >= this.width) {
break;
}
this.data[this.index(y, x)] = v;
x++;
}
x = ox;
y++;
}
this._dirty = 'front';
} else {
const [matrix] = args as [Matrix<H, W, A, T>];
const copy = this._device.createCommandEncoder();
copy.copyBufferToBuffer(matrix.buffer(), 0, this.buffer(), 0, this.byteLength);
this._device.queue.submit([copy.finish()]);
this._dirty = 'back';
}
}
public async get(y: number, x: number): Promise<T>;
public async get(): Promise<CPUMatrix<H, W, T>>;
public async get(y?: number, x?: number): Promise<T | CPUMatrix<H, W, T>> {
if (this._dirty === 'back') {
await this.read();
}
if (typeof x === 'number') {
return this.data[this.index(y!, x!)] as T;
} else {
const result: CPUMatrix<H, W, T> = [] as unknown as CPUMatrix<H, W, T>;
for (let y: number = 0; y < this.height; y++) {
const row: CPUVector<W, T> = [] as unknown as CPUVector<W, T>;
for (let x: number = 0; x < this.width; x++) {
row.push(this.data[this.index(y, x)] as T);
}
result.push(row);
}
return result;
}
}
protected static override assertPackable(pm: number, ...ms: Matrix<number, number, Uint32Array | Int32Array | Float32Array>[]): void {
for (const m of ms) {
if (m.height % Matrix.dunit !== 0 || m.width % (Matrix.dunit * pm) !== 0) {
throw new Error(`Matrix size ${m.height}x${m.width} must be divisible by ${Matrix.dunit}x${Matrix.dunit * pm}.`);
}
}
}
public static sum<
H extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>> = NonNullable<ReturnType<A['at']>>,
>(l: Matrix<H, W, A, T>, r: Matrix<H, W, A, T>, result: Matrix<H, W, A, T> = new Matrix(l._device, l._array, l.height, l.width)): Matrix<H, W, A, T> {
return GPUStruct.csum(l, r, result);
}
public static sub<
H extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>> = NonNullable<ReturnType<A['at']>>,
>(l: Matrix<H, W, A, T>, r: Matrix<H, W, A, T>, result: Matrix<H, W, A, T> = new Matrix(l._device, l._array, l.height, l.width)): Matrix<H, W, A, T> {
return GPUStruct.csub(l, r, result);
}
public static isum<
H extends number,
I extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>>,
>(
lis: Vector<I, Int32Array>,
m: Matrix<H, W, A, T>,
result: Matrix<H, W, A, T> = new Matrix(m._device, m._array, m.height, m.width),
): Matrix<H, W, A, T> {
this.assertUnique(m, result);
this.assertColocated(m, result);
this.assertPackable(4, m, result);
return result.calculate('isum', SHADER_ISUM, [lis, m], false, 4);
}
public static scale<
H extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>> = NonNullable<ReturnType<A['at']>>,
>(l: Scalar<A, T>, r: Matrix<H, W, A, T>, result: Matrix<H, W, A, T> = new Matrix(r._device, r._array, r.height, r.width)): Matrix<H, W, A, T> {
return GPUStruct.cscale(l, r, result);
}
public static mul<
H extends number,
S extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>>,
>(
l: Matrix<H, S, A, T>,
r: Matrix<S, W, A, T>,
result: Matrix<H, W, A, T> = new Matrix(l._device, l._array, l.height, r.width),
): Matrix<H, W, A, T> {
this.assertUnique(l, r, result);
this.assertColocated(l, r, result);
this.assertPackable(1, l, r, result);
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<ReturnType<A['at']>>,
>(
l: Matrix<H, S, A, T>,
tr: Matrix<W, S, A, T>,
result: Matrix<H, W, A, T> = new Matrix(l._device, l._array, l.height, tr.height),
): Matrix<H, W, A, T> {
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,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>>,
>(
m: Matrix<H, W, A, T>,
result: Matrix<W, H, A, T> = new Matrix(m._device, m._array, m.width, m.height),
): Matrix<W, H, A, T> {
this.assertUnique(m, result);
this.assertColocated(m, result);
this.assertPackable(1, m, result);
return result.calculate('transpose', SHADER_TRANS, [m]);
}
public static unavg<
L extends number,
H extends number,
W extends number,
A extends Uint32Array | Int32Array | Float32Array,
T extends NonNullable<ReturnType<A['at']>>,
>(
l: L,
m: Matrix<H, W, A, T>,
result: MatrixBatch<L, H, W, A, T> = new MatrixBatch(m._device, m._array, l, m.height, m.width),
): MatrixBatch<L, H, W, A, T> {
this.assertUnique(m, result);
this.assertColocated(m, result);
this.assertPackable(4, m);
return result.calculate('unavg', SHADER_UNAVG, [m], false, 4);
}
}