diff --git a/packages/typegpu/src/data/compiledIO.ts b/packages/typegpu/src/data/compiledIO.ts index 117ab03682..2a5ba6c017 100644 --- a/packages/typegpu/src/data/compiledIO.ts +++ b/packages/typegpu/src/data/compiledIO.ts @@ -233,14 +233,19 @@ export function buildWriter( const components = ['x', 'y', 'z', 'w']; const wgslElementStride = rowStride / 4; - return Array.from({ length: matSize * matSize }, (_, idx) => { - const col = Math.floor(idx / matSize); - const row = idx % matSize; + return Array.from({ length: matSize * wgslElementStride }, (_, idx) => { + const col = Math.floor(idx / wgslElementStride); + const row = idx % wgslElementStride; + const offsetAddition = col * rowStride + row * 4; + if (row >= matSize) { + // mat3x3f padding + return emitWrite(writeFunc, `(${offsetExpr} + ${offsetAddition})`, '0'); + } const packedIndex = col * matSize + row; const wgslIndex = col * wgslElementStride + row; return emitWrite( writeFunc, - `(${offsetExpr} + ${col * rowStride + row * 4})`, + `(${offsetExpr} + ${offsetAddition})`, `ArrayBuffer.isView(${valueExpr}) ? ${valueExpr}[${ wgslIndex }] : Array.isArray(${valueExpr}) ? ${valueExpr}[${packedIndex}] : ${valueExpr}.columns[${ diff --git a/packages/typegpu/src/data/dataIO.ts b/packages/typegpu/src/data/dataIO.ts index 23c89a08d2..75dc5f23da 100644 --- a/packages/typegpu/src/data/dataIO.ts +++ b/packages/typegpu/src/data/dataIO.ts @@ -25,8 +25,7 @@ import { vec4i, vec4u, } from './vector.ts'; -import type * as wgsl from './wgslTypes.ts'; -import { isWgslArray, type BaseData } from './wgslTypes.ts'; +import * as wgsl from './wgslTypes.ts'; import type { BufferWriteOptions } from '../core/buffer/buffer.ts'; import { getCompiledWriter } from './compiledIO.ts'; import { getName } from '../shared/meta.ts'; @@ -37,7 +36,7 @@ import { readFloat16, writeFloat16 } from './float16Conversion.ts'; type DataWriter = ( output: ISerialOutput, schema: TSchema, - value: Infer, + value: InferInput, ) => void; type DataReader = ( @@ -58,42 +57,42 @@ const dataWriters = { throw new Error('Booleans are not host-shareable'); }, - f32(output, _schema: wgsl.F32, value: number) { + f32(output, _schema, value) { output.writeFloat32(value); }, - f16(output, _schema: wgsl.F16, value: number) { + f16(output, _schema, value) { writeFloat16(output, value); }, - i32(output, _schema: wgsl.I32, value: number) { + i32(output, _schema, value) { output.writeInt32(value); }, - u32(output, _schema: wgsl.U32, value: number) { + u32(output, _schema, value) { output.writeUint32(value); }, - u16(output, _schema: wgsl.U16, value: number) { + u16(output, _schema, value) { output.writeUint16(value); }, - vec2f(output, _, value: wgsl.v2f) { + vec2f(output, _, value) { output.writeFloat32(value[0]); output.writeFloat32(value[1]); }, - vec2h(output, _, value: wgsl.v2h) { + vec2h(output, _, value) { writeFloat16(output, value[0]); writeFloat16(output, value[1]); }, - vec2i(output, _, value: wgsl.v2i) { + vec2i(output, _, value) { output.writeInt32(value[0]); output.writeInt32(value[1]); }, - vec2u(output, _, value: wgsl.v2u) { + vec2u(output, _, value) { output.writeUint32(value[0]); output.writeUint32(value[1]); }, @@ -102,25 +101,25 @@ const dataWriters = { throw new Error('Booleans are not host-shareable'); }, - vec3f(output, _, value: wgsl.v3f) { + vec3f(output, _, value) { output.writeFloat32(value[0]); output.writeFloat32(value[1]); output.writeFloat32(value[2]); }, - vec3h(output, _, value: wgsl.v3h) { + vec3h(output, _, value) { writeFloat16(output, value[0]); writeFloat16(output, value[1]); writeFloat16(output, value[2]); }, - vec3i(output, _, value: wgsl.v3i) { + vec3i(output, _, value) { output.writeInt32(value[0]); output.writeInt32(value[1]); output.writeInt32(value[2]); }, - vec3u(output, _, value: wgsl.v3u) { + vec3u(output, _, value) { output.writeUint32(value[0]); output.writeUint32(value[1]); output.writeUint32(value[2]); @@ -130,28 +129,28 @@ const dataWriters = { throw new Error('Booleans are not host-shareable'); }, - vec4f(output, _, value: wgsl.v4f) { + vec4f(output, _, value) { output.writeFloat32(value[0]); output.writeFloat32(value[1]); output.writeFloat32(value[2]); output.writeFloat32(value[3]); }, - vec4h(output, _, value: wgsl.v4h) { + vec4h(output, _, value) { writeFloat16(output, value[0]); writeFloat16(output, value[1]); writeFloat16(output, value[2]); writeFloat16(output, value[3]); }, - vec4i(output, _, value: wgsl.v4i) { + vec4i(output, _, value) { output.writeInt32(value[0]); output.writeInt32(value[1]); output.writeInt32(value[2]); output.writeInt32(value[3]); }, - vec4u(output, _, value: wgsl.v4u) { + vec4u(output, _, value) { output.writeUint32(value[0]); output.writeUint32(value[1]); output.writeUint32(value[2]); @@ -162,25 +161,34 @@ const dataWriters = { throw new Error('Booleans are not host-shareable'); }, - mat2x2f(output, _, value: wgsl.m2x2f) { + mat2x2f(output, _, value) { for (let i = 0; i < value.length; ++i) { output.writeFloat32(value[i] as number); } }, - mat3x3f(output, _, value: wgsl.m3x3f) { - for (let i = 0; i < value.length; ++i) { - output.writeFloat32(value[i] as number); + mat3x3f(output, _, value) { + if (value instanceof Float32Array || wgsl.isMatInstance(value)) { + for (let i = 0; i < value.length; ++i) { + output.writeFloat32(value[i] as number); + } + } else { + for (let i = 0; i < 3; ++i) { + for (let j = 0; j < 3; ++j) { + output.writeFloat32(value[3 * i + j] as number); + } + output.writeFloat32(0); + } } }, - mat4x4f(output, _, value: wgsl.m4x4f) { + mat4x4f(output, _, value) { for (let i = 0; i < value.length; ++i) { output.writeFloat32(value[i] as number); } }, - struct(output, schema: wgsl.WgslStruct, value: InferRecord>) { + struct(output, schema, value) { const alignment = alignmentOf(schema); alignIO(output, alignment); @@ -192,7 +200,7 @@ const dataWriters = { alignIO(output, alignment); }, - array(output, schema: wgsl.WgslArray, value: Infer[]) { + array(output, schema, value) { if (schema.elementCount === 0) { throw new Error('Cannot write using a runtime-sized schema.'); } @@ -211,11 +219,11 @@ const dataWriters = { throw new Error('Pointers are not host-shareable'); }, - atomic(output, schema: wgsl.Atomic, value: number) { + atomic(output, schema, value) { dataWriters[schema.inner.type]?.(output, schema, value); }, - decorated(output, schema: wgsl.Decorated, value: unknown) { + decorated(output, schema, value) { const alignment = customAlignmentOf(schema); alignIO(output, alignment); @@ -226,178 +234,178 @@ const dataWriters = { // Loose Types - uint8(output, _, value: number) { + uint8(output, _, value) { output.writeUint8(value); }, - uint8x2(output, _, value: wgsl.v2u) { + uint8x2(output, _, value) { output.writeUint8(value.x); output.writeUint8(value.y); }, - uint8x4(output, _, value: wgsl.v4u) { + uint8x4(output, _, value) { output.writeUint8(value.x); output.writeUint8(value.y); output.writeUint8(value.z); output.writeUint8(value.w); }, - sint8(output, _, value: number) { + sint8(output, _, value) { output.writeInt8(value); }, - sint8x2(output, _, value: wgsl.v2i) { + sint8x2(output, _, value) { output.writeInt8(value.x); output.writeInt8(value.y); }, - sint8x4(output, _, value: wgsl.v4i) { + sint8x4(output, _, value) { output.writeInt8(value.x); output.writeInt8(value.y); output.writeInt8(value.z); output.writeInt8(value.w); }, - unorm8(output, _, value: number) { + unorm8(output, _, value) { output.writeUint8(Math.round(value * 255)); }, - unorm8x2(output, _, value: wgsl.v2f) { + unorm8x2(output, _, value) { output.writeUint8(Math.round(value.x * 255)); output.writeUint8(Math.round(value.y * 255)); }, - unorm8x4(output, _, value: wgsl.v4f) { + unorm8x4(output, _, value) { output.writeUint8(Math.round(value.x * 255)); output.writeUint8(Math.round(value.y * 255)); output.writeUint8(Math.round(value.z * 255)); output.writeUint8(Math.round(value.w * 255)); }, - snorm8(output, _, value: number) { + snorm8(output, _, value) { output.writeInt8(Math.round(value * 127)); }, - snorm8x2(output, _, value: wgsl.v2f) { + snorm8x2(output, _, value) { output.writeInt8(Math.round(value.x * 127)); output.writeInt8(Math.round(value.y * 127)); }, - snorm8x4(output, _, value: wgsl.v4f) { + snorm8x4(output, _, value) { output.writeInt8(Math.round(value.x * 127)); output.writeInt8(Math.round(value.y * 127)); output.writeInt8(Math.round(value.z * 127)); output.writeInt8(Math.round(value.w * 127)); }, - uint16(output, _, value: number) { + uint16(output, _, value) { output.writeUint16(value); }, - uint16x2(output, _, value: wgsl.v2u) { + uint16x2(output, _, value) { output.writeUint16(value.x); output.writeUint16(value.y); }, - uint16x4(output, _, value: wgsl.v4u) { + uint16x4(output, _, value) { output.writeUint16(value.x); output.writeUint16(value.y); output.writeUint16(value.z); output.writeUint16(value.w); }, - sint16(output, _, value: number) { + sint16(output, _, value) { output.writeInt16(value); }, - sint16x2(output, _, value: wgsl.v2i) { + sint16x2(output, _, value) { output.writeInt16(value.x); output.writeInt16(value.y); }, - sint16x4(output, _, value: wgsl.v4i) { + sint16x4(output, _, value) { output.writeInt16(value.x); output.writeInt16(value.y); output.writeInt16(value.z); output.writeInt16(value.w); }, - unorm16(output, _, value: number) { + unorm16(output, _, value) { output.writeUint16(value * 65535); }, - unorm16x2(output, _, value: wgsl.v2f) { + unorm16x2(output, _, value) { output.writeUint16(value.x * 65535); output.writeUint16(value.y * 65535); }, - unorm16x4(output, _, value: wgsl.v4f) { + unorm16x4(output, _, value) { output.writeUint16(value.x * 65535); output.writeUint16(value.y * 65535); output.writeUint16(value.z * 65535); output.writeUint16(value.w * 65535); }, - snorm16(output, _, value: number) { + snorm16(output, _, value) { output.writeInt16(Math.round(value * 32767)); }, - snorm16x2(output, _, value: wgsl.v2f) { + snorm16x2(output, _, value) { output.writeInt16(Math.round(value.x * 32767)); output.writeInt16(Math.round(value.y * 32767)); }, - snorm16x4(output, _, value: wgsl.v4f) { + snorm16x4(output, _, value) { output.writeInt16(Math.round(value.x * 32767)); output.writeInt16(Math.round(value.y * 32767)); output.writeInt16(Math.round(value.z * 32767)); output.writeInt16(Math.round(value.w * 32767)); }, - float16(output, _, value: number) { + float16(output, _, value) { writeFloat16(output, value); }, - float16x2(output, _, value: wgsl.v2f) { + float16x2(output, _, value) { writeFloat16(output, value.x); writeFloat16(output, value.y); }, - float16x4(output, _, value: wgsl.v4f) { + float16x4(output, _, value) { writeFloat16(output, value.x); writeFloat16(output, value.y); writeFloat16(output, value.z); writeFloat16(output, value.w); }, - float32(output, _, value: number) { + float32(output, _, value) { output.writeFloat32(value); }, - float32x2(output, _, value: wgsl.v2f) { + float32x2(output, _, value) { output.writeFloat32(value.x); output.writeFloat32(value.y); }, - float32x3(output, _, value: wgsl.v3f) { + float32x3(output, _, value) { output.writeFloat32(value.x); output.writeFloat32(value.y); output.writeFloat32(value.z); }, - float32x4(output, _, value: wgsl.v4f) { + float32x4(output, _, value) { output.writeFloat32(value.x); output.writeFloat32(value.y); output.writeFloat32(value.z); output.writeFloat32(value.w); }, - uint32(output, _, value: number) { + uint32(output, _, value) { output.writeUint32(value); }, - uint32x2(output, _, value: wgsl.v2u) { + uint32x2(output, _, value) { output.writeUint32(value.x); output.writeUint32(value.y); }, - uint32x3(output, _, value: wgsl.v3u) { + uint32x3(output, _, value) { output.writeUint32(value.x); output.writeUint32(value.y); output.writeUint32(value.z); }, - uint32x4(output, _, value: wgsl.v4u) { + uint32x4(output, _, value) { output.writeUint32(value.x); output.writeUint32(value.y); output.writeUint32(value.z); output.writeUint32(value.w); }, - sint32(output, _, value: number) { + sint32(output, _, value) { output.writeInt32(value); }, - sint32x2(output, _, value: wgsl.v2i) { + sint32x2(output, _, value) { output.writeInt32(value.x); output.writeInt32(value.y); }, - sint32x3(output, _, value: wgsl.v3i) { + sint32x3(output, _, value) { output.writeInt32(value.x); output.writeInt32(value.y); output.writeInt32(value.z); }, - sint32x4(output, _, value: wgsl.v4i) { + sint32x4(output, _, value) { output.writeInt32(value.x); output.writeInt32(value.y); output.writeInt32(value.z); output.writeInt32(value.w); }, - 'unorm10-10-10-2'(output, _, value: wgsl.v4f) { + 'unorm10-10-10-2'(output, _, value) { let packed = 0; packed |= ((value.x * 1023) & 1023) << 22; // r (10 bits) packed |= ((value.y * 1023) & 1023) << 12; // g (10 bits) @@ -405,14 +413,14 @@ const dataWriters = { packed |= (value.w * 3) & 3; // a (2 bits) output.writeUint32(packed); }, - 'unorm8x4-bgra'(output, _, value: wgsl.v4f) { + 'unorm8x4-bgra'(output, _, value) { output.writeUint8(value.z * 255); output.writeUint8(value.y * 255); output.writeUint8(value.x * 255); output.writeUint8(value.w * 255); }, - disarray(output, schema: Disarray, value: unknown[]) { + disarray(output, schema, value) { const alignment = alignmentOf(schema); alignIO(output, alignment); @@ -432,7 +440,7 @@ const dataWriters = { } }, - 'loose-decorated'(output, schema: LooseDecorated, value: unknown) { + 'loose-decorated'(output, schema, value) { const alignment = customAlignmentOf(schema); alignIO(output, alignment); @@ -451,7 +459,7 @@ const dataWriters = { export function writeData( output: ISerialOutput, schema: TData, - value: Infer, + value: InferInput, ): void { if (ArrayBuffer.isView(value)) { const src = value as ArrayBufferView; @@ -800,7 +808,7 @@ export function readData( const endianness = getSystemEndianness(); -export function calculateOffsets( +export function calculateOffsets( options: BufferWriteOptions | undefined, schema: T, data: InferInput | ArrayBuffer | ArrayBufferView, @@ -808,7 +816,7 @@ export function calculateOffsets( const bufferSize = sizeOf(schema); const startOffset = options?.startOffset ?? 0; let naturalSize: number | undefined = undefined; - if (isWgslArray(schema) && Array.isArray(data)) { + if (wgsl.isWgslArray(schema) && Array.isArray(data)) { const arrayData = data as unknown[]; naturalSize = arrayData.length * roundUp(sizeOf(schema.elementType), alignmentOf(schema.elementType)); @@ -823,7 +831,12 @@ export function calculateOffsets( return { startOffset, endOffset }; } -export function writeToArrayBuffer( +/** + * A function for filling in buffers with data based on TypeGPU schemas. + * When schema instances or JS arrays are passed, the padding is applied automatically. + * ArrayBuffers are expected to be already padded. + */ +export function writeToArrayBuffer( buffer: ArrayBuffer, schema: T, data: InferInput | ArrayBuffer | ArrayBufferView, @@ -872,9 +885,12 @@ export function writeToArrayBuffer( const writer = new BufferWriter(buffer); writer.seekTo(startOffset); - writeData(writer, schema, data as Infer); + writeData(writer, schema, data); } -export function readFromArrayBuffer(buffer: ArrayBuffer, schema: T): Infer { +export function readFromArrayBuffer( + buffer: ArrayBuffer, + schema: T, +): Infer { return readData(new BufferReader(buffer), schema); } diff --git a/packages/typegpu/tests/arrayBufferIO.test.ts b/packages/typegpu/tests/arrayBufferIO.test.ts index c0de7ecc42..a72f0996e3 100644 --- a/packages/typegpu/tests/arrayBufferIO.test.ts +++ b/packages/typegpu/tests/arrayBufferIO.test.ts @@ -77,6 +77,19 @@ describe('arrayBufferIO', () => { expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0]); }); + it('overrides all values when handling mat3x3f', () => { + const buffer = new ArrayBuffer(48); + writeToArrayBuffer( + buffer, + d.arrayOf(d.u32, 12), + [-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1], + ); + + writeToArrayBuffer(buffer, d.mat3x3f, [1, 2, 3, 4, 5, 6, 7, 8, 9]); + + expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0]); + }); + it('expects a particular number of elements for vectors and matrices', () => { const buffer = new ArrayBuffer(48); diff --git a/packages/typegpu/tests/internal/arrayBufferIO.noEval.test.ts b/packages/typegpu/tests/internal/arrayBufferIO.noEval.test.ts new file mode 100644 index 0000000000..5a5007be68 --- /dev/null +++ b/packages/typegpu/tests/internal/arrayBufferIO.noEval.test.ts @@ -0,0 +1,137 @@ +import { describe, expect, vi } from 'vitest'; +import { d, patchArrayBuffer, readFromArrayBuffer, writeToArrayBuffer } from 'typegpu'; +import { it } from 'typegpu-testing-utility'; + +vi.mock('../../src/data/compiledIO.ts', () => ({ + EVAL_ALLOWED_IN_ENV: false, + getCompiledWriter: () => undefined, + buildWriter: () => '', +})); + +describe('arrayBufferIO', () => { + it('handles d.vec input', () => { + const buffer = new ArrayBuffer(16); + + writeToArrayBuffer(buffer, d.vec4u, d.vec4u(1, 2, 3, 4)); + + expect([...new Uint32Array(buffer)]).toStrictEqual([1, 2, 3, 4]); + }); + + it('handles plain array input', () => { + const buffer = new ArrayBuffer(16); + + writeToArrayBuffer(buffer, d.vec4u, [1, 2, 3, 4]); + + expect([...new Uint32Array(buffer)]).toStrictEqual([1, 2, 3, 4]); + }); + + it('handles typed array input', () => { + const buffer = new ArrayBuffer(16); + + writeToArrayBuffer(buffer, d.vec4u, new Uint32Array([1, 2, 3, 4])); + + expect([...new Uint32Array(buffer)]).toStrictEqual([1, 2, 3, 4]); + }); + + it('handles ArrayBuffer input', () => { + const buffer = new ArrayBuffer(16); + + writeToArrayBuffer(buffer, d.vec4u, new Uint32Array([1, 2, 3, 4]).buffer); + + expect([...new Uint32Array(buffer)]).toStrictEqual([1, 2, 3, 4]); + }); + + it('respects startOffset', () => { + const buffer = new Uint32Array([10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10]); + const Numbers = d.arrayOf(d.u32, 12); + const layout = d.memoryLayoutOf(Numbers, (a) => a[4]); + + writeToArrayBuffer(buffer.buffer, Numbers, [1, 2, 3, 4], { startOffset: layout.offset }); + + expect([...new Uint32Array(buffer)]).toStrictEqual([ + 10, 10, 10, 10, 1, 2, 3, 4, 10, 10, 10, 10, + ]); + }); + + it('handles structs', () => { + const buffer = new ArrayBuffer(32); + const Boid = d.struct({ pos: d.vec2u, id: d.u32 }); + const Boids = d.arrayOf(Boid, 2); + + writeToArrayBuffer(buffer, Boids, [ + Boid({ pos: d.vec2u(1, 2), id: 3 }), + Boid({ pos: d.vec2u(4, 5), id: 6 }), + ]); + + expect([...new Uint32Array(buffer)]).toStrictEqual([1, 2, 3, 0, 4, 5, 6, 0]); + }); + + it('handles vec3f', () => { + const buffer = new ArrayBuffer(16); + + writeToArrayBuffer(buffer, d.vec3f, [1, 2, 3]); + + expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0]); + }); + + it('does not override 4th element when handling vec3f', () => { + const buffer = new ArrayBuffer(16); + writeToArrayBuffer(buffer, d.vec4f, [1, 2, 3, 4]); + + writeToArrayBuffer(buffer, d.vec3f, [6, 7, 8]); + + expect([...new Float32Array(buffer)]).toEqual([6, 7, 8, 4]); + }); + + it('handles mat3x3f with array input', () => { + const buffer = new ArrayBuffer(48); + + writeToArrayBuffer(buffer, d.mat3x3f, [1, 2, 3, 4, 5, 6, 7, 8, 9]); + + expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0]); + }); + + it('handles mat3x3f with matrix input', () => { + const buffer = new ArrayBuffer(48); + + writeToArrayBuffer(buffer, d.mat3x3f, d.mat3x3f(1, 2, 3, 4, 5, 6, 7, 8, 9)); + + expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0]); + }); + + it('handles mat3x3f with float32Array input', () => { + const buffer = new ArrayBuffer(48); + + writeToArrayBuffer(buffer, d.mat3x3f, new Float32Array([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0])); + + expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0]); + }); + + it('overrides all values when handling mat3x3f', () => { + const buffer = new ArrayBuffer(48); + writeToArrayBuffer( + buffer, + d.arrayOf(d.u32, 12), + [-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1], + ); + + writeToArrayBuffer(buffer, d.mat3x3f, [1, 2, 3, 4, 5, 6, 7, 8, 9]); + + expect([...new Float32Array(buffer)]).toEqual([1, 2, 3, 0, 4, 5, 6, 0, 7, 8, 9, 0]); + }); + + it('handles structs with attributes', () => { + const buffer = new ArrayBuffer(64); + const Boid = d.struct({ pos: d.vec2u, id: d.align(16, d.u32) }); + const Boids = d.arrayOf(Boid, 2); + + writeToArrayBuffer(buffer, Boids, [ + Boid({ pos: d.vec2u(1, 2), id: 3 }), + Boid({ pos: d.vec2u(4, 5), id: 6 }), + ]); + + expect([...new Uint32Array(buffer)]).toStrictEqual([ + 1, 2, 0, 0, 3, 0, 0, 0, 4, 5, 0, 0, 6, 0, 0, 0, + ]); + }); +});