pub enum KernelKey {
Nvfp4Matmul {
buffers: u32,
cooperative: bool,
subgroup: bool,
},
Elementwise {
op: ElementwiseOp,
float: Storage,
broadcast: bool,
workgroup: u32,
buffers: u32,
},
Reduce {
op: ReduceOp,
float: Storage,
workgroup: u32,
buffers: u32,
},
Matmul {
input: Storage,
output: Storage,
tile: u32,
buffers: u32,
},
MatmulStream {
rhs: Storage,
output: Storage,
buffers: u32,
},
MaxPool {
nan_mode: NanMode,
float: Storage,
workgroup: u32,
buffers: u32,
},
Cast {
input: Storage,
output: Storage,
workgroup: u32,
buffers: u32,
},
Move {
storage: Storage,
contiguous: bool,
workgroup: u32,
buffers: u32,
},
}Expand description
One assembled kernel variant. Everything that changes instructions is in the key; everything that changes only numbers is a specialization constant.
Variants§
Nvfp4Matmul
F32 activations times row-major packed E2M1 weights and E4M3 block scales.
Elementwise
Elementwise lanes over count output elements; broadcast selects the strided
multi-index addressing, otherwise every operand shares the output’s linear index. float
is the storage of the operator’s floating-point tensors (Word for FP32, Half for
FP16); BOOL lanes are byte storage in either variant. Binary16 lanes evaluate in
binary32 and narrow once, except the integer NEGATE/ABS sign lanes (ADR 0008).
Reduce
Axis reduction: one invocation per output element, sequential ascending-axis fold. The
input is read at float storage; sums and products fold in binary32 (the TOSA
accumulator width), and the output is stored back at float storage.
Matmul
Batched, register-tiled matrix multiplication: a tile × tile workgroup computes a
matmul_block-sided output square from workgroup-shared binary32 slabs. Both
operands are read at input storage and accumulate in binary32 — the accumulator width
TOSA assigns FP16 and FP8 MATMUL — and the result is stored at output storage. The two
differ only for the FP8 tier, where TOSA defines MATMUL as (FP8, FP8) -> FP16.
MatmulStream
Split-k streaming MATMUL for m ≤ STREAM_ROWS rows:
a 1-D workgroup of STREAM_WORKGROUP invocations over STREAM_COLUMNS columns of
rhs (read at rhs storage, a word per invocation), the lhs read as binary32 words —
lowering widens a narrower lhs beforehand — and one fixed-order reduction of the k
slices at the end. Same accumulator width as Self::Matmul.
MaxPool
NHWC max pooling with padding excluded from the window, at float storage.
Cast
Elementwise float conversion: read at input storage, write at output. One dispatch
per CAST between float dtypes, including both FP8 directions (ADR 0009).
Move
Strided copy over a rank-MAX_RANK iteration space (TRANSPOSE, REVERSE, CONCAT
segments); contiguous collapses to a linear copy.
Implementations§
Source§impl KernelKey
impl KernelKey
Sourcepub fn every_variant() -> Vec<KernelKey>
pub fn every_variant() -> Vec<KernelKey>
Every kernel variant this backend can assemble, at one representative tuning. The SPIR-V validation sweep and the specialization-count test both walk this list, so a new variant is validated the moment it is added here.
Sourcepub const fn spec_constant_count(self) -> u32
pub const fn spec_constant_count(self) -> u32
Number of specialization constants the module declares, ids 0..count.
Sourcepub const fn local_size(self) -> [u32; 3]
pub const fn local_size(self) -> [u32; 3]
OpExecutionMode LocalSize of the module.