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_UNAVG = fs.readFileSync('shader/unavg.wgsl').toString(); export class Matrix< H extends number, W extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, > extends GPUStruct { 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); public constructor(device: GPUDevice, array: new (size: number) => A, h: H, w: W, data?: CPUMatrix); 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] = [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 | 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; public async set(data: CPUMatrix): Promise; public async set(slice: T[][], oy?: number, ox?: number): Promise; public async set(m: Matrix): Promise; public async set(...args: unknown[]): Promise { 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]; 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; public async get(): Promise>; public async get(y?: number, x?: number): Promise> { if (this._dirty === 'back') { await this.read(); } if (typeof x === 'number') { return this.data[this.index(y!, x!)] as T; } else { const result: CPUMatrix = [] as unknown as CPUMatrix; for (let y: number = 0; y < this.height; y++) { const row: CPUVector = [] as unknown as CPUVector; 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[]): 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> = NonNullable>, >(l: Matrix, r: Matrix, result: Matrix = new Matrix(l._device, l._array, l.height, l.width)): Matrix { return GPUStruct.csum(l, r, result); } public static sub< H extends number, W extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: Matrix, r: Matrix, result: Matrix = new Matrix(l._device, l._array, l.height, l.width)): Matrix { 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>, >( lis: Vector, m: Matrix, result: Matrix = new Matrix(m._device, m._array, m.height, m.width), ): Matrix { 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> = NonNullable>, >(l: Scalar, r: Matrix, result: Matrix = new Matrix(r._device, r._array, r.height, r.width)): Matrix { 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>, >( l: Matrix, r: Matrix, result: Matrix = new Matrix(l._device, l._array, l.height, r.width), ): Matrix { 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 transpose< H extends number, W extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable>, >( m: Matrix, result: Matrix = new Matrix(m._device, m._array, m.width, m.height), ): Matrix { 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>, >( l: L, m: Matrix, result: MatrixBatch = new MatrixBatch(m._device, m._array, l, m.height, m.width), ): MatrixBatch { this.assertUnique(m, result); this.assertColocated(m, result); this.assertPackable(4, m); return result.calculate('unavg', SHADER_UNAVG, [m], false, 4); } }