import fs from 'node:fs'; import { Vector as CPUVector } from '../math'; import { GPUStruct } from './gpu-struct'; import { ScalarBatch } from './scalar-batch'; import { Vector } from './vector'; const SHADER_AVG = fs.readFileSync('shader/avg.wgsl').toString(); const SHADER_TOT = fs.readFileSync('shader/tot.wgsl').toString(); const SHADER_SCALE = fs.readFileSync('shader/bscale.wgsl').toString(); const SHADER_DOT = fs.readFileSync('shader/mul.wgsl').toString(); export class VectorBatch< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, > extends GPUStruct { public get dims(): D { return this.size[0] as D; } public get layers(): L { return this.size[2] as L; } protected override workgroups(pm: number = 1): [number, number, number] { return [this.dims / GPUStruct.dunit / pm, 1, this.layers / GPUStruct.lunit]; } protected override workgroup(): [number, number, number] { return [GPUStruct.dunit, 1, GPUStruct.lunit]; } public constructor(device: GPUDevice, array: new (size: number) => A, data: CPUVector>); public constructor(device: GPUDevice, array: new (size: number) => A, l: L, d: D, data?: CPUVector>); public constructor(device: GPUDevice, array: new (size: number) => A, ...args: unknown[]) { let [l, d]: [number, number] = [1, 1]; if (typeof args[0] === 'number') { [l, d] = args as [L, D]; } else { const [data] = args as [CPUVector>]; [l, d] = [data.length, data[0].length]; } super(device, array, `${d}[${l}]`, [d, 1, l, 1]); VectorBatch.assertPackable(1, this); const data = (typeof args[0] === 'number' ? args[2] : args[0]) as CPUVector> | undefined; if (data) { this.set(data); } } public index(l: number, i: number): number { return l * this.dims + i; } public async set(l: number, i: number, v: T): Promise; public async set(data: CPUVector>): Promise; public async set(slice: T[][], ol?: number, oi?: number): Promise; public async set(b: VectorBatch): Promise; public async set(...args: unknown[]): Promise { if (typeof args[0] === 'number') { const [l, i, v]: number[] = args as number[]; if (this._dirty === 'back') { await this.read(); } this.data[this.index(l, i)] = v; this._dirty = 'front'; } else if (Array.isArray(args[0])) { const [data, ol = 0, oi = 0] = args as [T[][], number | undefined, number | undefined]; if (this._dirty === 'back' && (ol || oi || data.length < this.layers || data[0].length < this.dims)) { await this.read(); } let l = ol; let i = oi; for (const vec of data ?? []) { if (l >= this.layers) { break; } for (const v of vec) { if (i >= this.dims) { break; } this.data[this.index(l, i)] = v; i++; } i = oi; l++; } this._dirty = 'front'; } else { const [batch] = args as [VectorBatch]; const copy = this._device.createCommandEncoder(); copy.copyBufferToBuffer(batch.buffer(), 0, this.buffer(), 0, this.byteLength); this._device.queue.submit([copy.finish()]); this._dirty = 'back'; } } public async get(l: number, i: number): Promise; public async get(l: number): Promise>; public async get(): Promise>>; public async get(l?: number, i?: number): Promise | CPUVector>> { if (this._dirty === 'back') { await this.read(); } if (typeof i === 'number') { return this.data[this.index(l!, i!)] as T; } else if (typeof l === 'number') { const result: CPUVector = [] as unknown as CPUVector; for (let i: number = 0; i < this.dims; i++) { result.push(this.data[this.index(l, i)] as T); } return result; } else { const result: CPUVector> = [] as unknown as CPUVector>; for (let l: number = 0; l < this.layers; l++) { const vec: CPUVector = [] as unknown as CPUVector; for (let i: number = 0; i < this.dims; i++) { vec.push(this.data[this.index(l, i)] as T); } result.push(vec); } return result; } } protected static override assertPackable(pm: number, ...ms: VectorBatch[]): void { for (const m of ms) { if (m.layers % VectorBatch.lunit !== 0 || m.dims % (VectorBatch.dunit * pm) !== 0) { throw new Error(`Batch size ${m.dims}[${m.layers}] must be divisible by ${VectorBatch.dunit * pm}[${VectorBatch.lunit}].`); } } } public static sum< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: VectorBatch, r: VectorBatch, result: VectorBatch = new VectorBatch(l._device, l._array, l.layers, l.dims)): VectorBatch { return GPUStruct.csum(l, r, result); } public static sum3< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: VectorBatch, m: VectorBatch, r: VectorBatch, result: VectorBatch = new VectorBatch(l._device, l._array, l.layers, l.dims)): VectorBatch { return GPUStruct.csum3(l, m, r, result); } public static sub< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: VectorBatch, r: VectorBatch, result: VectorBatch = new VectorBatch(l._device, l._array, l.layers, l.dims)): VectorBatch { return GPUStruct.csub(l, r, result); } public static mul< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: ScalarBatch, r: VectorBatch, result: VectorBatch = new VectorBatch(r._device, r._array, r.layers, r.dims)): VectorBatch { this.assertUnique(l, r); this.assertUnique(l, result); this.assertColocated(l, r, result); this.assertPackable(4, r, result); return result.calculate('scale', SHADER_SCALE, result === r ? [l] : [l, r], result === r, 4); } public static dot< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: VectorBatch, r: VectorBatch, result: ScalarBatch = new ScalarBatch(l._device, l._array, l.layers)): ScalarBatch { throw new Error('Not implemented.'); this.assertUnique(l, r); this.assertUnique(l, result); this.assertColocated(l, r, result); // this.assertPackable(4, r, result); return result.calculate('dot', SHADER_DOT, [l, r]); } public static avg< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable>, >( m: VectorBatch, result: Vector = new Vector(m._device, m._array, m.dims), ): Vector { this.assertUnique(m, result); this.assertColocated(m, result); this.assertPackable(4, m); return result.calculate('avg', SHADER_AVG, [m], false, 4); } public static tot< L extends number, D extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable>, >( m: VectorBatch, result: Vector = new Vector(m._device, m._array, m.dims), ): Vector { this.assertUnique(m, result); this.assertColocated(m, result); this.assertPackable(4, m); return result.calculate('tot', SHADER_TOT, [m], false, 4); } }