import fs from 'node:fs'; import { Vector as CPUVector } from '../math'; import { GPUStruct } from './gpu-struct'; import { Scalar } from './scalar'; const SHADER_AVG = fs.readFileSync('shader/avg.wgsl').toString(); const SHADER_TOT = fs.readFileSync('shader/tot.wgsl').toString(); export class ScalarBatch< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, > extends GPUStruct { public get layers(): L { return this.size[2] as L; } protected override workgroups(): [number, number, number] { return [1, 1, this.layers / GPUStruct.lunit]; } protected override workgroup(): [number, number, number] { return [1, 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, data?: CPUVector); public constructor(device: GPUDevice, array: new (size: number) => A, ...args: unknown[]) { const l: number = typeof args[0] === 'number' ? args[0] as L : (args[0] as CPUVector).length; super(device, array, `[${l}]`, [1, 1, l, 1]); ScalarBatch.assertPackable(1, this); const data = (typeof args[0] === 'number' ? args[1] : args[0]) as CPUVector | undefined; if (data) { this.set(data); } } public async set(l: number, v: T): Promise; public async set(data: CPUVector): Promise; public async set(slice: T[], ol?: number): Promise; public async set(b: ScalarBatch): Promise; public async set(...args: unknown[]): Promise { if (typeof args[0] === 'number') { const [l, v]: number[] = args as number[]; if (this._dirty === 'back') { await this.read(); } this.data[l] = v; this._dirty = 'front'; } else if (Array.isArray(args[0])) { const [data, ol = 0] = args as [T[], number | undefined]; if (this._dirty === 'back' && (ol || data.length < this.layers)) { await this.read(); } let l = ol; for (const v of data ?? []) { if (l >= this.layers) { break; } this.data[l] = v; l++; } this._dirty = 'front'; } else { const [batch] = args as [ScalarBatch]; 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): Promise; public async get(): Promise>; public async get(l?: number): Promise> { if (this._dirty === 'back') { await this.read(); } if (typeof l === 'number') { return this.data[l] as T; } else { const result: CPUVector = [] as unknown as CPUVector; for (let l: number = 0; l < this.layers; l++) { result.push(this.data[l] as T); } return result; } } protected static override assertPackable(_: number, ...ms: ScalarBatch[]): void { for (const m of ms) { if (m.layers % ScalarBatch.dunit !== 0) { throw new Error(`Batch size [${m.layers}] must be divisible by [${ScalarBatch.dunit}].`); } } } public static sum< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: ScalarBatch, r: ScalarBatch, result: ScalarBatch = new ScalarBatch(l._device, l._array, l.layers)): ScalarBatch { return GPUStruct.csum(l, r, result); } public static sub< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: ScalarBatch, r: ScalarBatch, result: ScalarBatch = new ScalarBatch(l._device, l._array, l.layers)): ScalarBatch { return GPUStruct.csub(l, r, result); } public static mul< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: ScalarBatch, r: ScalarBatch, result: ScalarBatch = new ScalarBatch(l._device, l._array, l.layers)): ScalarBatch { return GPUStruct.cmul(l, r, result); } public static div< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: ScalarBatch, r: ScalarBatch, result: ScalarBatch = new ScalarBatch(l._device, l._array, l.layers)): ScalarBatch { return GPUStruct.cdiv(l, r, result); } public static scale< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable> = NonNullable>, >(l: Scalar, r: ScalarBatch, result: ScalarBatch = new ScalarBatch(r._device, r._array, r.layers)): ScalarBatch { return GPUStruct.cscale(l, r, result); } public static avg< L extends number, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable>, >( m: ScalarBatch, result: Scalar = new Scalar(m._device, m._array), ): Scalar { 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, A extends Uint32Array | Int32Array | Float32Array, T extends NonNullable>, >( m: ScalarBatch, result: Scalar = new Scalar(m._device, m._array), ): Scalar { this.assertUnique(m, result); this.assertColocated(m, result); this.assertPackable(4, m); return result.calculate('tot', SHADER_TOT, [m], false, 4); } }