From c77522e4804c0b08ee93624d5b684e2fb5f8a332 Mon Sep 17 00:00:00 2001 From: Iwo Plaza Date: Thu, 6 Aug 2026 11:52:27 +0200 Subject: [PATCH] feat: Route function calls through the shader generator --- packages/typegpu-gl/src/glslGenerator.ts | 95 ++++++++++++++ .../typegpu-gl/tests/glslGenerator.test.ts | 120 ++++++++++++++++++ .../typegpu/src/core/function/dualImpl.ts | 11 +- packages/typegpu/src/std/array.ts | 5 +- packages/typegpu/src/std/bitcast.ts | 22 ++-- packages/typegpu/src/std/boolean.ts | 2 +- packages/typegpu/src/std/numeric.ts | 2 +- packages/typegpu/src/tgsl/shaderGenerator.ts | 3 +- packages/typegpu/src/tgsl/wgslGenerator.ts | 12 ++ 9 files changed, 249 insertions(+), 23 deletions(-) diff --git a/packages/typegpu-gl/src/glslGenerator.ts b/packages/typegpu-gl/src/glslGenerator.ts index 31171ad578..7e9e29d154 100644 --- a/packages/typegpu-gl/src/glslGenerator.ts +++ b/packages/typegpu-gl/src/glslGenerator.ts @@ -8,6 +8,8 @@ import { type ResolutionCtx, type TgpuShaderStage, type FunctionDefinitionOptions, + type Snippet, + snip, } from 'typegpu/~internal'; // ---------- @@ -68,6 +70,21 @@ ${Object.entries(struct.propTypes) return id; } +function correspondingBooleanVectorSchema(dataType: d.BaseData) { + if (dataType.type.includes('2')) { + return d.vec2b; + } + if (dataType.type.includes('3')) { + return d.vec3b; + } + if (dataType.type.includes('4')) { + return d.vec4b; + } + throw new Error( + `Internal error: schema of type '${dataType.type}' does not have a corresponding boolean vector.`, + ); +} + const gl_PositionSnippet = tgpu['~unstable'].rawCodeSnippet('gl_Position', d.vec4f, 'private'); interface EntryFnState { @@ -105,6 +122,84 @@ export class GlslGenerator extends WgslGenerator { return super.typeAnnotation(data); } + override call( + name: string, + templateParams: readonly Snippet[], + args: readonly Snippet[], + ): string { + if (name === 'bitcast') { + const [target] = templateParams; + if (!target || !d.isWgslData(target.value)) { + throw new Error(`Expected bitcast() to be called with a data type template parameter`); + } + const [source] = args; + if (!source || source.dataType === UnknownData) { + throw new Error(`Invalid argument passed to bitcast()`); + } + const targetSchema = target.value; + const sourceSchema = source.dataType; + const targetPrimitive = targetSchema.type.startsWith('vec') + ? (targetSchema as d.Vec3f).primitive + : targetSchema; + const sourcePrimitive = sourceSchema.type.startsWith('vec') + ? (sourceSchema as d.Vec3f).primitive + : sourceSchema; + + if (sourcePrimitive.type === 'u32' && targetPrimitive.type === 'f32') { + return super.call('uintBitsToFloat', [], [source]); + } + if (sourcePrimitive.type === 'i32' && targetPrimitive.type === 'f32') { + return super.call('intBitsToFloat', [], [source]); + } + if (sourcePrimitive.type === 'f32' && targetPrimitive.type === 'u32') { + return super.call('floatBitsToUint', [], [source]); + } + if (sourcePrimitive.type === 'f32' && targetPrimitive.type === 'i32') { + return super.call('floatBitsToInt', [], [source]); + } + if (sourceSchema.type === targetSchema.type) { + return this.ctx.resolveSnippet(source).value; + } + + throw new Error(`Cannot bitcast from ${String(sourceSchema)} to ${String(targetSchema)}`); + } + + if (name === 'select') { + const [falsy, truthy, cond] = args; + if (!falsy || !truthy || !cond) { + throw new Error(`Invalid number of arguments for 'select'`); + } + + if (falsy.dataType !== UnknownData && falsy.dataType.type.startsWith('vec')) { + if (cond.dataType !== UnknownData && cond.dataType.type.startsWith('vec')) { + return super.call('mix', templateParams, args); + } + return super.call('mix', templateParams, [ + falsy, + truthy, + this.typeInstantiation(correspondingBooleanVectorSchema(falsy.dataType), [cond]), + ]); + } + + // Generating a ternary expression, which is supported in GLSL (scalar condition only) + if (cond.dataType !== UnknownData && cond.dataType.type.startsWith('vec')) { + throw new Error(`GLSL select() with scalar branches requires a scalar boolean condition`); + } + + return `(${this.ctx.resolveSnippet(cond).value} ? ${this.ctx.resolveSnippet(truthy).value} : ${this.ctx.resolveSnippet(falsy).value})`; + } + + if (name === 'saturate') { + const [arg] = args; + if (!arg) { + throw new Error(`Invalid number of arguments for 'saturate'`); + } + return super.call('clamp', [], [arg, snip(0, d.f32, 'constant'), snip(1, d.f32, 'constant')]); + } + + return super.call(name, templateParams, args); + } + override _emitVarDecl( _keyword: 'var' | 'let' | 'const', name: string, diff --git a/packages/typegpu-gl/tests/glslGenerator.test.ts b/packages/typegpu-gl/tests/glslGenerator.test.ts index c0f67f7936..b924665bba 100644 --- a/packages/typegpu-gl/tests/glslGenerator.test.ts +++ b/packages/typegpu-gl/tests/glslGenerator.test.ts @@ -97,6 +97,126 @@ describe('GlslGenerator - variable declarations', () => { }); }); +describe('GlslGenerator - standard function calls', () => { + it('translates scalar `select()` to ternary expression', () => { + function foo() { + 'use gpu'; + const cond = false; + return std.select(0, 1, cond); + } + + expect(tgpu.resolve([foo], glOptions())).toMatchInlineSnapshot(` + "int foo() { + bool cond = false; + return (cond ? 1i : 0i); + }" + `); + }); + + it('translates vector `select()` to mix()', () => { + function foo() { + 'use gpu'; + const cond = false; + const vecCond = d.vec3b(false, true, false); + const bar = std.select(d.vec3f(0), d.vec3f(1), cond); // `cond` should be coerced to a boolean vector + const baz = std.select(d.vec3f(1), d.vec3f(0), vecCond); + } + + expect(tgpu.resolve([foo], glOptions())).toMatchInlineSnapshot(` + "void foo() { + bool cond = false; + bvec3 vecCond = bvec3(false, true, false); + vec3 bar = mix(vec3(), vec3(1), bvec3(cond)); + vec3 baz = mix(vec3(1), vec3(), vecCond); + }" + `); + }); + + it('should throw on select() with vector cond and scalar branches', () => { + function foo() { + 'use gpu'; + const cond = d.vec3b(false, true, false); + // @ts-ignore + return std.select(0, 1, cond); + } + + expect(() => tgpu.resolve([foo], glOptions())).toThrowErrorMatchingInlineSnapshot(` + [Error: Resolution of the following tree failed: + - + - fn*:foo + - fn*:foo() + - fn:select: GLSL select() with scalar branches requires a scalar boolean condition] + `); + }); + + it('translates `saturate(v)` to `clamp(v, 0.0, 1.0)`', () => { + function foo() { + 'use gpu'; + const scalar = 2; + const vec3 = d.vec3f(1, 2, 3); + std.saturate(scalar); + std.saturate(vec3); + } + + expect(tgpu.resolve([foo], glOptions())).toMatchInlineSnapshot(` + "void foo() { + int scalar = 2; + vec3 vec3_1 = vec3(1, 2, 3); + clamp(float(scalar), 0f, 1f); + clamp(vec3_1, 0f, 1f); + }" + `); + }); + + it('translates bitcast', () => { + function foo() { + 'use gpu'; + const f = d.f32(1.5); + const f2 = d.vec2f(1.5); + const u = d.u32(15); + const u2 = d.vec2u(15); + const i = d.i32(-5); + const i2 = d.vec2i(-5); + + std.bitcast(d.f32, d.f32)(f); //no-op + std.bitcast(d.u32, d.u32)(u); //no-op + std.bitcast(d.i32, d.i32)(i); //no-op + + std.bitcast(d.f32, d.u32)(f); + std.bitcast(d.f32, d.i32)(f); + std.bitcast(d.u32, d.f32)(u); + std.bitcast(d.i32, d.f32)(i); + + std.bitcast(d.vec2f, d.vec2u)(f2); + std.bitcast(d.vec2f, d.vec2i)(f2); + std.bitcast(d.vec2u, d.vec2f)(u2); + std.bitcast(d.vec2i, d.vec2f)(i2); + } + + expect(tgpu.resolve([foo], glOptions())).toMatchInlineSnapshot(` + "void foo() { + float f = 1.5f; + vec2 f2 = vec2(1.5); + uint u = 15u; + uvec2 u2 = uvec2(15); + int i = -5i; + ivec2 i2 = ivec2(-5); + f; + u; + i; + floatBitsToUint(f); + floatBitsToInt(f); + uintBitsToFloat(u); + intBitsToFloat(i); + floatBitsToUint(f2); + floatBitsToInt(f2); + uintBitsToFloat(u2); + intBitsToFloat(i2); + }" + `); + }); +}); + describe('GlslGenerator - function definitions', () => { it('generates proper function signatures', () => { function add(a: number, b: number) { diff --git a/packages/typegpu/src/core/function/dualImpl.ts b/packages/typegpu/src/core/function/dualImpl.ts index 564feb5d6b..ccb0f6cebf 100644 --- a/packages/typegpu/src/core/function/dualImpl.ts +++ b/packages/typegpu/src/core/function/dualImpl.ts @@ -12,7 +12,11 @@ type AnyFn = (...args: never[]) => unknown; interface DualImplOptions { readonly name: string | undefined; readonly normalImpl: T | string; - readonly codegenImpl: (ctx: ResolutionCtx, args: MapValueToSnippet>) => string; + readonly codegenImpl: ( + ctx: ResolutionCtx, + args: MapValueToSnippet>, + returnType: BaseData, + ) => string; readonly signature: | { argTypes: (BaseData | BaseData[])[]; @@ -117,9 +121,10 @@ export function dualImpl(options: DualImplOptions): DualFn a.possibleSideEffects); + const concreteReturnType = concretize(returnType); return snip( - options.codegenImpl(ctx, converted), - concretize(returnType), + options.codegenImpl(ctx, converted, concreteReturnType), + concreteReturnType, // Functions give up ownership of their return value /* origin */ 'runtime', possibleSideEffects, diff --git a/packages/typegpu/src/std/array.ts b/packages/typegpu/src/std/array.ts index 6fe62f6f55..8fbb639c68 100644 --- a/packages/typegpu/src/std/array.ts +++ b/packages/typegpu/src/std/array.ts @@ -1,5 +1,4 @@ import { dualImpl } from '../core/function/dualImpl.ts'; -import { stitch } from '../core/resolve/stitch.ts'; import { abstractInt, u32 } from '../data/numeric.ts'; import { ptrFn } from '../data/ptr.ts'; import { type _ref as ref, isRef } from '../data/ref.ts'; @@ -18,9 +17,9 @@ export const arrayLength = dualImpl({ }; }, normalImpl: (a: unknown[] | ref) => (isRef(a) ? a.$.length : a.length), - codegenImpl(_ctx, [a]) { + codegenImpl(ctx, [a]) { const length = sizeOfPointedToArray(a.dataType); - return length > 0 ? `${length}` : stitch`arrayLength(${a})`; + return length > 0 ? `${length}` : ctx.gen.call('arrayLength', [], [a]); }, sideEffects: false, }); diff --git a/packages/typegpu/src/std/bitcast.ts b/packages/typegpu/src/std/bitcast.ts index 2b7fac5498..f0856ea4fe 100644 --- a/packages/typegpu/src/std/bitcast.ts +++ b/packages/typegpu/src/std/bitcast.ts @@ -1,5 +1,4 @@ import { dualImpl } from '../core/function/dualImpl.ts'; -import { stitch } from '../core/resolve/stitch.ts'; import { bitcastF32toU32Impl, bitcastU32toF32Impl, @@ -44,6 +43,7 @@ import { SignatureNotSupportedError } from '../errors.ts'; import { getName } from '../shared/meta.ts'; import type { Infer } from '../shared/repr.ts'; import { comptime } from '../core/function/comptime.ts'; +import { coerceToSnippet } from '../tgsl/generationHelpers.ts'; type BitcastU32toF32Overload = ( value: T, @@ -64,10 +64,8 @@ export const bitcastU32toF32 = dualImpl({ } return VectorOps.bitcastU32toF32[value.kind](value); }) as BitcastU32toF32Overload, - codegenImpl: (_ctx, [n]) => { - return isVec(n.dataType) - ? stitch`bitcast(${n})` - : stitch`bitcast(${n})`; + codegenImpl: (ctx, [n], returnType) => { + return ctx.gen.call('bitcast', [coerceToSnippet(returnType)], [n]); }, signature: (...arg) => { const uargs = unifyStrict(arg, u32AllowedSchemas); @@ -103,10 +101,8 @@ export const bitcastU32toI32 = dualImpl({ } return VectorOps.bitcastU32toI32[value.kind](value); }) as BitcastU32toI32Overload, - codegenImpl: (_ctx, [n]) => { - return isVec(n.dataType) - ? stitch`bitcast(${n})` - : stitch`bitcast(${n})`; + codegenImpl: (ctx, [n], returnType) => { + return ctx.gen.call('bitcast', [coerceToSnippet(returnType)], [n]); }, signature: (...arg) => { const uargs = unifyStrict(arg, u32AllowedSchemas); @@ -144,10 +140,8 @@ export const bitcastF32toU32 = dualImpl({ } return VectorOps.bitcastF32toU32[value.kind](value); }) as BitcastF32toU32Overload, - codegenImpl: (_ctx, [n]) => { - return isVec(n.dataType) - ? stitch`bitcast(${n})` - : stitch`bitcast(${n})`; + codegenImpl: (ctx, [n], returnType) => { + return ctx.gen.call('bitcast', [coerceToSnippet(returnType)], [n]); }, signature: (...arg) => { const uargs = unifyStrict(arg, f32AllowedSchemas); @@ -283,7 +277,7 @@ function bitcastFor(inType, outType), - codegenImpl: (_ctx, [n]) => stitch`bitcast<${outType.type}>(${n})`, + codegenImpl: (ctx, [n]) => ctx.gen.call('bitcast', [coerceToSnippet(outType)], [n]), signature: (arg) => { const uarg = unifyStrict([arg], [inType]); if (!uarg) { diff --git a/packages/typegpu/src/std/boolean.ts b/packages/typegpu/src/std/boolean.ts index 8d6d673f93..04d945a216 100644 --- a/packages/typegpu/src/std/boolean.ts +++ b/packages/typegpu/src/std/boolean.ts @@ -443,7 +443,7 @@ export const select = dualImpl({ }, normalImpl: cpuSelect, codegenImpl: (ctx, [f, t, cond]) => { - const result = stitch`select(${f}, ${t}, ${cond})`; + const result = ctx.gen.call('select', [], [f, t, cond]); if ( !validSelectBranchTypes.includes(f.dataType as AnyWgslData) || !validSelectBranchTypes.includes(t.dataType as AnyWgslData) diff --git a/packages/typegpu/src/std/numeric.ts b/packages/typegpu/src/std/numeric.ts index 50928e1418..2619040904 100644 --- a/packages/typegpu/src/std/numeric.ts +++ b/packages/typegpu/src/std/numeric.ts @@ -1072,7 +1072,7 @@ export const saturate = dualImpl({ name: 'saturate', signature: unifyRestrictedSignature(anyFloat), normalImpl: cpuSaturate, - codegenImpl: (_ctx, [value]) => stitch`saturate(${value})`, + codegenImpl: (ctx, [value]) => ctx.gen.call('saturate', [], [value]), sideEffects: false, }); diff --git a/packages/typegpu/src/tgsl/shaderGenerator.ts b/packages/typegpu/src/tgsl/shaderGenerator.ts index 96985e7b11..03f31f148b 100644 --- a/packages/typegpu/src/tgsl/shaderGenerator.ts +++ b/packages/typegpu/src/tgsl/shaderGenerator.ts @@ -64,6 +64,7 @@ export interface ShaderGenerator { functionDefinition(options: FunctionDefinitionOptions): string; typeInstantiation(schema: BaseData, args: readonly Snippet[]): ResolvedSnippet; - typeAnnotation(schema: BaseData): string; numericLiteral(value: number, schema: BaseData): ResolvedSnippet; + typeAnnotation(schema: BaseData): string; + call(name: string, templateParams: readonly Snippet[], args: readonly Snippet[]): string; } diff --git a/packages/typegpu/src/tgsl/wgslGenerator.ts b/packages/typegpu/src/tgsl/wgslGenerator.ts index b0136457b6..139830b179 100644 --- a/packages/typegpu/src/tgsl/wgslGenerator.ts +++ b/packages/typegpu/src/tgsl/wgslGenerator.ts @@ -1089,6 +1089,18 @@ ${this.ctx.pre}}`; return snip(base, schema, /* origin */ 'constant', false); } + public call(name: string, templateParams: readonly Snippet[], args: readonly Snippet[]): string { + const resolvedTemplateParams = templateParams + .map((arg) => this.ctx.resolveSnippet(arg).value) + .join(', '); + const resolvedArgs = args.map((arg) => this.ctx.resolveSnippet(arg).value).join(', '); + + if (resolvedTemplateParams.length > 0) { + return `${name}<${resolvedTemplateParams}>(${resolvedArgs})`; + } + return `${name}(${resolvedArgs})`; + } + protected _return(statement: tinyest.Return): string { const returnNode = statement[1];