import fs from 'node:fs'; import type { Scalar } from './scalar'; interface ComputePipeline { pipeline: GPUComputePipeline; layout: GPUBindGroupLayout; } const PIPELINES: Map> = new Map(); const COMPWISE_SHADER = { sum: fs.readFileSync('shader/csum.wgsl').toString(), sub: fs.readFileSync('shader/csub.wgsl').toString(), mul: fs.readFileSync('shader/cmul.wgsl').toString(), div: fs.readFileSync('shader/cdiv.wgsl').toString(), scale: fs.readFileSync('shader/cscale.wgsl').toString(), }; export function createComputePipeline( device: GPUDevice, name: string, shader: () => string, entry: string = 'main', ) { if (!PIPELINES.has(device)) { PIPELINES.set(device, new Map()); } const d = PIPELINES.get(device)!; if (!d.has(name)) { const pipeline = device.createComputePipeline({ label: name, layout: 'auto', compute: { module: device.createShaderModule({ code: shader() }), entryPoint: entry, }, }); const layout = pipeline.getBindGroupLayout(0); d.set(name, { pipeline, layout }); } return d.get(name)!; } export abstract class GPUStruct { public static dunit = 8; public static lunit = 2; protected readonly _label: string; protected readonly _device: GPUDevice; protected readonly _array: new (size: number | ArrayBufferLike | ArrayLike) => A; protected readonly _inbuffer: GPUBuffer; protected readonly _dummybuffer: GPUBuffer; protected readonly _outbuffer: GPUBuffer; protected _op: 'reading' | 'writing' | null = null; protected _dirty: 'front' | 'back' | null = 'front'; protected abstract workgroups(pm?: number): [number, number, number]; protected abstract workgroup(): [number, number, number]; public readonly size: Int32Array; public readonly data: A; public get type(): 'u32' | 'i32' | 'f32' { switch (this._array as unknown) { case Uint32Array: return 'u32'; case Int32Array: return 'i32'; case Float32Array: return 'f32'; default: throw new Error('Unknown array type.'); } } public get length(): number { return this.data.length; } public get byteLength(): number { return this.size.byteLength + this.data.byteLength; } public constructor( device: GPUDevice, array: new (size: number) => A, label: string, size: [number, number, number, number], ) { this._device = device; this._array = array; this._label = label ?? 'struct'; this.size = new Int32Array(size); this.data = new array(size[0] * size[1] * size[2] * size[3]); this._inbuffer = device.createBuffer({ label: `${this._label}in`, size: Math.max(32, this.byteLength), usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST, }); this._dummybuffer = device.createBuffer({ label: `${label}dummy`, size: 0, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST, }); this._outbuffer = this._device.createBuffer({ label: `${label}out`, size: Math.max(32, this.byteLength), usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ, }); } public write(): void { if (!this._dirty) { return; } if (this._op) { throw new Error('Sync in progress.'); } this._op = 'writing'; this._device.queue.writeBuffer(this._inbuffer, 0, new Uint8Array(this.size.buffer)); this._device.queue.writeBuffer(this._inbuffer, this.size.byteLength, new Uint8Array(this.data.buffer)); this._dirty = null; this._op = null; } public async read(): Promise { if (!this._dirty) { return; } if (this._op) { throw new Error('Sync in progress.'); } this._op = 'reading'; const copy = this._device.createCommandEncoder(); copy.copyBufferToBuffer(this._inbuffer, 0, this._outbuffer, 0, this.byteLength); this._device.queue.submit([copy.finish()]); await this._device.queue.onSubmittedWorkDone(); await this._outbuffer.mapAsync(GPUMapMode.READ); this.data.set(new this._array(this._outbuffer.getMappedRange()).slice(this.size.length)); this._outbuffer.unmap(); this._dirty = null; this._op = null; } public abstract get(): Promise; public buffer(): GPUBuffer { if (this._dirty === 'front') { this.write(); } return this._inbuffer; } protected workgroupify(shader: string): string { const [x, y, z] = this.workgroup(); if (!x || !y || !z) { throw new Error(`Invalid workgroup size: (${x}, ${y}, ${z}).`); } return shader.replace(/@workgroup_size\(.*?\)/g, `@workgroup_size(${x}, ${y}, ${z})`); } protected inplacify(shader: string): string { let r = ''; return shader.replace(/\w+_or_result/g, v => r = r ? 'result' : v); // The first occurrence must be the declaration. } protected typify(shader: string): string { return shader.replace(/alias number = .+?;/g, `alias number = ${this.type};`); } public calculate( name: string, shader: string, args: GPUStruct[], ip = false, pm = 1, ): typeof this { const encoder = this._device.createCommandEncoder(); const pass = encoder.beginComputePass(); const { pipeline, layout } = createComputePipeline(this._device, `${name}${ip ? '_inplace_' : '_'}${this.type}`, () => this.workgroupify(this.typify(ip ? this.inplacify(shader) : shader))); pass.setPipeline(pipeline); pass.setBindGroup(0, this._device.createBindGroup({ layout, entries: [ ...args.map((s, i) => ({ binding: i, resource: { buffer: s.buffer() } })), { binding: args.length + +ip, resource: { buffer: this.buffer() } }, ], })); const [x, y, z] = this.workgroups(pm); if (!x || !y || !z || Math.floor(x) !== x || Math.floor(y) !== y || Math.floor(z) !== z) { throw new Error(`Invalid workgroup count: (${x}, ${y}, ${z}).`); } pass.dispatchWorkgroups(x, y, z); pass.end(); this._device.queue.submit([encoder.finish()]); this._dirty = 'back'; return this; } protected static assertUnique(...ms: GPUStruct[]): void { const map: Map, boolean> = new Map(); for (const m of ms) { if (map.has(m)) { throw new Error('Can not use the same GPU structure multiple times.'); } map.set(m, true); } } protected static assertColocated(...ms: GPUStruct[]): void { const map: Map = new Map(); for (const m of ms) { map.set(m._device, true); if (map.size > 1) { throw new Error('Can not use GPU structures stored on multiple devices.'); } } } protected static assertSimilar(...ms: GPUStruct[]): void { for (const m of ms) { if (m.length !== ms[0].length || m.workgroup()[0] !== ms[0].workgroup()[0] || m.workgroup()[1] !== ms[0].workgroup()[1] || m.workgroup()[2] !== ms[0].workgroup()[2]) { throw new Error('Can not use GPU structures with different dimensions.'); } } } protected static assertPackable(pm: number, ...ms: GPUStruct[]): void { for (const m of ms) { if (m.workgroups(pm).some(v => v < 1)) { throw new Error(`Struct ${m._label} size must be divisible by ${pm}.`); } } } public static compwise< A extends Uint32Array | Int32Array | Float32Array, S extends GPUStruct, >(op: keyof typeof COMPWISE_SHADER, l: S, r: S, result: S): S { this.assertUnique(l, r); this.assertUnique(l, result); this.assertColocated(l, r, result); this.assertPackable(4, l, r, result); return result.calculate(op, COMPWISE_SHADER[op], result === r ? [l] : [l, r], result === r, 4); } public static csum< A extends Uint32Array | Int32Array | Float32Array, S extends GPUStruct, >(l: S, r: S, result: S): S { return GPUStruct.compwise('sum', l, r, result); } public static csub< A extends Uint32Array | Int32Array | Float32Array, S extends GPUStruct, >(l: S, r: S, result: S): S { return GPUStruct.compwise('sub', l, r, result); } public static cmul< A extends Uint32Array | Int32Array | Float32Array, S extends GPUStruct, >(l: S, r: S, result: S): S { return GPUStruct.compwise('mul', l, r, result); } public static cdiv< A extends Uint32Array | Int32Array | Float32Array, S extends GPUStruct, >(l: S, r: S, result: S): S { return GPUStruct.compwise('div', l, r, result); } public static cscale< A extends Uint32Array | Int32Array | Float32Array, S extends GPUStruct, >(l: Scalar, r: S, result: S): S { this.assertUnique(l, r); this.assertUnique(l, result); this.assertColocated(l, r, result); this.assertPackable(4, r, result); return result.calculate('scale', COMPWISE_SHADER.scale, result === r ? [l] : [l, r], result === r, 4); } }