Skip to main content

Module shader

Module shader 

Source
Expand description

The crate-authored SPIR-V compute kernels (ADR 0003, ADR 0007).

Every module the backend hands to a driver is assembled here from fixed templates and specialized at load_program through specialization constants alone. Guest bytes never reach the driver’s shader compiler: a TOSA artifact selects a KernelKey and supplies validated shape parameters as a specialization payload, nothing else (docs/threat-model.md, transient-compile budget). The only template parameters that change a module’s instructions are device properties fixed when the backend opens a device — the workgroup size, the MATMUL tile, the length of the storage-buffer descriptor array — plus the operator selection itself.

§Operand addressing

Each kernel reads and writes tensors through one descriptor: set 0, binding 0, an array of { uint words[]; } storage-buffer blocks. An Operand names an array element and a base offset in 32-bit words, both specialization constants, so one module per kernel serves every binding layout, the program-owned arena, and CONCAT with any input count. Word-storage tensors (FP32, INT32) are addressed one element per word; half-storage tensors (FP16) are packed two elements per word and byte-storage tensors (BOOL) four per word, and both are read with word loads and written with OpAtomicAnd/OpAtomicOr, so a kernel never modifies bytes outside the elements it owns even at a tensor’s unaligned tail.

§Binary16 kernels (ADR 0008)

FP16 kernels are separate KernelKey variants (float: Storage::Half) built from the same 32-bit-only instruction set as the FP32 kernels — no 16-bit types, no device features, no driver float-controls to trust. Packed binary16 tensors are unpacked with integer ops and widened to binary32 exactly (Builder::widen_f16); every float lane evaluates in binary32 — which TOSA 1.0 §1.10.3 explicitly permits (“fp16_t operations [may] be implemented using the fp32_t datatype”) — and results narrow back through crate-owned round-to-nearest-even integer code (Builder::narrow_f16) that produces subnormals on every device. For ADD/SUB/MUL that is the correctly rounded binary16 result (their exact results fit the binary32 significand); the comparison and selection lanes are exact; RECIPROCAL stays within TOSA’s tolerance; the transcendental lanes keep ADR 0007’s crate-owned binary32 numerics; MATMUL and reduction folds accumulate in binary32, the accumulator width TOSA assigns FP16. NEGATE/ABS are integer sign masks on the packed lane, and data movement copies the 16-bit lanes as integers, so NaN payloads and subnormals move bit-exactly. Because the conversions are integer code, no compiler can demote the binary32 arithmetic back to f16, and binary16 results are bit-identical on every device.

§Numerics policy

Every floating-point arithmetic result carries NoContraction, so no driver may fuse a multiply and an add: the same TOSA graph yields the same bits on every conformant device. SIN, COS, TANH, and ERF are evaluated by crate-authored range reductions and polynomials rather than the driver’s built-ins, whose precision Vulkan specifies loosely (sin/cos: absolute error 2⁻¹¹) or not at all (tanh). EXP, LOG, POW, and RSQRT use the GLSL.std.450 built-ins, whose relative-error bounds Vulkan does specify. NaN-mode attributes (PROPAGATE/IGNORE) follow the TOSA 1.0 pseudocode literally, with explicit OpIsNan selects instead of the driver’s undefined NaN handling for FMax/FMin.

Structs§

ElementwiseSpec
Specialization payload of an elementwise dispatch. The declaration order in the elementwise kernel is: count, each input operand, the output operand, then (broadcast only) the output dims and each input’s element strides, then the clamp bounds.
MatmulGeometry
Shape of one MATMUL workgroup: tile_x × tile_y invocations, each accumulating a micro_m × micro_n register block, over shared-memory slabs depth deep in k. Every dimension is a power of two and the two slabs stage evenly over the invocations.
MoveGeometry
Iteration space of a strided copy: dims (leading ones to MAX_RANK), element strides and element offsets on each side. Strides are wrapping u32 so a reversed axis is -inner.
Nvfp4MatmulSpec
Specialization payload for native row-major NVFP4: activation, packed weights, block scales, tensor scale, output, then m, n, k, epilogue.
Operand
Where a kernel operand lives: a descriptor-array element and a base offset in 32-bit words.
PoolGeometry
NHWC pooling geometry: batch, input height/width, channels, output height/width, kernel, stride, and the top/left pads (the bottom/right pads only shape the output).

Enums§

ElementwiseOp
Elementwise operator lanes: the scalar function one invocation applies per output element.
Fp8Format
Which of the two TOSA FP8 encodings a byte-storage float tensor carries. They differ in exponent width, bias, and specials: E4M3 has no infinity (0x7f/0xff are its only NaNs, finite max 448) while E5M2 is IEEE-shaped (infinity at 0x7c, finite max 57344).
KernelKey
One assembled kernel variant. Everything that changes instructions is in the key; everything that changes only numbers is a specialization constant.
NanMode
TOSA NaN-propagation attribute value a kernel is specialized for.
ReduceOp
Reduction operator lanes.
Storage
How a tensor’s scalars are laid out in the storage words a kernel addresses.

Constants§

MATMUL_MICRO
Outputs per invocation per side of the register-tiled MATMUL: each invocation of a tile × tile workgroup accumulates a MATMUL_MICRO × MATMUL_MICRO block, so the workgroup covers a matmul_block-sided square of the result.
MAX_ELEMENTWISE_INPUTS
Elementwise kernels take at most three tensor inputs (SELECT).
MAX_RANK
TOSA level 8K rank bound the strided kernels are sized for.
SPIRV_MAGIC
The SPIR-V magic number.
SPIRV_VERSION_1_3
SPIR-V 1.3: the version every Vulkan 1.1+ implementation must consume, and the first with the StorageBuffer storage class in core.
STREAM_COLUMNS
Output columns one streaming MATMUL workgroup covers.
STREAM_ROWS
Rows the streaming MATMUL kernel carries per invocation, and the row count at or below which lowering selects it over the register-tiled kernel.
STREAM_WORKGROUP
Invocations of a streaming MATMUL workgroup: STREAM_COLUMNS / lanes weight words times STREAM_WORKGROUP / that slices of k.

Functions§

f16_to_f32
f32_to_f16_bits
The binary16 bit pattern nearest to value, round-to-nearest-even. NaN is canonicalized to the quiet 0x7e00 payload with its sign; magnitudes at or above 65520 round to infinity, and subnormals are produced, never flushed).
f32_to_fp8_bits
The exact binary32 value of a binary16 bit pattern (host side of the kernels’ unpack: every binary16 value, subnormals included, is exactly representable in binary32; NaN payloads are preserved). Narrow binary32 to one FP8 encoding with round-to-nearest, ties-to-even.
linear_workgroups
Number of workgroups covering count items at workgroup invocations each, capped at limit; the kernels loop with a grid stride so the cap only trades parallelism, never coverage.
matmul_block
Side of the output square one MATMUL workgroup of tile × tile invocations computes.
matmul_shared_bytes
Bytes of workgroup-shared memory the MATMUL kernels at tile declare, whichever kernel needs more.
matmul_spec
Specialization payload of a MATMUL: lhs, rhs, output, then m, n, k, batch.
matmul_workgroups
Workgroup counts of a tiled MATMUL over m rows, n columns, and batch batches.
max_pool_spec
Specialization payload of a MAX_POOL2D dispatch.
move_spec
Specialization payload of a copy dispatch.
reduce_spec
Specialization payload of a reduction: input, output, outer, axis, inner.
stream_matmul_shared_bytes
Bytes of workgroup-shared memory the streaming kernel declares for its final reduction: every invocation’s STREAM_ROWS × lanes partial sums, at the widest lane count.
stream_matmul_workgroups
Workgroup counts of a streaming MATMUL over n columns and batch batches.

Type Aliases§

Id
A SPIR-V result id.