Files
nn/src/data/gpu/vector.ts
T
2024-07-27 14:39:08 +00:00

165 lines
5.5 KiB
TypeScript

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