Skip to main content

virtio_accel_vulkan/
shader.rs

1//! The crate-authored SPIR-V compute kernels (ADR 0003, ADR 0007).
2//!
3//! Every module the backend hands to a driver is assembled here from fixed templates and
4//! specialized at `load_program` through specialization constants alone. Guest bytes never reach
5//! the driver's shader compiler: a TOSA artifact selects a [`KernelKey`] and supplies validated
6//! shape parameters as a specialization payload, nothing else (`docs/threat-model.md`,
7//! transient-compile budget). The only template parameters that change a module's instructions
8//! are device properties fixed when the backend opens a device — the workgroup size, the MATMUL
9//! tile, the length of the storage-buffer descriptor array — plus the operator selection itself.
10//!
11//! ## Operand addressing
12//!
13//! Each kernel reads and writes tensors through one descriptor: set 0, binding 0, an array of
14//! `{ uint words[]; }` storage-buffer blocks. An [`Operand`] names an array element and a base
15//! offset in 32-bit words, both specialization constants, so one module per kernel serves every
16//! binding layout, the program-owned arena, and `CONCAT` with any input count. Word-storage
17//! tensors (`FP32`, `INT32`) are addressed one element per word; half-storage tensors (`FP16`)
18//! are packed two elements per word and byte-storage tensors (`BOOL`) four per word, and both
19//! are read with word loads and written with `OpAtomicAnd`/`OpAtomicOr`, so a kernel never
20//! modifies bytes outside the elements it owns even at a tensor's unaligned tail.
21//!
22//! ## Binary16 kernels (ADR 0008)
23//!
24//! FP16 kernels are separate [`KernelKey`] variants (`float: Storage::Half`) built from the same
25//! 32-bit-only instruction set as the FP32 kernels — no 16-bit types, no device features, no
26//! driver float-controls to trust. Packed binary16 tensors are unpacked with integer ops and
27//! widened to binary32 exactly (`Builder::widen_f16`); every float lane evaluates in binary32
28//! — which TOSA 1.0 §1.10.3 explicitly permits ("fp16_t operations \[may\] be implemented using
29//! the fp32_t datatype") — and results narrow back through crate-owned round-to-nearest-even
30//! integer code (`Builder::narrow_f16`) that produces subnormals on every device. For
31//! `ADD`/`SUB`/`MUL` that is the correctly rounded binary16 result (their exact results fit the
32//! binary32 significand); the comparison and selection lanes are exact; `RECIPROCAL` stays
33//! within TOSA's tolerance; the transcendental lanes keep ADR 0007's crate-owned binary32
34//! numerics; MATMUL and reduction folds accumulate in binary32, the accumulator width TOSA
35//! assigns FP16. `NEGATE`/`ABS` are integer sign masks on the packed lane, and data movement
36//! copies the 16-bit lanes as integers, so NaN payloads and subnormals move bit-exactly. Because
37//! the conversions are integer code, no compiler can demote the binary32 arithmetic back to
38//! f16, and binary16 results are bit-identical on every device.
39//!
40//! ## Numerics policy
41//!
42//! Every floating-point arithmetic result carries `NoContraction`, so no driver may fuse a
43//! multiply and an add: the same TOSA graph yields the same bits on every conformant device.
44//! `SIN`, `COS`, `TANH`, and `ERF` are evaluated by crate-authored range reductions and
45//! polynomials rather than the driver's built-ins, whose precision Vulkan specifies loosely
46//! (`sin`/`cos`: absolute error 2⁻¹¹) or not at all (`tanh`). `EXP`, `LOG`, `POW`, and
47//! `RSQRT` use the `GLSL.std.450` built-ins, whose relative-error bounds Vulkan does specify.
48//! NaN-mode attributes (`PROPAGATE`/`IGNORE`) follow the TOSA 1.0 pseudocode literally, with
49//! explicit `OpIsNan` selects instead of the driver's undefined NaN handling for `FMax`/`FMin`.
50
51use std::collections::HashMap;
52
53/// SPIR-V 1.3: the version every Vulkan 1.1+ implementation must consume, and the first with the
54/// `StorageBuffer` storage class in core.
55pub const SPIRV_VERSION_1_3: u32 = 0x0001_0300;
56/// The SPIR-V magic number.
57pub const SPIRV_MAGIC: u32 = 0x0723_0203;
58
59/// The leading fraction bits of 2/π, most significant first, preceded by one zero word.
60///
61/// The zero word biases the Payne–Hanek bit window so its offset is never negative for any
62/// argument the reduction handles; eleven words of 2/π (352 bits) cover every binary32 exponent
63/// with the 128-bit window the reduction extracts (`docs/adr/0007-fp32-operator-tier.md`).
64const TWO_OVER_PI_BITS: &[u32] = &[
65    0x0000_0000,
66    0xa2f9_836e,
67    0x4e44_1529,
68    0xfc27_57d1,
69    0xf534_ddc0,
70    0xdb62_9599,
71    0x3c43_9041,
72    0xfe51_63ab,
73    0xdebb_c561,
74    0xb724_6e3a,
75    0x424d_d2e0,
76    0x0649_2eea,
77];
78
79/// Below this magnitude the three-part π/4 subtraction is exact enough; at or above it the
80/// Payne–Hanek reduction takes over.
81const SINCOS_FAST_RANGE: f32 = 8192.0;
82
83/// TOSA level 8K rank bound the strided kernels are sized for.
84pub const MAX_RANK: usize = 6;
85/// Elementwise kernels take at most three tensor inputs (`SELECT`).
86pub const MAX_ELEMENTWISE_INPUTS: usize = 3;
87
88// SPIR-V opcodes (Unified specification, section 3.52).
89const OP_EXT_INST_IMPORT: u16 = 11;
90const OP_EXT_INST: u16 = 12;
91const OP_EXTENSION: u16 = 10;
92const OP_MEMORY_MODEL: u16 = 14;
93const OP_ENTRY_POINT: u16 = 15;
94const OP_EXECUTION_MODE: u16 = 16;
95const OP_CAPABILITY: u16 = 17;
96const OP_TYPE_VOID: u16 = 19;
97const OP_TYPE_BOOL: u16 = 20;
98const OP_TYPE_INT: u16 = 21;
99const OP_TYPE_FLOAT: u16 = 22;
100const OP_TYPE_VECTOR: u16 = 23;
101const OP_TYPE_ARRAY: u16 = 28;
102const OP_TYPE_RUNTIME_ARRAY: u16 = 29;
103const OP_TYPE_STRUCT: u16 = 30;
104const OP_TYPE_POINTER: u16 = 32;
105const OP_TYPE_FUNCTION: u16 = 33;
106const OP_CONSTANT_FALSE: u16 = 42;
107const OP_CONSTANT: u16 = 43;
108const OP_CONSTANT_COMPOSITE: u16 = 44;
109const OP_SPEC_CONSTANT: u16 = 50;
110const OP_FUNCTION: u16 = 54;
111const OP_FUNCTION_END: u16 = 56;
112const OP_VARIABLE: u16 = 59;
113const OP_LOAD: u16 = 61;
114const OP_STORE: u16 = 62;
115const OP_ACCESS_CHAIN: u16 = 65;
116const OP_DECORATE: u16 = 71;
117const OP_MEMBER_DECORATE: u16 = 72;
118const OP_CONVERT_F_TO_U: u16 = 109;
119const OP_CONVERT_U_TO_F: u16 = 112;
120const OP_F_CONVERT: u16 = 115;
121const OP_BITCAST: u16 = 124;
122const OP_F_NEGATE: u16 = 127;
123const OP_I_ADD: u16 = 128;
124const OP_F_ADD: u16 = 129;
125const OP_I_SUB: u16 = 130;
126const OP_F_SUB: u16 = 131;
127const OP_I_MUL: u16 = 132;
128const OP_F_MUL: u16 = 133;
129const OP_U_DIV: u16 = 134;
130const OP_F_DIV: u16 = 136;
131const OP_U_MOD: u16 = 137;
132const OP_IS_NAN: u16 = 156;
133const OP_LOGICAL_NOT_EQUAL: u16 = 165;
134const OP_LOGICAL_OR: u16 = 166;
135const OP_LOGICAL_AND: u16 = 167;
136const OP_LOGICAL_NOT: u16 = 168;
137const OP_SELECT: u16 = 169;
138const OP_I_EQUAL: u16 = 170;
139const OP_I_NOT_EQUAL: u16 = 171;
140const OP_U_GREATER_THAN_EQUAL: u16 = 174;
141const OP_U_LESS_THAN: u16 = 176;
142const OP_F_ORD_EQUAL: u16 = 180;
143const OP_F_ORD_LESS_THAN: u16 = 184;
144const OP_F_ORD_GREATER_THAN: u16 = 186;
145const OP_F_ORD_GREATER_THAN_EQUAL: u16 = 190;
146const OP_SHIFT_RIGHT_LOGICAL: u16 = 194;
147const OP_SHIFT_LEFT_LOGICAL: u16 = 196;
148const OP_BITWISE_OR: u16 = 197;
149const OP_BITWISE_XOR: u16 = 198;
150const OP_BITWISE_AND: u16 = 199;
151const OP_NOT: u16 = 200;
152const OP_CONTROL_BARRIER: u16 = 224;
153const OP_ATOMIC_AND: u16 = 240;
154const OP_ATOMIC_OR: u16 = 241;
155const OP_LOOP_MERGE: u16 = 246;
156const OP_SELECTION_MERGE: u16 = 247;
157const OP_LABEL: u16 = 248;
158const OP_BRANCH: u16 = 249;
159const OP_BRANCH_CONDITIONAL: u16 = 250;
160const OP_RETURN: u16 = 253;
161const OP_GROUP_NON_UNIFORM_F_ADD: u16 = 350;
162const OP_TYPE_COOPERATIVE_MATRIX_KHR: u16 = 4456;
163const OP_COOPERATIVE_MATRIX_LOAD_KHR: u16 = 4457;
164const OP_COOPERATIVE_MATRIX_STORE_KHR: u16 = 4458;
165const OP_COOPERATIVE_MATRIX_MUL_ADD_KHR: u16 = 4459;
166
167// Enumerants (section 3).
168const CAPABILITY_SHADER: u32 = 1;
169const CAPABILITY_FLOAT16: u32 = 9;
170const CAPABILITY_GROUP_NON_UNIFORM_ARITHMETIC: u32 = 63;
171const CAPABILITY_VULKAN_MEMORY_MODEL: u32 = 5345;
172const CAPABILITY_COOPERATIVE_MATRIX_KHR: u32 = 6022;
173const ADDRESSING_MODEL_LOGICAL: u32 = 0;
174const MEMORY_MODEL_GLSL450: u32 = 1;
175const EXECUTION_MODEL_GL_COMPUTE: u32 = 5;
176const EXECUTION_MODE_LOCAL_SIZE: u32 = 17;
177const STORAGE_CLASS_INPUT: u32 = 1;
178const STORAGE_CLASS_WORKGROUP: u32 = 4;
179const STORAGE_CLASS_PRIVATE: u32 = 6;
180const STORAGE_CLASS_FUNCTION: u32 = 7;
181const STORAGE_CLASS_STORAGE_BUFFER: u32 = 12;
182const DECORATION_SPEC_ID: u32 = 1;
183const DECORATION_BLOCK: u32 = 2;
184const DECORATION_ARRAY_STRIDE: u32 = 6;
185const DECORATION_BUILT_IN: u32 = 11;
186const DECORATION_BINDING: u32 = 33;
187const DECORATION_DESCRIPTOR_SET: u32 = 34;
188const DECORATION_OFFSET: u32 = 35;
189const DECORATION_NO_CONTRACTION: u32 = 42;
190const BUILT_IN_NUM_WORKGROUPS: u32 = 24;
191const BUILT_IN_WORKGROUP_ID: u32 = 26;
192const BUILT_IN_LOCAL_INVOCATION_ID: u32 = 27;
193const BUILT_IN_GLOBAL_INVOCATION_ID: u32 = 28;
194const FUNCTION_CONTROL_NONE: u32 = 0;
195const SELECTION_CONTROL_NONE: u32 = 0;
196const LOOP_CONTROL_NONE: u32 = 0;
197const SCOPE_DEVICE: u32 = 1;
198const SCOPE_WORKGROUP: u32 = 2;
199const MEMORY_SEMANTICS_RELAXED: u32 = 0;
200const MEMORY_SEMANTICS_ACQUIRE_RELEASE_WORKGROUP: u32 = 0x8 | 0x100;
201
202// GLSL.std.450 extended instructions.
203const GLSL_ROUND_EVEN: u32 = 2;
204const GLSL_FABS: u32 = 4;
205const GLSL_FLOOR: u32 = 8;
206const GLSL_CEIL: u32 = 9;
207const GLSL_POW: u32 = 26;
208const GLSL_EXP: u32 = 27;
209const GLSL_LOG: u32 = 28;
210const GLSL_INVERSE_SQRT: u32 = 32;
211const GLSL_UMIN: u32 = 38;
212const GLSL_FMA: u32 = 50;
213const GLSL_FIND_U_MSB: u32 = 75;
214
215/// A SPIR-V result id.
216pub type Id = u32;
217
218/// Which of the two TOSA FP8 encodings a byte-storage float tensor carries. They differ in
219/// exponent width, bias, and specials: E4M3 has no infinity (`0x7f`/`0xff` are its only NaNs,
220/// finite max 448) while E5M2 is IEEE-shaped (infinity at `0x7c`, finite max 57344).
221#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
222pub enum Fp8Format {
223    E4M3,
224    E5M2,
225}
226
227/// How a tensor's scalars are laid out in the storage words a kernel addresses.
228#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
229pub enum Storage {
230    /// One 32-bit word per element (`FP32`, `INT32`).
231    Word,
232    /// One byte per element, four to a word (`BOOL`).
233    Byte,
234    /// One byte per element, four to a word, holding an FP8 encoding. Geometrically identical
235    /// to [`Self::Byte`]; it is a separate variant because the byte is a raw float pattern, not
236    /// a canonical `0`/`1`, so it is loaded through an exact widening and moved as raw bits
237    /// (ADR 0009).
238    Quarter(Fp8Format),
239    /// Two bytes per element, two to a word (`FP16`); the 16-bit lanes are unpacked on load and
240    /// repacked with `OpAtomicAnd`/`OpAtomicOr` on store, so a kernel never modifies the
241    /// neighbouring element of its word.
242    Half,
243}
244
245impl Storage {
246    /// Elements packed into one 32-bit storage word.
247    pub const fn lanes(self) -> u32 {
248        match self {
249            Self::Word => 1,
250            Self::Half => 2,
251            Self::Byte | Self::Quarter(_) => 4,
252        }
253    }
254}
255
256/// TOSA NaN-propagation attribute value a kernel is specialized for.
257#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
258pub enum NanMode {
259    /// A NaN operand yields NaN (`apply_max_s`/`apply_min_s` propagate).
260    Propagate,
261    /// A NaN operand is ignored in favor of the other operand.
262    Ignore,
263}
264
265/// Elementwise operator lanes: the scalar function one invocation applies per output element.
266#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
267pub enum ElementwiseOp {
268    Abs,
269    Ceil,
270    Cos,
271    Erf,
272    Exp,
273    Floor,
274    Log,
275    Negate,
276    Reciprocal,
277    Rsqrt,
278    Sin,
279    Sigmoid,
280    Tanh,
281    /// `apply_min(apply_max(x, lo), hi)`; the bounds arrive as two trailing specialization
282    /// constants holding FP32 bit patterns.
283    Clamp(NanMode),
284    Add,
285    Sub,
286    Mul,
287    Pow,
288    Maximum(NanMode),
289    Minimum(NanMode),
290    Equal,
291    Greater,
292    GreaterEqual,
293    LogicalAnd,
294    LogicalOr,
295    LogicalXor,
296    LogicalNot,
297    Select,
298    /// Byte-exact copy of a `BOOL` tensor (`IDENTITY`, `RESHAPE` on byte storage).
299    CopyBytes,
300}
301
302impl ElementwiseOp {
303    /// Storage of each tensor input, in operand order. `Word` marks a floating-point lane: the
304    /// kernel resolves it to its own float storage (`Word` for FP32, `Half` for FP16); `Byte`
305    /// lanes are always `BOOL`.
306    pub const fn inputs(self) -> &'static [Storage] {
307        match self {
308            Self::Abs
309            | Self::Ceil
310            | Self::Cos
311            | Self::Erf
312            | Self::Exp
313            | Self::Floor
314            | Self::Log
315            | Self::Negate
316            | Self::Reciprocal
317            | Self::Rsqrt
318            | Self::Sin
319            | Self::Sigmoid
320            | Self::Tanh
321            | Self::Clamp(_) => &[Storage::Word],
322            Self::Add
323            | Self::Sub
324            | Self::Mul
325            | Self::Pow
326            | Self::Maximum(_)
327            | Self::Minimum(_)
328            | Self::Equal
329            | Self::Greater
330            | Self::GreaterEqual => &[Storage::Word, Storage::Word],
331            Self::LogicalAnd | Self::LogicalOr | Self::LogicalXor => {
332                &[Storage::Byte, Storage::Byte]
333            }
334            Self::LogicalNot | Self::CopyBytes => &[Storage::Byte],
335            Self::Select => &[Storage::Byte, Storage::Word, Storage::Word],
336        }
337    }
338
339    /// Storage of the output tensor, with the same float-lane convention as [`Self::inputs`].
340    pub const fn output(self) -> Storage {
341        match self {
342            Self::Equal
343            | Self::Greater
344            | Self::GreaterEqual
345            | Self::LogicalAnd
346            | Self::LogicalOr
347            | Self::LogicalXor
348            | Self::LogicalNot
349            | Self::CopyBytes => Storage::Byte,
350            _ => Storage::Word,
351        }
352    }
353
354    /// Trailing operator-specific specialization constants (after the operand and shape block).
355    pub const fn extra_spec_constants(self) -> u32 {
356        match self {
357            Self::Clamp(_) => 2,
358            _ => 0,
359        }
360    }
361}
362
363/// Reduction operator lanes.
364#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
365pub enum ReduceOp {
366    Sum,
367    Product,
368    Max(NanMode),
369    Min(NanMode),
370    /// `INT32` index of the maximum along the axis (lowest index on ties).
371    ArgMax(NanMode),
372}
373
374/// One assembled kernel variant. Everything that changes instructions is in the key; everything
375/// that changes only numbers is a specialization constant.
376#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
377pub enum KernelKey {
378    /// F32 activations times row-major packed E2M1 weights and E4M3 block scales.
379    Nvfp4Matmul {
380        buffers: u32,
381        cooperative: bool,
382        subgroup: bool,
383    },
384    /// Elementwise lanes over `count` output elements; `broadcast` selects the strided
385    /// multi-index addressing, otherwise every operand shares the output's linear index. `float`
386    /// is the storage of the operator's floating-point tensors (`Word` for FP32, `Half` for
387    /// FP16); `BOOL` lanes are byte storage in either variant. Binary16 lanes evaluate in
388    /// binary32 and narrow once, except the integer `NEGATE`/`ABS` sign lanes (ADR 0008).
389    Elementwise {
390        op: ElementwiseOp,
391        float: Storage,
392        broadcast: bool,
393        workgroup: u32,
394        buffers: u32,
395    },
396    /// Axis reduction: one invocation per output element, sequential ascending-axis fold. The
397    /// input is read at `float` storage; sums and products fold in binary32 (the TOSA
398    /// accumulator width), and the output is stored back at `float` storage.
399    Reduce {
400        op: ReduceOp,
401        float: Storage,
402        workgroup: u32,
403        buffers: u32,
404    },
405    /// Batched, register-tiled matrix multiplication: a `tile × tile` workgroup computes a
406    /// [`matmul_block`]-sided output square from workgroup-shared binary32 slabs. Both
407    /// operands are read at `input` storage and accumulate in binary32 — the accumulator width
408    /// TOSA assigns FP16 and FP8 MATMUL — and the result is stored at `output` storage. The two
409    /// differ only for the FP8 tier, where TOSA defines MATMUL as `(FP8, FP8) -> FP16`.
410    Matmul {
411        input: Storage,
412        output: Storage,
413        tile: u32,
414        buffers: u32,
415    },
416    /// Split-`k` streaming MATMUL for `m ≤` [`STREAM_ROWS`] rows:
417    /// a 1-D workgroup of [`STREAM_WORKGROUP`] invocations over [`STREAM_COLUMNS`] columns of
418    /// `rhs` (read at `rhs` storage, a word per invocation), the lhs read as binary32 words —
419    /// lowering widens a narrower lhs beforehand — and one fixed-order reduction of the `k`
420    /// slices at the end. Same accumulator width as [`Self::Matmul`].
421    MatmulStream {
422        rhs: Storage,
423        output: Storage,
424        buffers: u32,
425    },
426    /// NHWC max pooling with padding excluded from the window, at `float` storage.
427    MaxPool {
428        nan_mode: NanMode,
429        float: Storage,
430        workgroup: u32,
431        buffers: u32,
432    },
433    /// Elementwise float conversion: read at `input` storage, write at `output`. One dispatch
434    /// per `CAST` between float dtypes, including both FP8 directions (ADR 0009).
435    Cast {
436        input: Storage,
437        output: Storage,
438        workgroup: u32,
439        buffers: u32,
440    },
441    /// Strided copy over a rank-`MAX_RANK` iteration space (`TRANSPOSE`, `REVERSE`, `CONCAT`
442    /// segments); `contiguous` collapses to a linear copy.
443    Move {
444        storage: Storage,
445        contiguous: bool,
446        workgroup: u32,
447        buffers: u32,
448    },
449}
450
451impl KernelKey {
452    /// Assemble the module for this variant.
453    pub fn assemble(self) -> Vec<u32> {
454        match self {
455            Self::Nvfp4Matmul {
456                buffers,
457                cooperative: false,
458                subgroup: false,
459            } => assemble_nvfp4_matmul(buffers),
460            Self::Nvfp4Matmul {
461                buffers,
462                cooperative: true,
463                ..
464            } => assemble_nvfp4_matmul_cooperative(buffers),
465            Self::Nvfp4Matmul {
466                buffers,
467                cooperative: false,
468                subgroup: true,
469            } => assemble_nvfp4_matmul_subgroup(buffers),
470            Self::Elementwise {
471                op,
472                float,
473                broadcast,
474                workgroup,
475                buffers,
476            } => assemble_elementwise(op, float, broadcast, workgroup, buffers),
477            Self::Reduce {
478                op,
479                float,
480                workgroup,
481                buffers,
482            } => assemble_reduce(op, float, workgroup, buffers),
483            Self::Matmul {
484                input,
485                output,
486                tile,
487                buffers,
488            } => assemble_matmul(input, output, MatmulGeometry::wide(tile), buffers),
489            Self::MatmulStream {
490                rhs,
491                output,
492                buffers,
493            } => assemble_matmul_stream(rhs, output, buffers),
494            Self::MaxPool {
495                nan_mode,
496                float,
497                workgroup,
498                buffers,
499            } => assemble_max_pool(nan_mode, float, workgroup, buffers),
500            Self::Move {
501                storage,
502                contiguous,
503                workgroup,
504                buffers,
505            } => assemble_move(storage, contiguous, workgroup, buffers),
506            Self::Cast {
507                input,
508                output,
509                workgroup,
510                buffers,
511            } => assemble_cast(input, output, workgroup, buffers),
512        }
513    }
514
515    /// Every kernel variant this backend can assemble, at one representative tuning. The
516    /// SPIR-V validation sweep and the specialization-count test both walk this list, so a new
517    /// variant is validated the moment it is added here.
518    pub fn every_variant() -> Vec<KernelKey> {
519        let mut keys = Vec::new();
520        keys.push(KernelKey::Nvfp4Matmul {
521            buffers: 17,
522            cooperative: false,
523            subgroup: false,
524        });
525        keys.push(KernelKey::Nvfp4Matmul {
526            buffers: 17,
527            cooperative: true,
528            subgroup: false,
529        });
530        keys.push(KernelKey::Nvfp4Matmul {
531            buffers: 17,
532            cooperative: false,
533            subgroup: true,
534        });
535        let ops = [
536            ElementwiseOp::Abs,
537            ElementwiseOp::Ceil,
538            ElementwiseOp::Cos,
539            ElementwiseOp::Erf,
540            ElementwiseOp::Exp,
541            ElementwiseOp::Floor,
542            ElementwiseOp::Log,
543            ElementwiseOp::Negate,
544            ElementwiseOp::Reciprocal,
545            ElementwiseOp::Rsqrt,
546            ElementwiseOp::Sin,
547            ElementwiseOp::Sigmoid,
548            ElementwiseOp::Tanh,
549            ElementwiseOp::Clamp(NanMode::Propagate),
550            ElementwiseOp::Clamp(NanMode::Ignore),
551            ElementwiseOp::Add,
552            ElementwiseOp::Sub,
553            ElementwiseOp::Mul,
554            ElementwiseOp::Pow,
555            ElementwiseOp::Maximum(NanMode::Propagate),
556            ElementwiseOp::Minimum(NanMode::Ignore),
557            ElementwiseOp::Equal,
558            ElementwiseOp::Greater,
559            ElementwiseOp::GreaterEqual,
560            ElementwiseOp::LogicalAnd,
561            ElementwiseOp::LogicalOr,
562            ElementwiseOp::LogicalXor,
563            ElementwiseOp::LogicalNot,
564            ElementwiseOp::Select,
565            ElementwiseOp::CopyBytes,
566        ];
567        for op in ops {
568            for float in [Storage::Word, Storage::Half] {
569                for broadcast in [false, true] {
570                    keys.push(KernelKey::Elementwise {
571                        op,
572                        float,
573                        broadcast,
574                        workgroup: 64,
575                        buffers: 17,
576                    });
577                }
578            }
579        }
580        for op in [
581            ReduceOp::Sum,
582            ReduceOp::Product,
583            ReduceOp::Max(NanMode::Propagate),
584            ReduceOp::Min(NanMode::Ignore),
585            ReduceOp::ArgMax(NanMode::Propagate),
586            ReduceOp::ArgMax(NanMode::Ignore),
587        ] {
588            for float in [Storage::Word, Storage::Half] {
589                keys.push(KernelKey::Reduce {
590                    op,
591                    float,
592                    workgroup: 64,
593                    buffers: 17,
594                });
595            }
596        }
597        for float in [Storage::Word, Storage::Half] {
598            keys.push(KernelKey::Matmul {
599                input: float,
600                output: float,
601                tile: 16,
602                buffers: 17,
603            });
604            keys.push(KernelKey::Matmul {
605                input: float,
606                output: float,
607                tile: 8,
608                buffers: 5,
609            });
610        }
611        for float in [Storage::Word, Storage::Half] {
612            keys.push(KernelKey::MatmulStream {
613                rhs: float,
614                output: float,
615                buffers: 17,
616            });
617        }
618        // The lhs widening lowering inserts ahead of a skinny FP16 MATMUL.
619        keys.push(KernelKey::Cast {
620            input: Storage::Half,
621            output: Storage::Word,
622            workgroup: 64,
623            buffers: 17,
624        });
625        keys.push(KernelKey::Cast {
626            input: Storage::Word,
627            output: Storage::Half,
628            workgroup: 64,
629            buffers: 17,
630        });
631        // The FP8 tier: TOSA's `(FP8, FP8) -> FP16` MATMUL, and exact FP8 data movement.
632        for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
633            keys.push(KernelKey::MatmulStream {
634                rhs: Storage::Quarter(format),
635                output: Storage::Half,
636                buffers: 17,
637            });
638            keys.push(KernelKey::Matmul {
639                input: Storage::Quarter(format),
640                output: Storage::Half,
641                tile: 16,
642                buffers: 17,
643            });
644            keys.push(KernelKey::Matmul {
645                input: Storage::Quarter(format),
646                output: Storage::Half,
647                tile: 8,
648                buffers: 5,
649            });
650        }
651        for nan_mode in [NanMode::Propagate, NanMode::Ignore] {
652            for float in [Storage::Word, Storage::Half] {
653                keys.push(KernelKey::MaxPool {
654                    nan_mode,
655                    float,
656                    workgroup: 64,
657                    buffers: 17,
658                });
659            }
660        }
661        // The FP8 tier's pooling and ARGMAX: pooling selects an existing encoding, ARGMAX
662        // compares widened values and emits an INT32 index.
663        for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
664            for nan_mode in [NanMode::Propagate, NanMode::Ignore] {
665                keys.push(KernelKey::MaxPool {
666                    nan_mode,
667                    float: Storage::Quarter(format),
668                    workgroup: 64,
669                    buffers: 17,
670                });
671                keys.push(KernelKey::Reduce {
672                    op: ReduceOp::ArgMax(nan_mode),
673                    float: Storage::Quarter(format),
674                    workgroup: 64,
675                    buffers: 17,
676                });
677            }
678        }
679        // Every CAST pair the FP8 tier lowers, both directions.
680        for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
681            for wide in [Storage::Word, Storage::Half] {
682                keys.push(KernelKey::Cast {
683                    input: Storage::Quarter(format),
684                    output: wide,
685                    workgroup: 64,
686                    buffers: 17,
687                });
688                keys.push(KernelKey::Cast {
689                    input: wide,
690                    output: Storage::Quarter(format),
691                    workgroup: 64,
692                    buffers: 17,
693                });
694            }
695        }
696        for storage in [
697            Storage::Word,
698            Storage::Byte,
699            Storage::Half,
700            Storage::Quarter(Fp8Format::E4M3),
701            Storage::Quarter(Fp8Format::E5M2),
702        ] {
703            for contiguous in [false, true] {
704                keys.push(KernelKey::Move {
705                    storage,
706                    contiguous,
707                    workgroup: 64,
708                    buffers: 17,
709                });
710            }
711        }
712        keys
713    }
714
715    /// Number of specialization constants the module declares, ids `0..count`.
716    pub const fn spec_constant_count(self) -> u32 {
717        match self {
718            Self::Elementwise { op, broadcast, .. } => {
719                let inputs = op.inputs().len() as u32;
720                let base = 1 + 2 * inputs + 2;
721                let shape = if broadcast {
722                    MAX_RANK as u32 * (1 + inputs)
723                } else {
724                    0
725                };
726                base + shape + op.extra_spec_constants()
727            }
728            Self::Reduce { .. } => 2 + 2 + 3,
729            Self::Matmul { .. } | Self::MatmulStream { .. } => 3 * 2 + 4,
730            // Activation, six packed bindings, six scale bindings, tensor scale and output,
731            // followed by m/n/k/epilogue/weight mode.
732            Self::Nvfp4Matmul { .. } => 15 * 2 + 5,
733            Self::MaxPool { .. } => 2 + 2 + 12,
734            Self::Move { contiguous, .. } => {
735                if contiguous {
736                    2 + 2 + 1
737                } else {
738                    2 + 2 + 1 + MAX_RANK as u32 * 3 + 2
739                }
740            }
741            Self::Cast { .. } => 2 + 2 + 1,
742        }
743    }
744
745    /// `OpExecutionMode LocalSize` of the module.
746    pub const fn local_size(self) -> [u32; 3] {
747        match self {
748            Self::Elementwise { workgroup, .. }
749            | Self::Reduce { workgroup, .. }
750            | Self::MaxPool { workgroup, .. }
751            | Self::Move { workgroup, .. }
752            | Self::Cast { workgroup, .. } => [workgroup, 1, 1],
753            Self::Matmul { tile, .. } => MatmulGeometry::wide(tile).local_size(),
754            Self::MatmulStream { .. } => [STREAM_WORKGROUP, 1, 1],
755            Self::Nvfp4Matmul {
756                cooperative: false,
757                subgroup: false,
758                ..
759            } => [STREAM_WORKGROUP, 1, 1],
760            Self::Nvfp4Matmul {
761                cooperative: true, ..
762            } => [32, 1, 1],
763            Self::Nvfp4Matmul {
764                cooperative: false,
765                subgroup: true,
766                ..
767            } => [32, 1, 1],
768        }
769    }
770}
771
772/// Where a kernel operand lives: a descriptor-array element and a base offset in 32-bit words.
773#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
774pub struct Operand {
775    pub buffer: u32,
776    pub base: u32,
777}
778
779impl Operand {
780    fn push(self, words: &mut Vec<u32>) {
781        words.push(self.buffer);
782        words.push(self.base);
783    }
784}
785
786/// Specialization payload of an elementwise dispatch. The declaration order in the elementwise
787/// kernel is: `count`, each input operand, the output operand, then (broadcast
788/// only) the output dims and each input's element strides, then the clamp bounds.
789#[derive(Clone, Debug, PartialEq, Eq)]
790pub struct ElementwiseSpec<'a> {
791    pub count: u32,
792    pub inputs: &'a [Operand],
793    pub output: Operand,
794    /// Output shape padded with leading ones to `MAX_RANK`.
795    pub dims: [u32; MAX_RANK],
796    /// Per input, element strides over `dims` (0 along a broadcast dimension).
797    pub strides: &'a [[u32; MAX_RANK]],
798    /// `CLAMP` bounds as FP32 bit patterns.
799    pub clamp: Option<[u32; 2]>,
800}
801
802impl ElementwiseSpec<'_> {
803    pub fn words(&self, broadcast: bool) -> Vec<u32> {
804        let mut words = vec![self.count];
805        for input in self.inputs {
806            input.push(&mut words);
807        }
808        self.output.push(&mut words);
809        if broadcast {
810            words.extend_from_slice(&self.dims);
811            for strides in self.strides {
812                words.extend_from_slice(strides);
813            }
814        }
815        if let Some(bounds) = self.clamp {
816            words.extend_from_slice(&bounds);
817        }
818        words
819    }
820}
821
822/// Specialization payload of a reduction: input, output, `outer`, `axis`, `inner`.
823pub fn reduce_spec(input: Operand, output: Operand, outer: u32, axis: u32, inner: u32) -> Vec<u32> {
824    let mut words = Vec::with_capacity(7);
825    input.push(&mut words);
826    output.push(&mut words);
827    words.extend_from_slice(&[outer, axis, inner]);
828    words
829}
830
831/// Specialization payload of a MATMUL: `lhs`, `rhs`, output, then `m`, `n`, `k`, `batch`.
832pub fn matmul_spec(
833    lhs: Operand,
834    rhs: Operand,
835    output: Operand,
836    m: u32,
837    n: u32,
838    k: u32,
839    batch: u32,
840) -> Vec<u32> {
841    let mut words = Vec::with_capacity(10);
842    lhs.push(&mut words);
843    rhs.push(&mut words);
844    output.push(&mut words);
845    words.extend_from_slice(&[m, n, k, batch]);
846    words
847}
848
849/// Specialization payload for native row-major NVFP4: activation, packed
850/// weights, block scales, tensor scale, output, then `m`, `n`, `k`, epilogue.
851pub struct Nvfp4MatmulSpec<'a> {
852    pub activation: Operand,
853    pub packed: &'a [Operand],
854    pub block_scales: &'a [Operand],
855    pub tensor_scale: Operand,
856    pub output: Operand,
857    pub m: u32,
858    pub n: u32,
859    pub k: u32,
860    pub epilogue: u32,
861    pub weight_mode: u32,
862}
863
864impl Nvfp4MatmulSpec<'_> {
865    pub fn words(&self) -> Vec<u32> {
866        assert!(
867            !self.packed.is_empty()
868                && self.packed.len() <= 6
869                && self.packed.len() == self.block_scales.len()
870        );
871        let mut words = Vec::with_capacity(35);
872        self.activation.push(&mut words);
873        for index in 0..6 {
874            self.packed[index.min(self.packed.len() - 1)].push(&mut words);
875        }
876        for index in 0..6 {
877            self.block_scales[index.min(self.block_scales.len() - 1)].push(&mut words);
878        }
879        self.tensor_scale.push(&mut words);
880        self.output.push(&mut words);
881        words.extend_from_slice(&[self.m, self.n, self.k, self.epilogue, self.weight_mode]);
882        words
883    }
884}
885
886/// NHWC pooling geometry: batch, input height/width, channels, output height/width, kernel,
887/// stride, and the top/left pads (the bottom/right pads only shape the output).
888#[derive(Clone, Copy, Debug, PartialEq, Eq)]
889pub struct PoolGeometry {
890    pub batch: u32,
891    pub height: u32,
892    pub width: u32,
893    pub channels: u32,
894    pub out_height: u32,
895    pub out_width: u32,
896    pub kernel: [u32; 2],
897    pub stride: [u32; 2],
898    pub pad_top: u32,
899    pub pad_left: u32,
900}
901
902/// Specialization payload of a MAX_POOL2D dispatch.
903pub fn max_pool_spec(input: Operand, output: Operand, geometry: PoolGeometry) -> Vec<u32> {
904    let mut words = Vec::with_capacity(16);
905    input.push(&mut words);
906    output.push(&mut words);
907    words.extend_from_slice(&[
908        geometry.batch,
909        geometry.height,
910        geometry.width,
911        geometry.channels,
912        geometry.out_height,
913        geometry.out_width,
914        geometry.kernel[0],
915        geometry.kernel[1],
916        geometry.stride[0],
917        geometry.stride[1],
918        geometry.pad_top,
919        geometry.pad_left,
920    ]);
921    words
922}
923
924/// Iteration space of a strided copy: `dims` (leading ones to `MAX_RANK`), element strides and
925/// element offsets on each side. Strides are wrapping `u32` so a reversed axis is `-inner`.
926#[derive(Clone, Copy, Debug, PartialEq, Eq)]
927pub struct MoveGeometry {
928    pub count: u32,
929    pub dims: [u32; MAX_RANK],
930    pub in_strides: [u32; MAX_RANK],
931    pub in_offset: u32,
932    pub out_strides: [u32; MAX_RANK],
933    pub out_offset: u32,
934}
935
936/// Specialization payload of a copy dispatch.
937pub fn move_spec(
938    input: Operand,
939    output: Operand,
940    geometry: MoveGeometry,
941    contiguous: bool,
942) -> Vec<u32> {
943    let mut words = Vec::with_capacity(25);
944    input.push(&mut words);
945    output.push(&mut words);
946    words.push(geometry.count);
947    if !contiguous {
948        words.extend_from_slice(&geometry.dims);
949        words.extend_from_slice(&geometry.in_strides);
950        words.push(geometry.in_offset);
951        words.extend_from_slice(&geometry.out_strides);
952        words.push(geometry.out_offset);
953    }
954    words
955}
956
957/// Number of workgroups covering `count` items at `workgroup` invocations each, capped at
958/// `limit`; the kernels loop with a grid stride so the cap only trades parallelism, never
959/// coverage.
960pub const fn linear_workgroups(count: u32, workgroup: u32, limit: u32) -> u32 {
961    let needed = count.div_ceil(workgroup);
962    if needed == 0 {
963        1
964    } else if needed > limit {
965        limit
966    } else {
967        needed
968    }
969}
970
971/// Outputs per invocation per side of the register-tiled MATMUL: each invocation of a
972/// `tile × tile` workgroup accumulates a `MATMUL_MICRO × MATMUL_MICRO` block, so the workgroup
973/// covers a [`matmul_block`]-sided square of the result.
974pub const MATMUL_MICRO: u32 = 4;
975
976/// Rows the streaming MATMUL kernel carries per invocation, and the row count at or below which
977/// lowering selects it over the register-tiled kernel.
978pub const STREAM_ROWS: u32 = 8;
979
980/// Output columns one streaming MATMUL workgroup covers.
981pub const STREAM_COLUMNS: u32 = 16;
982
983/// Invocations of a streaming MATMUL workgroup: `STREAM_COLUMNS / lanes` weight words times
984/// `STREAM_WORKGROUP / that` slices of `k`.
985pub const STREAM_WORKGROUP: u32 = 64;
986
987/// Bytes of workgroup-shared memory the streaming kernel declares for its final reduction:
988/// every invocation's `STREAM_ROWS × lanes` partial sums, at the widest lane count.
989pub const fn stream_matmul_shared_bytes() -> u32 {
990    STREAM_WORKGROUP * STREAM_ROWS * 4 * 4
991}
992
993/// Workgroup counts of a streaming MATMUL over `n` columns and `batch` batches.
994pub const fn stream_matmul_workgroups(n: u32, batch: u32) -> [u32; 3] {
995    [n.div_ceil(STREAM_COLUMNS), 1, batch]
996}
997
998/// Side of the output square one MATMUL workgroup of `tile × tile` invocations computes.
999pub const fn matmul_block(tile: u32) -> u32 {
1000    tile * MATMUL_MICRO
1001}
1002
1003/// Bytes of workgroup-shared memory the MATMUL kernels at `tile` declare, whichever kernel
1004/// needs more.
1005pub const fn matmul_shared_bytes(tile: u32) -> u32 {
1006    let wide = MatmulGeometry::wide(tile).shared_bytes();
1007    let stream = stream_matmul_shared_bytes();
1008    if wide > stream { wide } else { stream }
1009}
1010
1011/// Workgroup counts of a tiled MATMUL over `m` rows, `n` columns, and `batch` batches.
1012pub const fn matmul_workgroups(m: u32, n: u32, batch: u32, tile: u32) -> [u32; 3] {
1013    let block = matmul_block(tile);
1014    [n.div_ceil(block), m.div_ceil(block), batch]
1015}
1016
1017/// Shape of one MATMUL workgroup: `tile_x × tile_y` invocations, each accumulating a
1018/// `micro_m × micro_n` register block, over shared-memory slabs `depth` deep in `k`. Every
1019/// dimension is a power of two and the two slabs stage evenly over the invocations.
1020#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1021pub struct MatmulGeometry {
1022    pub tile_x: u32,
1023    pub tile_y: u32,
1024    pub micro_m: u32,
1025    pub micro_n: u32,
1026    pub depth: u32,
1027}
1028
1029impl MatmulGeometry {
1030    /// The square register-tiled geometry: `tile × tile` invocations, [`MATMUL_MICRO`]² each,
1031    /// `tile` deep — a 64 × 64 block at the preferred tile.
1032    pub const fn wide(tile: u32) -> Self {
1033        Self {
1034            tile_x: tile,
1035            tile_y: tile,
1036            micro_m: MATMUL_MICRO,
1037            micro_n: MATMUL_MICRO,
1038            depth: tile,
1039        }
1040    }
1041
1042    pub const fn invocations(self) -> u32 {
1043        self.tile_x * self.tile_y
1044    }
1045
1046    pub const fn block_m(self) -> u32 {
1047        self.tile_y * self.micro_m
1048    }
1049
1050    pub const fn block_n(self) -> u32 {
1051        self.tile_x * self.micro_n
1052    }
1053
1054    pub const fn local_size(self) -> [u32; 3] {
1055        [self.tile_x, self.tile_y, 1]
1056    }
1057
1058    /// Two binary32 slabs: lhs `block_m × depth` (stored transposed) and rhs `depth × block_n`.
1059    pub const fn shared_bytes(self) -> u32 {
1060        (self.block_m() + self.block_n()) * self.depth * 4
1061    }
1062}
1063
1064// ---------------------------------------------------------------------------------------------
1065// Binary16 host conversions
1066// ---------------------------------------------------------------------------------------------
1067
1068/// The exact binary32 value of a binary16 bit pattern (host side of the kernels' unpack: every
1069/// binary16 value, subnormals included, is exactly representable in binary32; NaN payloads are
1070/// preserved).
1071/// Narrow binary32 to one FP8 encoding with round-to-nearest, ties-to-even.
1072///
1073/// **Overflow policy.** TOSA 1.0 through 1.2 leave float-to-FP8 overflow undefined, so this is
1074/// the crate's policy rather than the spec's. A magnitude too large for E4M3 — which has no
1075/// infinity — becomes NaN, not the finite maximum. The reason is that the alternatives are not
1076/// symmetric: a consumer who wants saturation can `CLAMP` in a wider dtype before the `CAST`
1077/// and get it exactly, while a consumer handed a saturated 448 cannot tell it from a value that
1078/// was always 448. NaN preserves the choice; saturation destroys it. It also matches what this
1079/// crate already does one format up, where `narrow_f16` signals unrepresentability as infinity
1080/// rather than clamping to 65504.
1081///
1082/// E5M2 needs no policy: it is IEEE-shaped, so overflow becomes infinity like any binary float.
1083/// NaN in becomes a canonical quiet NaN out, sign preserved, for both encodings.
1084pub fn f32_to_fp8_bits(format: Fp8Format, value: f32) -> u8 {
1085    let (mantissa_bits, bias, nan_out, max_finite): (u32, u32, u8, u32) = match format {
1086        // E4M3's `0x7f` is its only NaN, so the finite encodings stop at `0x7e` (448).
1087        Fp8Format::E4M3 => (3, 7, 0x7f, 0x7e),
1088        // E5M2's `0x7c` is infinity; finite encodings stop at `0x7b` (57344).
1089        Fp8Format::E5M2 => (2, 15, 0x7e, 0x7b),
1090    };
1091    let bits = value.to_bits();
1092    let sign = ((bits >> 24) as u8) & 0x80;
1093    let magnitude = bits & 0x7fff_ffff;
1094    if magnitude > 0x7f80_0000 {
1095        return sign | nan_out;
1096    }
1097    let shift = 23 - mantissa_bits;
1098    let normal_floor = (128 - bias) << 23;
1099    let body = if magnitude >= normal_floor {
1100        // Rebias into the target's exponent range, then round the significand with the
1101        // add-half-ulp-plus-guard trick; a significand carry increments the exponent by itself.
1102        let adjusted = magnitude - ((127 - bias) << 23);
1103        let guard = (adjusted >> shift) & 1;
1104        let half_ulp = (1 << (shift - 1)) - 1;
1105        (adjusted + half_ulp + guard) >> shift
1106    } else {
1107        // Subnormal or zero: scale into integer range — exact for every representable
1108        // subnormal — and round to the nearest integer, ties to even. A rounding carry reaches
1109        // the smallest normal encoding on its own.
1110        let scale = f32::from_bits((127 + bias - 1 + mantissa_bits) << 23);
1111        (f32::from_bits(magnitude) * scale).round_ties_even() as u32
1112    };
1113    // Rounding has already happened, so landing above the last finite encoding *is* the
1114    // overflow test; infinity in reaches it too.
1115    if body > max_finite {
1116        let overflow = match format {
1117            Fp8Format::E4M3 => nan_out,
1118            Fp8Format::E5M2 => 0x7c,
1119        };
1120        return sign | overflow;
1121    }
1122    sign | body as u8
1123}
1124
1125pub fn f16_to_f32(bits: u16) -> f32 {
1126    let bits = u32::from(bits);
1127    let sign = (bits & 0x8000) << 16;
1128    let exponent = (bits >> 10) & 0x1f;
1129    let mantissa = bits & 0x3ff;
1130    let converted = if exponent == 0 {
1131        // Subnormal or zero: `mantissa · 2^-24`, exact for every 10-bit mantissa.
1132        (mantissa as f32) * (1.0 / 16_777_216.0)
1133    } else if exponent == 31 {
1134        f32::from_bits(0x7f80_0000 | (mantissa << 13))
1135    } else {
1136        f32::from_bits(((exponent + 112) << 23) | (mantissa << 13))
1137    };
1138    f32::from_bits(converted.to_bits() | sign)
1139}
1140
1141/// The binary16 bit pattern nearest to `value`, round-to-nearest-even. NaN is canonicalized to
1142/// the quiet `0x7e00` payload with its sign;
1143/// magnitudes at or above 65520 round to infinity, and subnormals are produced, never flushed).
1144pub fn f32_to_f16_bits(value: f32) -> u16 {
1145    let bits = value.to_bits();
1146    let sign = ((bits >> 16) & 0x8000) as u16;
1147    let magnitude = bits & 0x7fff_ffff;
1148    if magnitude > 0x7f80_0000 {
1149        return sign | 0x7e00;
1150    }
1151    if magnitude >= 0x477f_f000 {
1152        // At or above 65520 — the midpoint between 65504 and the overflow binade — round to
1153        // infinity (65504 itself is `0x477f_e000`, finite).
1154        return sign | 0x7c00;
1155    }
1156    if magnitude >= 0x3880_0000 {
1157        // Normal binary16: rebias the exponent, then round the significand with the
1158        // add-half-ulp-plus-guard trick; a mantissa carry increments the exponent on its own.
1159        let adjusted = magnitude - 0x3800_0000;
1160        let rounded = adjusted + 0x0000_0fff + ((adjusted >> 13) & 1);
1161        return sign | (rounded >> 13) as u16;
1162    }
1163    // Subnormal binary16 or zero: scale into integer range (exact while the value is at or
1164    // above 2^-25; anything smaller rounds to zero either way) and round to the nearest integer.
1165    let scaled = f32::from_bits(magnitude) * 16_777_216.0;
1166    let truncated = scaled as u16;
1167    let fraction = scaled - f32::from(truncated);
1168    let rounded = if fraction > 0.5 || (fraction == 0.5 && truncated & 1 != 0) {
1169        truncated + 1
1170    } else {
1171        truncated
1172    };
1173    sign | rounded
1174}
1175
1176// ---------------------------------------------------------------------------------------------
1177// Module builder
1178// ---------------------------------------------------------------------------------------------
1179
1180#[derive(Clone, PartialEq, Eq, Hash)]
1181enum TypeKey {
1182    Void,
1183    Bool,
1184    U32,
1185    F32,
1186    F16,
1187    Vector(Id, u32),
1188    Pointer(u32, Id),
1189    RuntimeArray(Id),
1190    Array(Id, Id),
1191    Function(Id),
1192}
1193
1194#[derive(Clone, Copy, PartialEq, Eq, Hash)]
1195enum ConstKey {
1196    U32(u32),
1197    F32(u32),
1198    False,
1199}
1200
1201/// Section-ordered SPIR-V assembler. Word counts and the id bound are derived; every instruction
1202/// is emitted opcode-first exactly as the specification tabulates it.
1203struct Builder {
1204    next_id: Id,
1205    capabilities: Vec<u32>,
1206    extensions: Vec<u32>,
1207    imports: Vec<u32>,
1208    memory_model: Vec<u32>,
1209    entry_point: Vec<u32>,
1210    execution_modes: Vec<u32>,
1211    annotations: Vec<u32>,
1212    declarations: Vec<u32>,
1213    functions: Vec<u32>,
1214    types: HashMap<TypeKey, Id>,
1215    constants: HashMap<ConstKey, Id>,
1216    glsl: Id,
1217    spec_next: u32,
1218    /// Where the entry block's `OpVariable Function` declarations are spliced in.
1219    local_variable_cursor: usize,
1220    interface: Vec<Id>,
1221    main: Id,
1222}
1223
1224fn instruction(target: &mut Vec<u32>, opcode: u16, operands: &[u32]) {
1225    let word_count = u32::try_from(operands.len() + 1).expect("instruction fits");
1226    target.push((word_count << 16) | u32::from(opcode));
1227    target.extend_from_slice(operands);
1228}
1229
1230/// Encode a literal string operand: UTF-8 bytes, NUL terminated, zero-padded to whole words.
1231fn literal_string(text: &str) -> Vec<u32> {
1232    let mut bytes = text.as_bytes().to_vec();
1233    bytes.push(0);
1234    while bytes.len() % 4 != 0 {
1235        bytes.push(0);
1236    }
1237    bytes
1238        .chunks_exact(4)
1239        .map(|chunk| u32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]))
1240        .collect()
1241}
1242
1243impl Builder {
1244    fn new() -> Self {
1245        let mut builder = Self {
1246            next_id: 1,
1247            capabilities: Vec::new(),
1248            extensions: Vec::new(),
1249            imports: Vec::new(),
1250            memory_model: Vec::new(),
1251            entry_point: Vec::new(),
1252            execution_modes: Vec::new(),
1253            annotations: Vec::new(),
1254            declarations: Vec::new(),
1255            functions: Vec::new(),
1256            types: HashMap::new(),
1257            constants: HashMap::new(),
1258            glsl: 0,
1259            spec_next: 0,
1260            local_variable_cursor: 0,
1261            interface: Vec::new(),
1262            main: 0,
1263        };
1264        instruction(
1265            &mut builder.capabilities,
1266            OP_CAPABILITY,
1267            &[CAPABILITY_SHADER],
1268        );
1269        builder.glsl = builder.id();
1270        let mut import = vec![builder.glsl];
1271        import.extend(literal_string("GLSL.std.450"));
1272        instruction(&mut builder.imports, OP_EXT_INST_IMPORT, &import);
1273        instruction(
1274            &mut builder.memory_model,
1275            OP_MEMORY_MODEL,
1276            &[ADDRESSING_MODEL_LOGICAL, MEMORY_MODEL_GLSL450],
1277        );
1278        builder.main = builder.id();
1279        builder
1280    }
1281
1282    fn id(&mut self) -> Id {
1283        let id = self.next_id;
1284        self.next_id += 1;
1285        id
1286    }
1287
1288    fn finish(mut self, local_size: [u32; 3]) -> Vec<u32> {
1289        let mut entry = vec![EXECUTION_MODEL_GL_COMPUTE, self.main];
1290        entry.extend(literal_string("main"));
1291        entry.extend_from_slice(&self.interface);
1292        instruction(&mut self.entry_point, OP_ENTRY_POINT, &entry);
1293        instruction(
1294            &mut self.execution_modes,
1295            OP_EXECUTION_MODE,
1296            &[
1297                self.main,
1298                EXECUTION_MODE_LOCAL_SIZE,
1299                local_size[0],
1300                local_size[1],
1301                local_size[2],
1302            ],
1303        );
1304        let mut words = vec![SPIRV_MAGIC, SPIRV_VERSION_1_3, 0, self.next_id, 0];
1305        for section in [
1306            &self.capabilities,
1307            &self.extensions,
1308            &self.imports,
1309            &self.memory_model,
1310            &self.entry_point,
1311            &self.execution_modes,
1312            &self.annotations,
1313            &self.declarations,
1314            &self.functions,
1315        ] {
1316            words.extend_from_slice(section);
1317        }
1318        words
1319    }
1320
1321    // -- types and constants ------------------------------------------------------------------
1322
1323    fn ty(&mut self, key: TypeKey) -> Id {
1324        if let Some(id) = self.types.get(&key) {
1325            return *id;
1326        }
1327        let id = self.id();
1328        match &key {
1329            TypeKey::Void => instruction(&mut self.declarations, OP_TYPE_VOID, &[id]),
1330            TypeKey::Bool => instruction(&mut self.declarations, OP_TYPE_BOOL, &[id]),
1331            TypeKey::U32 => instruction(&mut self.declarations, OP_TYPE_INT, &[id, 32, 0]),
1332            TypeKey::F32 => instruction(&mut self.declarations, OP_TYPE_FLOAT, &[id, 32]),
1333            TypeKey::F16 => instruction(&mut self.declarations, OP_TYPE_FLOAT, &[id, 16]),
1334            TypeKey::Vector(element, count) => {
1335                instruction(
1336                    &mut self.declarations,
1337                    OP_TYPE_VECTOR,
1338                    &[id, *element, *count],
1339                );
1340            }
1341            TypeKey::Pointer(class, pointee) => {
1342                instruction(
1343                    &mut self.declarations,
1344                    OP_TYPE_POINTER,
1345                    &[id, *class, *pointee],
1346                );
1347            }
1348            TypeKey::RuntimeArray(element) => {
1349                instruction(
1350                    &mut self.declarations,
1351                    OP_TYPE_RUNTIME_ARRAY,
1352                    &[id, *element],
1353                );
1354            }
1355            TypeKey::Array(element, length) => {
1356                instruction(
1357                    &mut self.declarations,
1358                    OP_TYPE_ARRAY,
1359                    &[id, *element, *length],
1360                );
1361            }
1362            TypeKey::Function(ret) => {
1363                instruction(&mut self.declarations, OP_TYPE_FUNCTION, &[id, *ret]);
1364            }
1365        }
1366        self.types.insert(key, id);
1367        id
1368    }
1369
1370    fn void(&mut self) -> Id {
1371        self.ty(TypeKey::Void)
1372    }
1373    fn bool_ty(&mut self) -> Id {
1374        self.ty(TypeKey::Bool)
1375    }
1376    fn u32_ty(&mut self) -> Id {
1377        self.ty(TypeKey::U32)
1378    }
1379    fn f32_ty(&mut self) -> Id {
1380        self.ty(TypeKey::F32)
1381    }
1382    fn f16_ty(&mut self) -> Id {
1383        self.ty(TypeKey::F16)
1384    }
1385    fn uvec3(&mut self) -> Id {
1386        let u32_ty = self.u32_ty();
1387        self.ty(TypeKey::Vector(u32_ty, 3))
1388    }
1389    fn pointer(&mut self, class: u32, pointee: Id) -> Id {
1390        self.ty(TypeKey::Pointer(class, pointee))
1391    }
1392
1393    fn constant(&mut self, key: ConstKey) -> Id {
1394        if let Some(id) = self.constants.get(&key) {
1395            return *id;
1396        }
1397        let id = self.id();
1398        match key {
1399            ConstKey::U32(value) => {
1400                let ty = self.u32_ty();
1401                instruction(&mut self.declarations, OP_CONSTANT, &[ty, id, value]);
1402            }
1403            ConstKey::F32(bits) => {
1404                let ty = self.f32_ty();
1405                instruction(&mut self.declarations, OP_CONSTANT, &[ty, id, bits]);
1406            }
1407            ConstKey::False => {
1408                let ty = self.bool_ty();
1409                instruction(&mut self.declarations, OP_CONSTANT_FALSE, &[ty, id]);
1410            }
1411        }
1412        self.constants.insert(key, id);
1413        id
1414    }
1415
1416    fn c_u32(&mut self, value: u32) -> Id {
1417        self.constant(ConstKey::U32(value))
1418    }
1419    fn c_f32(&mut self, value: f32) -> Id {
1420        self.constant(ConstKey::F32(value.to_bits()))
1421    }
1422    fn c_false(&mut self) -> Id {
1423        self.constant(ConstKey::False)
1424    }
1425
1426    /// Declare the next `u32` specialization constant (ids are assigned in declaration order).
1427    fn spec_u32(&mut self, default: u32) -> Id {
1428        let ty = self.u32_ty();
1429        let id = self.id();
1430        instruction(&mut self.declarations, OP_SPEC_CONSTANT, &[ty, id, default]);
1431        instruction(
1432            &mut self.annotations,
1433            OP_DECORATE,
1434            &[id, DECORATION_SPEC_ID, self.spec_next],
1435        );
1436        self.spec_next += 1;
1437        id
1438    }
1439
1440    fn spec_operand(&mut self) -> (Id, Id) {
1441        let buffer = self.spec_u32(0);
1442        let base = self.spec_u32(0);
1443        (buffer, base)
1444    }
1445
1446    fn spec_dims(&mut self) -> [Id; MAX_RANK] {
1447        let mut ids = [0; MAX_RANK];
1448        for id in &mut ids {
1449            *id = self.spec_u32(1);
1450        }
1451        ids
1452    }
1453
1454    fn spec_strides(&mut self) -> [Id; MAX_RANK] {
1455        let mut ids = [0; MAX_RANK];
1456        for id in &mut ids {
1457            *id = self.spec_u32(0);
1458        }
1459        ids
1460    }
1461
1462    // -- global variables ---------------------------------------------------------------------
1463
1464    fn builtin_uvec3(&mut self, builtin: u32) -> Id {
1465        let uvec3 = self.uvec3();
1466        let pointer = self.pointer(STORAGE_CLASS_INPUT, uvec3);
1467        let id = self.id();
1468        instruction(
1469            &mut self.declarations,
1470            OP_VARIABLE,
1471            &[pointer, id, STORAGE_CLASS_INPUT],
1472        );
1473        instruction(
1474            &mut self.annotations,
1475            OP_DECORATE,
1476            &[id, DECORATION_BUILT_IN, builtin],
1477        );
1478        self.interface.push(id);
1479        id
1480    }
1481
1482    /// The descriptor: set 0, binding 0, `buffers` storage-buffer blocks of `{ uint words[]; }`.
1483    fn buffer_array(&mut self, buffers: u32) -> Id {
1484        let u32_ty = self.u32_ty();
1485        let words = self.ty(TypeKey::RuntimeArray(u32_ty));
1486        instruction(
1487            &mut self.annotations,
1488            OP_DECORATE,
1489            &[words, DECORATION_ARRAY_STRIDE, 4],
1490        );
1491        let block = self.id();
1492        instruction(&mut self.declarations, OP_TYPE_STRUCT, &[block, words]);
1493        instruction(
1494            &mut self.annotations,
1495            OP_DECORATE,
1496            &[block, DECORATION_BLOCK],
1497        );
1498        instruction(
1499            &mut self.annotations,
1500            OP_MEMBER_DECORATE,
1501            &[block, 0, DECORATION_OFFSET, 0],
1502        );
1503        let length = self.c_u32(buffers);
1504        let array = self.ty(TypeKey::Array(block, length));
1505        let pointer = self.pointer(STORAGE_CLASS_STORAGE_BUFFER, array);
1506        let variable = self.id();
1507        instruction(
1508            &mut self.declarations,
1509            OP_VARIABLE,
1510            &[pointer, variable, STORAGE_CLASS_STORAGE_BUFFER],
1511        );
1512        instruction(
1513            &mut self.annotations,
1514            OP_DECORATE,
1515            &[variable, DECORATION_DESCRIPTOR_SET, 0],
1516        );
1517        instruction(
1518            &mut self.annotations,
1519            OP_DECORATE,
1520            &[variable, DECORATION_BINDING, 0],
1521        );
1522        variable
1523    }
1524
1525    /// A `Private` array of `u32` initialized with `values`; indexable at run time.
1526    fn private_u32_array(&mut self, values: &[u32]) -> Id {
1527        let u32_ty = self.u32_ty();
1528        let length = self.c_u32(values.len() as u32);
1529        let array = self.ty(TypeKey::Array(u32_ty, length));
1530        let elements: Vec<Id> = values.iter().map(|value| self.c_u32(*value)).collect();
1531        let composite = self.id();
1532        let mut operands = vec![array, composite];
1533        operands.extend_from_slice(&elements);
1534        instruction(&mut self.declarations, OP_CONSTANT_COMPOSITE, &operands);
1535        let pointer = self.pointer(STORAGE_CLASS_PRIVATE, array);
1536        let variable = self.id();
1537        instruction(
1538            &mut self.declarations,
1539            OP_VARIABLE,
1540            &[pointer, variable, STORAGE_CLASS_PRIVATE, composite],
1541        );
1542        variable
1543    }
1544
1545    /// A workgroup-shared `float[length]` array.
1546    fn shared_f32_array(&mut self, length: u32) -> Id {
1547        let f32_ty = self.f32_ty();
1548        let length = self.c_u32(length);
1549        let array = self.ty(TypeKey::Array(f32_ty, length));
1550        let pointer = self.pointer(STORAGE_CLASS_WORKGROUP, array);
1551        let variable = self.id();
1552        instruction(
1553            &mut self.declarations,
1554            OP_VARIABLE,
1555            &[pointer, variable, STORAGE_CLASS_WORKGROUP],
1556        );
1557        variable
1558    }
1559
1560    /// A workgroup-shared binary16 array used as a cooperative-matrix staging tile.
1561    fn shared_f16_array(&mut self, length: u32) -> Id {
1562        let f16_ty = self.f16_ty();
1563        let length = self.c_u32(length);
1564        let array = self.ty(TypeKey::Array(f16_ty, length));
1565        let pointer = self.pointer(STORAGE_CLASS_WORKGROUP, array);
1566        let variable = self.id();
1567        instruction(
1568            &mut self.declarations,
1569            OP_VARIABLE,
1570            &[pointer, variable, STORAGE_CLASS_WORKGROUP],
1571        );
1572        variable
1573    }
1574
1575    fn enable_cooperative_matrix(&mut self) {
1576        instruction(&mut self.capabilities, OP_CAPABILITY, &[CAPABILITY_FLOAT16]);
1577        instruction(
1578            &mut self.capabilities,
1579            OP_CAPABILITY,
1580            &[CAPABILITY_COOPERATIVE_MATRIX_KHR],
1581        );
1582        instruction(
1583            &mut self.capabilities,
1584            OP_CAPABILITY,
1585            &[CAPABILITY_VULKAN_MEMORY_MODEL],
1586        );
1587        let memory_extension = literal_string("SPV_KHR_vulkan_memory_model");
1588        instruction(&mut self.extensions, OP_EXTENSION, &memory_extension);
1589        let extension = literal_string("SPV_KHR_cooperative_matrix");
1590        instruction(&mut self.extensions, OP_EXTENSION, &extension);
1591        self.memory_model.clear();
1592        instruction(
1593            &mut self.memory_model,
1594            OP_MEMORY_MODEL,
1595            &[ADDRESSING_MODEL_LOGICAL, 3], // VulkanKHR
1596        );
1597    }
1598
1599    fn enable_subgroup_arithmetic(&mut self) {
1600        instruction(
1601            &mut self.capabilities,
1602            OP_CAPABILITY,
1603            &[CAPABILITY_GROUP_NON_UNIFORM_ARITHMETIC],
1604        );
1605    }
1606
1607    fn cooperative_matrix_ty(&mut self, component: Id, rows: u32, columns: u32, usage: u32) -> Id {
1608        let ty = self.id();
1609        let scope = self.c_u32(3); // Scope Subgroup
1610        let rows = self.c_u32(rows);
1611        let columns = self.c_u32(columns);
1612        let usage = self.c_u32(usage);
1613        instruction(
1614            &mut self.declarations,
1615            OP_TYPE_COOPERATIVE_MATRIX_KHR,
1616            &[ty, component, scope, rows, columns, usage],
1617        );
1618        ty
1619    }
1620
1621    // -- function body ------------------------------------------------------------------------
1622
1623    fn begin_main(&mut self) {
1624        let void = self.void();
1625        let fn_type = self.ty(TypeKey::Function(void));
1626        instruction(
1627            &mut self.functions,
1628            OP_FUNCTION,
1629            &[void, self.main, FUNCTION_CONTROL_NONE, fn_type],
1630        );
1631        let entry = self.id();
1632        instruction(&mut self.functions, OP_LABEL, &[entry]);
1633        self.local_variable_cursor = self.functions.len();
1634    }
1635
1636    fn end_main(&mut self) {
1637        instruction(&mut self.functions, OP_RETURN, &[]);
1638        instruction(&mut self.functions, OP_FUNCTION_END, &[]);
1639    }
1640
1641    /// A function-scope variable of `ty`, declared at the top of the entry block.
1642    fn local(&mut self, ty: Id) -> Id {
1643        let pointer = self.pointer(STORAGE_CLASS_FUNCTION, ty);
1644        let id = self.id();
1645        let mut declaration = Vec::with_capacity(4);
1646        instruction(
1647            &mut declaration,
1648            OP_VARIABLE,
1649            &[pointer, id, STORAGE_CLASS_FUNCTION],
1650        );
1651        let cursor = self.local_variable_cursor;
1652        self.functions
1653            .splice(cursor..cursor, declaration.iter().copied());
1654        self.local_variable_cursor += declaration.len();
1655        id
1656    }
1657
1658    fn emit(&mut self, opcode: u16, operands: &[u32]) {
1659        instruction(&mut self.functions, opcode, operands);
1660    }
1661
1662    fn value(&mut self, opcode: u16, ty: Id, operands: &[u32]) -> Id {
1663        let id = self.id();
1664        let mut all = Vec::with_capacity(operands.len() + 2);
1665        all.push(ty);
1666        all.push(id);
1667        all.extend_from_slice(operands);
1668        instruction(&mut self.functions, opcode, &all);
1669        id
1670    }
1671
1672    fn float_value(&mut self, opcode: u16, operands: &[u32]) -> Id {
1673        let f32_ty = self.f32_ty();
1674        self.float_typed(f32_ty, opcode, operands)
1675    }
1676
1677    /// A `NoContraction`-decorated float result of `ty`, so no driver may fuse a multiply and an
1678    /// add: the same TOSA graph yields the same bits on every conformant device.
1679    fn float_typed(&mut self, ty: Id, opcode: u16, operands: &[u32]) -> Id {
1680        let id = self.value(opcode, ty, operands);
1681        instruction(
1682            &mut self.annotations,
1683            OP_DECORATE,
1684            &[id, DECORATION_NO_CONTRACTION],
1685        );
1686        id
1687    }
1688
1689    fn load(&mut self, ty: Id, pointer: Id) -> Id {
1690        self.value(OP_LOAD, ty, &[pointer])
1691    }
1692    fn store(&mut self, pointer: Id, value: Id) {
1693        self.emit(OP_STORE, &[pointer, value]);
1694    }
1695    fn access_chain(&mut self, pointer_ty: Id, base: Id, indices: &[Id]) -> Id {
1696        let mut operands = vec![base];
1697        operands.extend_from_slice(indices);
1698        self.value(OP_ACCESS_CHAIN, pointer_ty, &operands)
1699    }
1700
1701    /// Component `index` of a builtin `uvec3`.
1702    fn builtin_component(&mut self, variable: Id, index: u32) -> Id {
1703        let u32_ty = self.u32_ty();
1704        let pointer = self.pointer(STORAGE_CLASS_INPUT, u32_ty);
1705        let component = self.c_u32(index);
1706        let chain = self.access_chain(pointer, variable, &[component]);
1707        self.load(u32_ty, chain)
1708    }
1709
1710    fn iadd(&mut self, a: Id, b: Id) -> Id {
1711        let ty = self.u32_ty();
1712        self.value(OP_I_ADD, ty, &[a, b])
1713    }
1714    fn isub(&mut self, a: Id, b: Id) -> Id {
1715        let ty = self.u32_ty();
1716        self.value(OP_I_SUB, ty, &[a, b])
1717    }
1718    fn imul(&mut self, a: Id, b: Id) -> Id {
1719        let ty = self.u32_ty();
1720        self.value(OP_I_MUL, ty, &[a, b])
1721    }
1722    fn udiv(&mut self, a: Id, b: Id) -> Id {
1723        let ty = self.u32_ty();
1724        self.value(OP_U_DIV, ty, &[a, b])
1725    }
1726    fn umod(&mut self, a: Id, b: Id) -> Id {
1727        let ty = self.u32_ty();
1728        self.value(OP_U_MOD, ty, &[a, b])
1729    }
1730    fn umin(&mut self, a: Id, b: Id) -> Id {
1731        let ty = self.u32_ty();
1732        let glsl = self.glsl;
1733        self.value(OP_EXT_INST, ty, &[glsl, GLSL_UMIN, a, b])
1734    }
1735    fn shl(&mut self, a: Id, shift: Id) -> Id {
1736        let ty = self.u32_ty();
1737        self.value(OP_SHIFT_LEFT_LOGICAL, ty, &[a, shift])
1738    }
1739    fn shr(&mut self, a: Id, shift: Id) -> Id {
1740        let ty = self.u32_ty();
1741        self.value(OP_SHIFT_RIGHT_LOGICAL, ty, &[a, shift])
1742    }
1743    fn bor(&mut self, a: Id, b: Id) -> Id {
1744        let ty = self.u32_ty();
1745        self.value(OP_BITWISE_OR, ty, &[a, b])
1746    }
1747    fn bxor(&mut self, a: Id, b: Id) -> Id {
1748        let ty = self.u32_ty();
1749        self.value(OP_BITWISE_XOR, ty, &[a, b])
1750    }
1751    fn band(&mut self, a: Id, b: Id) -> Id {
1752        let ty = self.u32_ty();
1753        self.value(OP_BITWISE_AND, ty, &[a, b])
1754    }
1755    fn bnot(&mut self, a: Id) -> Id {
1756        let ty = self.u32_ty();
1757        self.value(OP_NOT, ty, &[a])
1758    }
1759    fn ult(&mut self, a: Id, b: Id) -> Id {
1760        let ty = self.bool_ty();
1761        self.value(OP_U_LESS_THAN, ty, &[a, b])
1762    }
1763    fn uge(&mut self, a: Id, b: Id) -> Id {
1764        let ty = self.bool_ty();
1765        self.value(OP_U_GREATER_THAN_EQUAL, ty, &[a, b])
1766    }
1767    fn ieq(&mut self, a: Id, b: Id) -> Id {
1768        let ty = self.bool_ty();
1769        self.value(OP_I_EQUAL, ty, &[a, b])
1770    }
1771    fn ine(&mut self, a: Id, b: Id) -> Id {
1772        let ty = self.bool_ty();
1773        self.value(OP_I_NOT_EQUAL, ty, &[a, b])
1774    }
1775    fn land(&mut self, a: Id, b: Id) -> Id {
1776        let ty = self.bool_ty();
1777        self.value(OP_LOGICAL_AND, ty, &[a, b])
1778    }
1779    fn lor(&mut self, a: Id, b: Id) -> Id {
1780        let ty = self.bool_ty();
1781        self.value(OP_LOGICAL_OR, ty, &[a, b])
1782    }
1783    fn lxor(&mut self, a: Id, b: Id) -> Id {
1784        let ty = self.bool_ty();
1785        self.value(OP_LOGICAL_NOT_EQUAL, ty, &[a, b])
1786    }
1787    fn lnot(&mut self, a: Id) -> Id {
1788        let ty = self.bool_ty();
1789        self.value(OP_LOGICAL_NOT, ty, &[a])
1790    }
1791    fn select(&mut self, ty: Id, condition: Id, then: Id, otherwise: Id) -> Id {
1792        self.value(OP_SELECT, ty, &[condition, then, otherwise])
1793    }
1794    fn select_u32(&mut self, condition: Id, then: Id, otherwise: Id) -> Id {
1795        let ty = self.u32_ty();
1796        self.select(ty, condition, then, otherwise)
1797    }
1798    fn select_f32(&mut self, condition: Id, then: Id, otherwise: Id) -> Id {
1799        let ty = self.f32_ty();
1800        self.select(ty, condition, then, otherwise)
1801    }
1802    fn bitcast_f32(&mut self, word: Id) -> Id {
1803        let ty = self.f32_ty();
1804        self.value(OP_BITCAST, ty, &[word])
1805    }
1806    fn bitcast_u32(&mut self, float: Id) -> Id {
1807        let ty = self.u32_ty();
1808        self.value(OP_BITCAST, ty, &[float])
1809    }
1810    fn f32_to_f16(&mut self, float: Id) -> Id {
1811        let ty = self.f16_ty();
1812        self.value(OP_F_CONVERT, ty, &[float])
1813    }
1814
1815    fn cooperative_load(&mut self, ty: Id, pointer: Id, stride: Id) -> Id {
1816        let row_major = self.c_u32(0);
1817        self.value(
1818            OP_COOPERATIVE_MATRIX_LOAD_KHR,
1819            ty,
1820            &[pointer, row_major, stride],
1821        )
1822    }
1823
1824    fn cooperative_store(&mut self, pointer: Id, value: Id, stride: Id) {
1825        let row_major = self.c_u32(0);
1826        self.emit(
1827            OP_COOPERATIVE_MATRIX_STORE_KHR,
1828            &[pointer, value, row_major, stride],
1829        );
1830    }
1831
1832    fn cooperative_mul_add(&mut self, ty: Id, a: Id, b: Id, c: Id) -> Id {
1833        self.value(OP_COOPERATIVE_MATRIX_MUL_ADD_KHR, ty, &[a, b, c])
1834    }
1835
1836    fn subgroup_sum_f32(&mut self, value: Id) -> Id {
1837        let ty = self.f32_ty();
1838        let subgroup_scope = self.c_u32(3);
1839        self.value(
1840            OP_GROUP_NON_UNIFORM_F_ADD,
1841            ty,
1842            &[subgroup_scope, 0, value], // GroupOperation Reduce
1843        )
1844    }
1845    fn u_to_f(&mut self, word: Id) -> Id {
1846        let ty = self.f32_ty();
1847        self.value(OP_CONVERT_U_TO_F, ty, &[word])
1848    }
1849    fn f_to_u(&mut self, float: Id) -> Id {
1850        let ty = self.u32_ty();
1851        self.value(OP_CONVERT_F_TO_U, ty, &[float])
1852    }
1853
1854    fn fadd(&mut self, a: Id, b: Id) -> Id {
1855        self.float_value(OP_F_ADD, &[a, b])
1856    }
1857    fn fsub(&mut self, a: Id, b: Id) -> Id {
1858        self.float_value(OP_F_SUB, &[a, b])
1859    }
1860    fn fmul(&mut self, a: Id, b: Id) -> Id {
1861        self.float_value(OP_F_MUL, &[a, b])
1862    }
1863    fn fdiv(&mut self, a: Id, b: Id) -> Id {
1864        self.float_value(OP_F_DIV, &[a, b])
1865    }
1866    fn fneg(&mut self, a: Id) -> Id {
1867        let ty = self.f32_ty();
1868        self.value(OP_F_NEGATE, ty, &[a])
1869    }
1870    /// `a · b + c` as two separately rounded operations, never contracted: the polynomial
1871    /// lanes were tuned against binary64 references with exactly this rounding.
1872    fn fma_free(&mut self, a: Id, b: Id, c: Id) -> Id {
1873        let product = self.fmul(a, b);
1874        self.fadd(product, c)
1875    }
1876    /// `a · b + c` as one fused multiply-add: the GLSL `Fma` instruction decorated
1877    /// `NoContraction`, which Vulkan defines as a single correctly rounded operation rather than
1878    /// a multiply the driver may or may not fuse with the add (ADR 0011).
1879    fn fma(&mut self, a: Id, b: Id, c: Id) -> Id {
1880        let ty = self.f32_ty();
1881        let glsl = self.glsl;
1882        let id = self.value(OP_EXT_INST, ty, &[glsl, GLSL_FMA, a, b, c]);
1883        instruction(
1884            &mut self.annotations,
1885            OP_DECORATE,
1886            &[id, DECORATION_NO_CONTRACTION],
1887        );
1888        id
1889    }
1890
1891    // -- crate-owned binary16 conversions (ADR 0008) --------------------------------------------
1892
1893    /// Widen packed FP8 bits (a `u32` in `0..=0xff`) to binary32, exactly. The kernel twin of
1894    /// [`virtio_accel_tosa::fp8e4m3_to_f32`] / [`virtio_accel_tosa::fp8e5m2_to_f32`]; every FP8
1895    /// value is representable in binary32, so all 256 patterns of each format widen exactly on
1896    /// every device, and no `OpFConvert` appears that a driver could demote.
1897    ///
1898    /// The exponent-offset construction (shared with [`Self::widen_f16`]): place the magnitude
1899    /// bits at the top of a binary32 significand under a fixed exponent `K` chosen so that a
1900    /// *normal* encoding reads off directly as `1.m · 2^(e - bias)`. A subnormal encoding then
1901    /// reads as `1.m · 2^(1 - bias)` and the value wanted is `0.m · 2^(1 - bias)`, which is
1902    /// `(x - 2^(1 - bias)) · 2` — two exact binary32 operations, since `x` and the constant
1903    /// share a binade. Nothing here produces a binary32 denormal (the smallest result is the
1904    /// format's smallest subnormal, `2^-9` or `2^-16`), so a device that flushes denormals
1905    /// computes the same bits. Specials are one compare and one select. About half the
1906    /// instructions of the integer-only expansion it replaced, which matters because the
1907    /// MATMUL kernels widen every staged element.
1908    fn widen_fp8(&mut self, format: Fp8Format, bits: Id) -> Id {
1909        // Mantissa width, the fixed exponent `K = 127 - bias`, and the top exponent field.
1910        let (mantissa_bits, k, top_exponent) = match format {
1911            Fp8Format::E4M3 => (3_u32, 120_u32, 15_u32),
1912            Fp8Format::E5M2 => (2, 112, 31),
1913        };
1914        let magnitude_mask = self.c_u32(0x7f);
1915        let shift = self.c_u32(23 - mantissa_bits);
1916        let k_bits = self.c_u32(k << 23);
1917        let mantissa_shift = self.c_u32(mantissa_bits);
1918        let zero = self.c_u32(0);
1919        let implicit_one = self.c_f32(f32::from_bits(k << 23));
1920        let sign_mask = self.c_u32(0x80);
1921        let twenty_four = self.c_u32(24);
1922        let magnitude = self.band(bits, magnitude_mask);
1923        let placed = self.shl(magnitude, shift);
1924        // An add, not an or: the encoding's exponent field and `K` overlap in the binary32
1925        // exponent bits, and the construction is `e + K`.
1926        let x_bits = self.iadd(placed, k_bits);
1927        let x = self.bitcast_f32(x_bits);
1928        let exponent = self.shr(magnitude, mantissa_shift);
1929        let is_subnormal = self.ieq(exponent, zero);
1930        let difference = self.fsub(x, implicit_one);
1931        let subnormal = self.fadd(difference, difference);
1932        let value = self.select_f32(is_subnormal, subnormal, x);
1933        let value = self.bitcast_u32(value);
1934        // Specials differ: E4M3's top exponent stays finite except for the all-ones fraction,
1935        // which is its only NaN; E5M2's top exponent is infinity or NaN as usual, payload kept.
1936        let value = match format {
1937            Fp8Format::E4M3 => {
1938                let nan_pattern = self.c_u32(0x7f);
1939                let quiet_nan = self.c_u32(0x7fc0_0000);
1940                let is_nan = self.ieq(magnitude, nan_pattern);
1941                self.select_u32(is_nan, quiet_nan, value)
1942            }
1943            Fp8Format::E5M2 => {
1944                let top = self.c_u32(top_exponent);
1945                let inf_bits = self.c_u32(0x7f80_0000);
1946                let fraction_mask = self.c_u32((1 << mantissa_bits) - 1);
1947                let is_infnan = self.ieq(exponent, top);
1948                let fraction = self.band(magnitude, fraction_mask);
1949                let payload = self.shl(fraction, shift);
1950                let infnan = self.bor(inf_bits, payload);
1951                self.select_u32(is_infnan, infnan, value)
1952            }
1953        };
1954        let sign = self.band(bits, sign_mask);
1955        let sign = self.shl(sign, twenty_four);
1956        let value = self.bor(value, sign);
1957        self.bitcast_f32(value)
1958    }
1959
1960    /// Widen packed binary16 bits (a `u32` in `0..=0xffff`) to binary32, exactly. This is the
1961    /// kernel twin of the host [`f16_to_f32`], by the exponent-offset construction described at
1962    /// [`Self::widen_fp8`] (`K = 112`, subnormals as `(x - 2^-15) · 2`): every value —
1963    /// subnormals included — is exact on every device, no binary32 denormal is ever formed, and
1964    /// there is no `OpFConvert` pattern a driver can demote back to f16 (ADR 0008).
1965    fn widen_f16(&mut self, bits: Id) -> Id {
1966        let magnitude_mask = self.c_u32(0x7fff);
1967        let thirteen = self.c_u32(13);
1968        let ten = self.c_u32(10);
1969        let sixteen = self.c_u32(16);
1970        let k_bits = self.c_u32(112 << 23);
1971        let implicit_one = self.c_f32(f32::from_bits(112 << 23));
1972        let zero = self.c_u32(0);
1973        let thirty_one = self.c_u32(31);
1974        let inf_bits = self.c_u32(0x7f80_0000);
1975        let mantissa_mask = self.c_u32(0x3ff);
1976        let sign_mask = self.c_u32(0x8000);
1977        let magnitude = self.band(bits, magnitude_mask);
1978        let placed = self.shl(magnitude, thirteen);
1979        // An add, not an or: the encoding's exponent field and `K` overlap in the binary32
1980        // exponent bits, and the construction is `e + K`.
1981        let x_bits = self.iadd(placed, k_bits);
1982        let x = self.bitcast_f32(x_bits);
1983        let exponent = self.shr(magnitude, ten);
1984        let is_subnormal = self.ieq(exponent, zero);
1985        let difference = self.fsub(x, implicit_one);
1986        let subnormal = self.fadd(difference, difference);
1987        let value = self.select_f32(is_subnormal, subnormal, x);
1988        let value = self.bitcast_u32(value);
1989        // Infinity or NaN: exponent all ones, payload preserved.
1990        let is_infnan = self.ieq(exponent, thirty_one);
1991        let mantissa = self.band(magnitude, mantissa_mask);
1992        let payload = self.shl(mantissa, thirteen);
1993        let infnan = self.bor(inf_bits, payload);
1994        let value = self.select_u32(is_infnan, infnan, value);
1995        let sign = self.band(bits, sign_mask);
1996        let sign = self.shl(sign, sixteen);
1997        let value = self.bor(value, sign);
1998        self.bitcast_f32(value)
1999    }
2000
2001    /// Narrow a binary32 value to packed binary16 bits (a `u32` in `0..=0xffff`),
2002    /// round-to-nearest-even, NaN canonicalized to the quiet `0x7e00` payload with its sign.
2003    /// This is the kernel twin of the host [`f32_to_f16_bits`]: integer and binary32 operations
2004    /// only, so subnormals are produced — never flushed — on every device (ADR 0008).
2005    /// Narrow binary32 to FP8, the kernel twin of [`f32_to_fp8_bits`] and bound by its tests.
2006    /// Integer rounding throughout, with one `RoundEven` for the subnormal range, so the result
2007    /// is identical on every device. The overflow policy is that function's: NaN for E4M3,
2008    /// infinity for E5M2 (ADR 0009).
2009    fn narrow_fp8(&mut self, format: Fp8Format, value: Id) -> Id {
2010        let (mantissa_bits, bias, nan_out, max_finite, overflow_out) = match format {
2011            Fp8Format::E4M3 => (3_u32, 7_u32, 0x7f_u32, 0x7e_u32, 0x7f_u32),
2012            Fp8Format::E5M2 => (2, 15, 0x7e, 0x7b, 0x7c),
2013        };
2014        let shift = 23 - mantissa_bits;
2015        let twenty_four = self.c_u32(24);
2016        let one = self.c_u32(1);
2017        let sign_mask = self.c_u32(0x80);
2018        let magnitude_mask = self.c_u32(0x7fff_ffff);
2019        let inf_bits = self.c_u32(0x7f80_0000);
2020        let normal_floor = self.c_u32((128 - bias) << 23);
2021        let rebias = self.c_u32((127 - bias) << 23);
2022        let shift_c = self.c_u32(shift);
2023        let half_ulp = self.c_u32((1 << (shift - 1)) - 1);
2024        let scale = self.c_f32(f32::from_bits((127 + bias - 1 + mantissa_bits) << 23));
2025        let max_finite_c = self.c_u32(max_finite);
2026        let nan_c = self.c_u32(nan_out);
2027        let overflow_c = self.c_u32(overflow_out);
2028        let bits = self.bitcast_u32(value);
2029        let sign = self.shr(bits, twenty_four);
2030        let sign = self.band(sign, sign_mask);
2031        let magnitude = self.band(bits, magnitude_mask);
2032        let is_nan = self.ult(inf_bits, magnitude);
2033        let is_normal = self.uge(magnitude, normal_floor);
2034        // Normal: rebias, then round the significand with the add-half-ulp-plus-guard trick; a
2035        // carry increments the exponent on its own. The subnormal branch discards the wrap.
2036        let adjusted = self.isub(magnitude, rebias);
2037        let guard = self.shr(adjusted, shift_c);
2038        let guard = self.band(guard, one);
2039        let rounding = self.iadd(half_ulp, guard);
2040        let rounded = self.iadd(adjusted, rounding);
2041        let normal_out = self.shr(rounded, shift_c);
2042        // Subnormal or zero: scale into integer range, exact for every representable
2043        // subnormal, and round to the nearest integer. A carry reaches the smallest normal.
2044        let safe_magnitude = self.select_u32(is_normal, normal_floor, magnitude);
2045        let scaled = self.bitcast_f32(safe_magnitude);
2046        let scaled = self.fmul(scaled, scale);
2047        let scaled = self.ext_f32(GLSL_ROUND_EVEN, &[scaled]);
2048        let subnormal_out = self.f_to_u(scaled);
2049        let body = self.select_u32(is_normal, normal_out, subnormal_out);
2050        // Rounding has happened, so landing past the last finite encoding is the overflow test.
2051        let overflowed = self.ult(max_finite_c, body);
2052        let body = self.select_u32(overflowed, overflow_c, body);
2053        let body = self.select_u32(is_nan, nan_c, body);
2054        self.bor(sign, body)
2055    }
2056
2057    fn narrow_f16(&mut self, value: Id) -> Id {
2058        let sixteen = self.c_u32(16);
2059        let thirteen = self.c_u32(13);
2060        let one = self.c_u32(1);
2061        let sign_mask = self.c_u32(0x8000);
2062        let magnitude_mask = self.c_u32(0x7fff_ffff);
2063        let inf_bits = self.c_u32(0x7f80_0000);
2064        let overflow_bits = self.c_u32(0x477f_f000);
2065        let normal_floor = self.c_u32(0x3880_0000);
2066        let rebias = self.c_u32(0x3800_0000);
2067        let half_ulp = self.c_u32(0x0000_0fff);
2068        let inf_out = self.c_u32(0x7c00);
2069        let nan_out = self.c_u32(0x7e00);
2070        let scale = self.c_f32(16_777_216.0);
2071        let bits = self.bitcast_u32(value);
2072        let sign = self.shr(bits, sixteen);
2073        let sign = self.band(sign, sign_mask);
2074        let magnitude = self.band(bits, magnitude_mask);
2075        let is_nan = self.ult(inf_bits, magnitude);
2076        // At or above 65520 — the midpoint between 65504 and the overflow binade — round to
2077        // infinity (65504 itself is `0x477f_e000`, finite).
2078        let overflow = self.uge(magnitude, overflow_bits);
2079        // Normal binary16: rebias the exponent, then round the significand with the
2080        // add-half-ulp-plus-guard trick; a mantissa carry increments the exponent on its own.
2081        let is_normal = self.uge(magnitude, normal_floor);
2082        let adjusted = self.isub(magnitude, rebias);
2083        let guard = self.shr(adjusted, thirteen);
2084        let guard = self.band(guard, one);
2085        let rounding = self.iadd(half_ulp, guard);
2086        let rounded = self.iadd(adjusted, rounding);
2087        let normal_out = self.shr(rounded, thirteen);
2088        // Subnormal binary16 or zero: scale into integer range (exact while the value is at or
2089        // above 2^-25; anything smaller rounds to zero either way) and round to the nearest
2090        // integer. The result is at most 1024 — the smallest normal encoding, correctly reached
2091        // by the rounding carry.
2092        let safe_magnitude = self.select_u32(is_normal, normal_floor, magnitude);
2093        let scaled = self.bitcast_f32(safe_magnitude);
2094        let scaled = self.fmul(scaled, scale);
2095        let scaled = self.ext_f32(GLSL_ROUND_EVEN, &[scaled]);
2096        let subnormal_out = self.f_to_u(scaled);
2097        let body = self.select_u32(is_normal, normal_out, subnormal_out);
2098        let body = self.select_u32(overflow, inf_out, body);
2099        let body = self.select_u32(is_nan, nan_out, body);
2100        self.bor(sign, body)
2101    }
2102    fn ext_f32(&mut self, op: u32, args: &[Id]) -> Id {
2103        let ty = self.f32_ty();
2104        let glsl = self.glsl;
2105        let mut operands = vec![glsl, op];
2106        operands.extend_from_slice(args);
2107        self.value(OP_EXT_INST, ty, &operands)
2108    }
2109    fn fabs(&mut self, a: Id) -> Id {
2110        self.ext_f32(GLSL_FABS, &[a])
2111    }
2112    fn is_nan(&mut self, a: Id) -> Id {
2113        let ty = self.bool_ty();
2114        self.value(OP_IS_NAN, ty, &[a])
2115    }
2116    fn foeq(&mut self, a: Id, b: Id) -> Id {
2117        let ty = self.bool_ty();
2118        self.value(OP_F_ORD_EQUAL, ty, &[a, b])
2119    }
2120    fn folt(&mut self, a: Id, b: Id) -> Id {
2121        let ty = self.bool_ty();
2122        self.value(OP_F_ORD_LESS_THAN, ty, &[a, b])
2123    }
2124    fn fogt(&mut self, a: Id, b: Id) -> Id {
2125        let ty = self.bool_ty();
2126        self.value(OP_F_ORD_GREATER_THAN, ty, &[a, b])
2127    }
2128    fn foge(&mut self, a: Id, b: Id) -> Id {
2129        let ty = self.bool_ty();
2130        self.value(OP_F_ORD_GREATER_THAN_EQUAL, ty, &[a, b])
2131    }
2132
2133    /// Element `index` of a `Private` `u32` array.
2134    fn private_element(&mut self, array: Id, index: Id) -> Id {
2135        let u32_ty = self.u32_ty();
2136        let pointer = self.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
2137        let chain = self.access_chain(pointer, array, &[index]);
2138        self.load(u32_ty, chain)
2139    }
2140
2141    /// Index of the most significant set bit of `value` (`-1` for zero, unused here).
2142    fn find_msb(&mut self, value: Id) -> Id {
2143        let ty = self.u32_ty();
2144        let glsl = self.glsl;
2145        self.value(OP_EXT_INST, ty, &[glsl, GLSL_FIND_U_MSB, value])
2146    }
2147
2148    /// `a * b` as a `(high, low)` pair of 32-bit words, built from four 16-bit products so no
2149    /// 64-bit integer capability is required.
2150    fn mul_wide(&mut self, a: Id, b: Id) -> (Id, Id) {
2151        let sixteen = self.c_u32(16);
2152        let mask = self.c_u32(0xffff);
2153        let a_lo = self.band(a, mask);
2154        let a_hi = self.shr(a, sixteen);
2155        let b_lo = self.band(b, mask);
2156        let b_hi = self.shr(b, sixteen);
2157        let ll = self.imul(a_lo, b_lo);
2158        let lh = self.imul(a_lo, b_hi);
2159        let hl = self.imul(a_hi, b_lo);
2160        let hh = self.imul(a_hi, b_hi);
2161        let ll_hi = self.shr(ll, sixteen);
2162        let lh_lo = self.band(lh, mask);
2163        let hl_lo = self.band(hl, mask);
2164        let middle = self.iadd(ll_hi, lh_lo);
2165        let middle = self.iadd(middle, hl_lo);
2166        let ll_lo = self.band(ll, mask);
2167        let middle_lo = self.band(middle, mask);
2168        let middle_shifted = self.shl(middle_lo, sixteen);
2169        let low = self.bor(ll_lo, middle_shifted);
2170        let lh_hi = self.shr(lh, sixteen);
2171        let hl_hi = self.shr(hl, sixteen);
2172        let middle_hi = self.shr(middle, sixteen);
2173        let high = self.iadd(hh, lh_hi);
2174        let high = self.iadd(high, hl_hi);
2175        let high = self.iadd(high, middle_hi);
2176        (high, low)
2177    }
2178
2179    /// `a + b` as `(sum, carry)`.
2180    fn add_carry(&mut self, a: Id, b: Id) -> (Id, Id) {
2181        let sum = self.iadd(a, b);
2182        let carried = self.ult(sum, a);
2183        let one = self.c_u32(1);
2184        let zero = self.c_u32(0);
2185        let carry = self.select_u32(carried, one, zero);
2186        (sum, carry)
2187    }
2188
2189    /// `2^exponent` as an FP32 bit pattern, for `exponent` inside the normal range.
2190    fn pow2(&mut self, exponent: Id) -> Id {
2191        let bias = self.c_u32(127);
2192        let biased = self.iadd(exponent, bias);
2193        let twenty_three = self.c_u32(23);
2194        let bits = self.shl(biased, twenty_three);
2195        self.bitcast_f32(bits)
2196    }
2197
2198    fn label(&mut self, id: Id) {
2199        self.emit(OP_LABEL, &[id]);
2200    }
2201    fn branch(&mut self, target: Id) {
2202        self.emit(OP_BRANCH, &[target]);
2203    }
2204    fn branch_conditional(&mut self, condition: Id, then: Id, otherwise: Id) {
2205        self.emit(OP_BRANCH_CONDITIONAL, &[condition, then, otherwise]);
2206    }
2207    fn workgroup_barrier(&mut self) {
2208        let scope = self.c_u32(SCOPE_WORKGROUP);
2209        let semantics = self.c_u32(MEMORY_SEMANTICS_ACQUIRE_RELEASE_WORKGROUP);
2210        self.emit(OP_CONTROL_BARRIER, &[scope, scope, semantics]);
2211    }
2212
2213    // -- structured loops ---------------------------------------------------------------------
2214
2215    /// Open `for (; *counter < limit; )`: emits the header and the body label. The returned
2216    /// [`LoopScope`] must be closed with [`Self::end_loop`], which adds `step` to the counter.
2217    fn begin_loop(&mut self, counter: Id, limit: Id) -> (LoopScope, Id) {
2218        let scope = LoopScope {
2219            header: self.id(),
2220            body: self.id(),
2221            cont: self.id(),
2222            merge: self.id(),
2223        };
2224        self.branch(scope.header);
2225        self.label(scope.header);
2226        let u32_ty = self.u32_ty();
2227        let current = self.load(u32_ty, counter);
2228        let in_range = self.ult(current, limit);
2229        self.emit(OP_LOOP_MERGE, &[scope.merge, scope.cont, LOOP_CONTROL_NONE]);
2230        self.branch_conditional(in_range, scope.body, scope.merge);
2231        self.label(scope.body);
2232        (scope, current)
2233    }
2234
2235    fn end_loop(&mut self, scope: LoopScope, counter: Id, step: Id) {
2236        self.branch(scope.cont);
2237        self.label(scope.cont);
2238        let u32_ty = self.u32_ty();
2239        let current = self.load(u32_ty, counter);
2240        let next = self.iadd(current, step);
2241        self.store(counter, next);
2242        self.branch(scope.header);
2243        self.label(scope.merge);
2244    }
2245
2246    /// `if condition { then(); }` with no value flowing out.
2247    fn if_then(&mut self, condition: Id, then: impl FnOnce(&mut Self)) {
2248        let then_label = self.id();
2249        let merge = self.id();
2250        self.emit(OP_SELECTION_MERGE, &[merge, SELECTION_CONTROL_NONE]);
2251        self.branch_conditional(condition, then_label, merge);
2252        self.label(then_label);
2253        then(self);
2254        self.branch(merge);
2255        self.label(merge);
2256    }
2257
2258    // -- tensor access -------------------------------------------------------------------------
2259
2260    /// Pointer to word `word_index` of the operand `(buffer, base)` inside `buffers`.
2261    fn word_pointer(&mut self, buffers: Id, operand: (Id, Id), word_index: Id) -> Id {
2262        let u32_ty = self.u32_ty();
2263        let pointer = self.pointer(STORAGE_CLASS_STORAGE_BUFFER, u32_ty);
2264        let zero = self.c_u32(0);
2265        let index = self.iadd(operand.1, word_index);
2266        self.access_chain(pointer, buffers, &[operand.0, zero, index])
2267    }
2268
2269    fn load_word(&mut self, buffers: Id, operand: (Id, Id), word_index: Id) -> Id {
2270        let pointer = self.word_pointer(buffers, operand, word_index);
2271        let u32_ty = self.u32_ty();
2272        self.load(u32_ty, pointer)
2273    }
2274
2275    fn load_f32(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2276        let word = self.load_word(buffers, operand, element);
2277        self.bitcast_f32(word)
2278    }
2279
2280    fn store_word(&mut self, buffers: Id, operand: (Id, Id), word_index: Id, value: Id) {
2281        let pointer = self.word_pointer(buffers, operand, word_index);
2282        self.store(pointer, value);
2283    }
2284
2285    fn store_f32(&mut self, buffers: Id, operand: (Id, Id), element: Id, value: Id) {
2286        let word = self.bitcast_u32(value);
2287        self.store_word(buffers, operand, element, word);
2288    }
2289
2290    /// The raw 8 bits of byte-storage `element` as a `u32` (`0..=0xff`): word load, shift the
2291    /// lane down, mask. No interpretation — an FP8 bit pattern passes through untouched.
2292    fn load_byte_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2293        let two = self.c_u32(2);
2294        let three = self.c_u32(3);
2295        let mask = self.c_u32(0xff);
2296        let word_index = self.shr(element, two);
2297        let word = self.load_word(buffers, operand, word_index);
2298        let lane = self.band(element, three);
2299        let eight = self.c_u32(8);
2300        let shift = self.imul(lane, eight);
2301        let shifted = self.shr(word, shift);
2302        self.band(shifted, mask)
2303    }
2304
2305    /// Byte `element` of a byte-storage operand as a boolean (any nonzero byte is true).
2306    fn load_bool(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2307        let zero = self.c_u32(0);
2308        let byte = self.load_byte_bits(buffers, operand, element);
2309        self.ine(byte, zero)
2310    }
2311
2312    /// Write the low 8 bits of `byte` into byte `element` without touching the other bytes of
2313    /// its word: clear with `OpAtomicAnd`, then set with `OpAtomicOr`. The neighbour-safe
2314    /// sequence `BOOL` and FP8 lanes share.
2315    fn store_byte_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id, byte: Id) {
2316        let two = self.c_u32(2);
2317        let three = self.c_u32(3);
2318        let mask = self.c_u32(0xff);
2319        let word_index = self.shr(element, two);
2320        let pointer = self.word_pointer(buffers, operand, word_index);
2321        let lane = self.band(element, three);
2322        let eight = self.c_u32(8);
2323        let shift = self.imul(lane, eight);
2324        let clear = self.shl(mask, shift);
2325        let clear = self.bnot(clear);
2326        let masked = self.band(byte, mask);
2327        let set = self.shl(masked, shift);
2328        let scope = self.c_u32(SCOPE_DEVICE);
2329        let semantics = self.c_u32(MEMORY_SEMANTICS_RELAXED);
2330        let u32_ty = self.u32_ty();
2331        self.value(OP_ATOMIC_AND, u32_ty, &[pointer, scope, semantics, clear]);
2332        self.value(OP_ATOMIC_OR, u32_ty, &[pointer, scope, semantics, set]);
2333    }
2334
2335    /// Write byte `element` of a byte-storage operand as canonical `0`/`1`.
2336    fn store_bool(&mut self, buffers: Id, operand: (Id, Id), element: Id, value: Id) {
2337        let one = self.c_u32(1);
2338        let zero = self.c_u32(0);
2339        let byte = self.select_u32(value, one, zero);
2340        self.store_byte_bits(buffers, operand, element, byte);
2341    }
2342
2343    /// The raw 16 bits of half-storage `element` as a `u32` (`0..=0xffff`): word load, shift
2344    /// the lane down, mask. No float conversion — bit patterns (NaN payloads, subnormals) pass
2345    /// through untouched.
2346    fn load_half_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2347        let one = self.c_u32(1);
2348        let sixteen = self.c_u32(16);
2349        let mask = self.c_u32(0xffff);
2350        let word_index = self.shr(element, one);
2351        let word = self.load_word(buffers, operand, word_index);
2352        let lane = self.band(element, one);
2353        let shift = self.imul(lane, sixteen);
2354        let shifted = self.shr(word, shift);
2355        self.band(shifted, mask)
2356    }
2357
2358    /// Write the low 16 bits of `bits` to half-storage `element` without touching the other
2359    /// element of its word: clear with `OpAtomicAnd`, then set with `OpAtomicOr` (the same
2360    /// neighbour-safe pattern as [`Self::store_bool`]).
2361    fn store_half_bits(&mut self, buffers: Id, operand: (Id, Id), element: Id, bits: Id) {
2362        let one = self.c_u32(1);
2363        let sixteen = self.c_u32(16);
2364        let mask = self.c_u32(0xffff);
2365        let word_index = self.shr(element, one);
2366        let pointer = self.word_pointer(buffers, operand, word_index);
2367        let lane = self.band(element, one);
2368        let shift = self.imul(lane, sixteen);
2369        let clear = self.shl(mask, shift);
2370        let clear = self.bnot(clear);
2371        let bits = self.band(bits, mask);
2372        let set = self.shl(bits, shift);
2373        let scope = self.c_u32(SCOPE_DEVICE);
2374        let semantics = self.c_u32(MEMORY_SEMANTICS_RELAXED);
2375        let u32_ty = self.u32_ty();
2376        self.value(OP_ATOMIC_AND, u32_ty, &[pointer, scope, semantics, clear]);
2377        self.value(OP_ATOMIC_OR, u32_ty, &[pointer, scope, semantics, set]);
2378    }
2379
2380    /// Store one binary32 float at `storage`, narrowing where the storage is narrower.
2381    fn store_float(
2382        &mut self,
2383        storage: Storage,
2384        buffers: Id,
2385        operand: (Id, Id),
2386        element: Id,
2387        value: Id,
2388    ) {
2389        match storage {
2390            Storage::Word => self.store_f32(buffers, operand, element, value),
2391            Storage::Half => {
2392                let bits = self.narrow_f16(value);
2393                self.store_half_bits(buffers, operand, element, bits);
2394            }
2395            Storage::Quarter(format) => {
2396                let bits = self.narrow_fp8(format, value);
2397                self.store_byte_bits(buffers, operand, element, bits);
2398            }
2399            Storage::Byte => unreachable!("BOOL is not a float storage"),
2400        }
2401    }
2402
2403    /// Load one float element as binary32: word storage bitcasts, half storage is unpacked and
2404    /// widened exactly by [`Self::widen_f16`].
2405    fn load_float(&mut self, storage: Storage, buffers: Id, operand: (Id, Id), element: Id) -> Id {
2406        match storage {
2407            Storage::Word => self.load_f32(buffers, operand, element),
2408            Storage::Quarter(format) => {
2409                let bits = self.load_byte_bits(buffers, operand, element);
2410                self.widen_fp8(format, bits)
2411            }
2412            Storage::Half => {
2413                let bits = self.load_half_bits(buffers, operand, element);
2414                self.widen_f16(bits)
2415            }
2416            Storage::Byte => unreachable!("byte storage is not a float lane"),
2417        }
2418    }
2419
2420    /// The raw bits of the `lanes` consecutive elements `base..base + lanes` of a `storage`
2421    /// operand, one `u32` per element (`0..=0xff` at byte storage, `0..=0xffff` at half storage,
2422    /// the whole word otherwise), loading every storage word the group touches exactly once.
2423    ///
2424    /// `base` must be a multiple of `lanes`. When the source packs no more elements per word
2425    /// than the group holds, the group starts on a word boundary and each lane's word and shift
2426    /// are compile-time constants; when it packs more (FP8 bytes feeding an FP16 pair), the group
2427    /// is a sub-span of one word at a run-time byte offset, and still one load.
2428    fn load_lane_group(
2429        &mut self,
2430        storage: Storage,
2431        buffers: Id,
2432        operand: (Id, Id),
2433        base: Id,
2434        lanes: u32,
2435    ) -> Vec<Id> {
2436        let source_lanes = storage.lanes();
2437        let lane_bits = 32 / source_lanes;
2438        let first = match source_lanes.trailing_zeros() {
2439            0 => base,
2440            shift => {
2441                let shift = self.c_u32(shift);
2442                self.shr(base, shift)
2443            }
2444        };
2445        let words: Vec<Id> = (0..lanes.div_ceil(source_lanes))
2446            .map(|word| {
2447                let word = self.c_u32(word);
2448                let index = self.iadd(first, word);
2449                self.load_word(buffers, operand, index)
2450            })
2451            .collect();
2452        if lane_bits == 32 {
2453            return words;
2454        }
2455        let mask = self.c_u32((1 << lane_bits) - 1);
2456        let offset = (source_lanes > lanes).then(|| {
2457            let modulus = self.c_u32(source_lanes - 1);
2458            self.band(base, modulus)
2459        });
2460        (0..lanes)
2461            .map(|lane| {
2462                let word = words[(lane / source_lanes) as usize];
2463                let shift = match offset {
2464                    None => self.c_u32((lane % source_lanes) * lane_bits),
2465                    Some(offset) => {
2466                        let lane = self.c_u32(lane % source_lanes);
2467                        let position = self.iadd(offset, lane);
2468                        let bits = self.c_u32(lane_bits);
2469                        self.imul(position, bits)
2470                    }
2471                };
2472                let shifted = self.shr(word, shift);
2473                self.band(shifted, mask)
2474            })
2475            .collect()
2476    }
2477
2478    /// Pack lane values — each already within `32 / values.len()` bits — into one storage word,
2479    /// lane 0 lowest.
2480    fn pack_lanes(&mut self, values: &[Id]) -> Id {
2481        let lane_bits = 32 / values.len() as u32;
2482        let mut word = values[0];
2483        for (lane, value) in values.iter().enumerate().skip(1) {
2484            let shift = self.c_u32(lane as u32 * lane_bits);
2485            let shifted = self.shl(*value, shift);
2486            word = self.bor(word, shifted);
2487        }
2488        word
2489    }
2490
2491    /// Store raw bits into one element of `storage` without touching its word neighbours.
2492    fn store_lane_bits(
2493        &mut self,
2494        storage: Storage,
2495        buffers: Id,
2496        operand: (Id, Id),
2497        element: Id,
2498        bits: Id,
2499    ) {
2500        match storage {
2501            Storage::Word => self.store_word(buffers, operand, element, bits),
2502            Storage::Half => self.store_half_bits(buffers, operand, element, bits),
2503            Storage::Byte | Storage::Quarter(_) => {
2504                self.store_byte_bits(buffers, operand, element, bits)
2505            }
2506        }
2507    }
2508
2509    /// Widen one element's raw storage bits to binary32.
2510    fn widen_bits(&mut self, storage: Storage, bits: Id) -> Id {
2511        match storage {
2512            Storage::Word => self.bitcast_f32(bits),
2513            Storage::Half => self.widen_f16(bits),
2514            Storage::Quarter(format) => self.widen_fp8(format, bits),
2515            Storage::Byte => unreachable!("byte storage is not a float lane"),
2516        }
2517    }
2518
2519    /// Narrow binary32 to the raw storage bits of `storage`, crate-owned rounding throughout.
2520    fn narrow_bits(&mut self, storage: Storage, value: Id) -> Id {
2521        match storage {
2522            Storage::Word => self.bitcast_u32(value),
2523            Storage::Half => self.narrow_f16(value),
2524            Storage::Quarter(format) => self.narrow_fp8(format, value),
2525            Storage::Byte => unreachable!("BOOL is not a float storage"),
2526        }
2527    }
2528
2529    /// Decompose linear index `i` over `dims` (last dimension fastest) and accumulate per-operand
2530    /// element indices from `strides`.
2531    fn strided_indices(
2532        &mut self,
2533        i: Id,
2534        dims: &[Id; MAX_RANK],
2535        strides: &[[Id; MAX_RANK]],
2536    ) -> Vec<Id> {
2537        let zero = self.c_u32(0);
2538        let mut indices = vec![zero; strides.len()];
2539        let mut remainder = i;
2540        for d in (0..MAX_RANK).rev() {
2541            let coordinate = self.umod(remainder, dims[d]);
2542            remainder = self.udiv(remainder, dims[d]);
2543            for (operand, stride) in strides.iter().enumerate() {
2544                let term = self.imul(coordinate, stride[d]);
2545                indices[operand] = self.iadd(indices[operand], term);
2546            }
2547        }
2548        indices
2549    }
2550
2551    /// The grid-stride prologue: `counter = gid.x`, returning `(counter, stride)` where stride is
2552    /// `NumWorkgroups.x * workgroup`.
2553    fn grid_stride(&mut self, workgroup: u32) -> (Id, Id) {
2554        let gid = self.builtin_uvec3(BUILT_IN_GLOBAL_INVOCATION_ID);
2555        let num_workgroups = self.builtin_uvec3(BUILT_IN_NUM_WORKGROUPS);
2556        self.begin_main();
2557        let u32_ty = self.u32_ty();
2558        let counter = self.local(u32_ty);
2559        let start = self.builtin_component(gid, 0);
2560        self.store(counter, start);
2561        let groups = self.builtin_component(num_workgroups, 0);
2562        let size = self.c_u32(workgroup);
2563        let stride = self.imul(groups, size);
2564        (counter, stride)
2565    }
2566
2567    // -- scalar lanes ---------------------------------------------------------------------------
2568
2569    /// TOSA `apply_max_s(a, b)`: `a >= b ? a : b` with the NaN mode applied first.
2570    fn apply_max(&mut self, a: Id, b: Id, nan_mode: NanMode) -> Id {
2571        let ordered = self.foge(a, b);
2572        let picked = self.select_f32(ordered, a, b);
2573        self.apply_nan_mode(a, b, picked, nan_mode)
2574    }
2575
2576    /// TOSA `apply_min_s(a, b)`: `a < b ? a : b` with the NaN mode applied first.
2577    fn apply_min(&mut self, a: Id, b: Id, nan_mode: NanMode) -> Id {
2578        let ordered = self.folt(a, b);
2579        let picked = self.select_f32(ordered, a, b);
2580        self.apply_nan_mode(a, b, picked, nan_mode)
2581    }
2582
2583    fn apply_nan_mode(&mut self, a: Id, b: Id, picked: Id, nan_mode: NanMode) -> Id {
2584        let a_nan = self.is_nan(a);
2585        let b_nan = self.is_nan(b);
2586        match nan_mode {
2587            NanMode::Propagate => {
2588                let nan = self.c_f32(f32::NAN);
2589                let any_nan = self.lor(a_nan, b_nan);
2590                self.select_f32(any_nan, nan, picked)
2591            }
2592            NanMode::Ignore => {
2593                let without_b = self.select_f32(b_nan, a, picked);
2594                self.select_f32(a_nan, b, without_b)
2595            }
2596        }
2597    }
2598
2599    /// Horner evaluation of `Σ coefficients[i] * x^i`, highest degree first in `coefficients`.
2600    fn horner(&mut self, x: Id, coefficients: &[f32]) -> Id {
2601        let mut acc = self.c_f32(coefficients[0]);
2602        for coefficient in &coefficients[1..] {
2603            let c = self.c_f32(*coefficient);
2604            acc = self.fma_free(acc, x, c);
2605        }
2606        acc
2607    }
2608
2609    /// Cephes octant reduction: the three-part π/4 subtraction, exact while the octant count is
2610    /// exactly representable. Returns `(octant & 7, z)` with the octant already bumped to even,
2611    /// so `z = |x| - octant·π/4` lies in `[-π/4, π/4]`.
2612    fn cephes_reduce(&mut self, magnitude: Id) -> (Id, Id) {
2613        let four_over_pi = self.c_f32(1.273_239_5);
2614        let scaled = self.fmul(magnitude, four_over_pi);
2615        let octant = self.f_to_u(scaled);
2616        let one = self.c_u32(1);
2617        let zero = self.c_u32(0);
2618        let odd = self.band(octant, one);
2619        let is_odd = self.ine(odd, zero);
2620        let bumped = self.iadd(octant, one);
2621        let octant = self.select_u32(is_odd, bumped, octant);
2622        let y = self.u_to_f(octant);
2623        // Extended-precision modular arithmetic: |x| - y·(DP1 + DP2 + DP3).
2624        let dp1 = self.c_f32(0.785_156_25);
2625        let dp2 = self.c_f32(2.418_756_5e-4);
2626        let dp3 = self.c_f32(3.774_895e-8);
2627        let t1 = self.fmul(y, dp1);
2628        let r = self.fsub(magnitude, t1);
2629        let t2 = self.fmul(y, dp2);
2630        let r = self.fsub(r, t2);
2631        let t3 = self.fmul(y, dp3);
2632        let z = self.fsub(r, t3);
2633        let seven = self.c_u32(7);
2634        let octant = self.band(octant, seven);
2635        (octant, z)
2636    }
2637
2638    /// Payne–Hanek octant reduction: exact for every finite magnitude, at the cost of a
2639    /// 128-bit window of 2/π and one 24×128-bit integer multiply.
2640    ///
2641    /// `|x| = m·2^(e-149)` with `m` the 24-bit significand and `e` the biased exponent, so
2642    /// `|x|·4/π = m·2^(e-148)·(2/π)`. Every 2/π bit whose product weight is an integer multiple
2643    /// of eight leaves the octant unchanged, so only the 128 bits starting at bit `e-120` of the
2644    /// biased stream contribute: their product with `m` carries the low three integer bits and
2645    /// 64 fraction bits of `|x|·4/π`. The fraction is renormalized before it becomes a float, so
2646    /// arguments that fall close to a multiple of π/4 keep their relative accuracy.
2647    ///
2648    /// Returns `(octant & 7, z)` on the same contract as [`Self::cephes_reduce`].
2649    fn payne_hanek_reduce(&mut self, magnitude: Id) -> (Id, Id) {
2650        let table = self.private_u32_array(TWO_OVER_PI_BITS);
2651        let bits = self.bitcast_u32(magnitude);
2652        let twenty_three = self.c_u32(23);
2653        let exponent_mask = self.c_u32(0xff);
2654        let exponent = self.shr(bits, twenty_three);
2655        let exponent = self.band(exponent, exponent_mask);
2656        let significand_mask = self.c_u32(0x007f_ffff);
2657        let implicit = self.c_u32(0x0080_0000);
2658        let significand = self.band(bits, significand_mask);
2659        let m = self.bor(significand, implicit);
2660
2661        // Window offset into the biased bit stream; the leading zero word keeps it positive for
2662        // every magnitude this path serves (|x| >= 8192, so the exponent is at least 140).
2663        let bias = self.c_u32(120);
2664        let offset = self.isub(exponent, bias);
2665        let five = self.c_u32(5);
2666        let thirty_one = self.c_u32(31);
2667        let base = self.shr(offset, five);
2668        let shift = self.band(offset, thirty_one);
2669        let thirty_two = self.c_u32(32);
2670        let complement = self.isub(thirty_two, shift);
2671        let complement = self.band(complement, thirty_one);
2672        let zero = self.c_u32(0);
2673        let aligned = self.ieq(shift, zero);
2674        let mut window = [0; 4];
2675        for (index, slot) in window.iter_mut().enumerate() {
2676            let step = self.c_u32(index as u32);
2677            let first = self.iadd(base, step);
2678            let one = self.c_u32(1);
2679            let second = self.iadd(first, one);
2680            let high = self.private_element(table, first);
2681            let low = self.private_element(table, second);
2682            let high = self.shl(high, shift);
2683            let low = self.shr(low, complement);
2684            let low = self.select_u32(aligned, zero, low);
2685            *slot = self.bor(high, low);
2686        }
2687
2688        // Schoolbook m × window, most significant window word first.
2689        let mut product = [0; 5];
2690        let mut carry = zero;
2691        for (index, word) in window.iter().rev().enumerate() {
2692            let (high, low) = self.mul_wide(m, *word);
2693            let (sum, overflow) = self.add_carry(low, carry);
2694            product[index] = sum;
2695            carry = self.iadd(high, overflow);
2696        }
2697        product[4] = carry;
2698
2699        // Bit 125 of the product is the binary point: three integer bits above it, the fraction
2700        // below.
2701        let twenty_nine = self.c_u32(29);
2702        let three = self.c_u32(3);
2703        let seven = self.c_u32(7);
2704        let integer_low = self.shr(product[3], twenty_nine);
2705        let integer_high = self.shl(product[4], three);
2706        let integer = self.bor(integer_low, integer_high);
2707        let integer = self.band(integer, seven);
2708        let fraction_high = self.shl(product[3], three);
2709        let carry_in = self.shr(product[2], twenty_nine);
2710        let fraction_high = self.bor(fraction_high, carry_in);
2711        let fraction_low = self.shl(product[2], three);
2712        let carry_in = self.shr(product[1], twenty_nine);
2713        let fraction_low = self.bor(fraction_low, carry_in);
2714
2715        // Bump an odd octant to the next even one; the fraction becomes negative.
2716        let one = self.c_u32(1);
2717        let odd = self.band(integer, one);
2718        let is_odd = self.ine(odd, zero);
2719        let bumped = self.iadd(integer, one);
2720        let octant = self.select_u32(is_odd, bumped, integer);
2721        let octant = self.band(octant, seven);
2722        // Two's complement of the 64-bit fraction when it is subtracted from one.
2723        let negated_low = self.bnot(fraction_low);
2724        let (negated_low, overflow) = self.add_carry(negated_low, one);
2725        let negated_high = self.bnot(fraction_high);
2726        let negated_high = self.iadd(negated_high, overflow);
2727        let magnitude_high = self.select_u32(is_odd, negated_high, fraction_high);
2728        let magnitude_low = self.select_u32(is_odd, negated_low, fraction_low);
2729
2730        // Renormalize: the leading set bit sets the exponent, the next 24 bits the significand.
2731        let nonzero_high = self.ine(magnitude_high, zero);
2732        let msb_high = self.find_msb(magnitude_high);
2733        let msb_low = self.find_msb(magnitude_low);
2734        let shift_high = self.isub(thirty_one, msb_high);
2735        let shift_low = self.isub(thirty_one, msb_low);
2736        let shift_low_total = self.iadd(shift_low, thirty_two);
2737        let leading = self.select_u32(nonzero_high, shift_high, shift_low_total);
2738        // `magnitude_high << leading` with the bits shifted in from `magnitude_low`, or the low
2739        // word alone once the high word is empty.
2740        let complement = self.isub(thirty_two, leading);
2741        let complement_masked = self.band(complement, thirty_one);
2742        let aligned = self.ieq(leading, zero);
2743        let top_high = self.shl(magnitude_high, leading);
2744        let carried = self.shr(magnitude_low, complement_masked);
2745        let carried = self.select_u32(aligned, zero, carried);
2746        let top_high = self.bor(top_high, carried);
2747        let top_low = self.shl(magnitude_low, shift_low);
2748        let top = self.select_u32(nonzero_high, top_high, top_low);
2749        let eight = self.c_u32(8);
2750        let significand = self.shr(top, eight);
2751        let value = self.u_to_f(significand);
2752        // The renormalized fraction is `significand · 2^-(24 + leading)`.
2753        let twenty_four = self.c_u32(24);
2754        let exponent = self.iadd(twenty_four, leading);
2755        let exponent = self.isub(zero, exponent);
2756        let scale = self.pow2(exponent);
2757        let fraction = self.fmul(value, scale);
2758        // A fraction of exactly zero renormalizes to nothing; the octant is already correct.
2759        let empty_low = self.ieq(magnitude_low, zero);
2760        let empty_high = self.ieq(magnitude_high, zero);
2761        let empty = self.land(empty_low, empty_high);
2762        let zero_f = self.c_f32(0.0);
2763        let fraction = self.select_f32(empty, zero_f, fraction);
2764        let negated = self.fneg(fraction);
2765        let fraction = self.select_f32(is_odd, negated, fraction);
2766        // z = fraction · π/4, split so the product keeps more than binary32 precision.
2767        let pi_over_four_high = self.c_f32(0.785_156_25);
2768        let pi_over_four_low = self.c_f32(2.419_134e-4);
2769        let high = self.fmul(fraction, pi_over_four_high);
2770        let low = self.fmul(fraction, pi_over_four_low);
2771        let z = self.fadd(high, low);
2772        (octant, z)
2773    }
2774
2775    /// `sin` or `cos` from crate-authored range reduction and minimax polynomials: the Cephes
2776    /// three-part subtraction below [`SINCOS_FAST_RANGE`], Payne–Hanek above it, and NaN for a
2777    /// non-finite argument. The driver's own `sin`/`cos` are never used: Vulkan bounds them only
2778    /// to an absolute 2⁻¹¹ inside an unspecified range, which is not a numerical contract a tier
2779    /// can advertise.
2780    fn sincos(&mut self, x: Id, cosine: bool) -> Id {
2781        let magnitude = self.fabs(x);
2782        let (fast_octant, fast_z) = self.cephes_reduce(magnitude);
2783        let (exact_octant, exact_z) = self.payne_hanek_reduce(magnitude);
2784        let threshold = self.c_f32(SINCOS_FAST_RANGE);
2785        let fast = self.folt(magnitude, threshold);
2786        let octant = self.select_u32(fast, fast_octant, exact_octant);
2787        let z = self.select_f32(fast, fast_z, exact_z);
2788
2789        let three = self.c_u32(3);
2790        let reflect = self.ult(three, octant);
2791        let four = self.c_u32(4);
2792        let reduced = self.isub(octant, four);
2793        let octant = self.select_u32(reflect, reduced, octant);
2794        let zz = self.fmul(z, z);
2795        // cos polynomial: 1 - zz/2 + zz² · (c0 zz² + c1 zz + c2)
2796        let cos_poly = self.horner(zz, &[2.443_315_7e-5, -1.388_731_6e-3, 4.166_664_6e-2]);
2797        let zz2 = self.fmul(zz, zz);
2798        let cos_tail = self.fmul(cos_poly, zz2);
2799        let half = self.c_f32(0.5);
2800        let half_zz = self.fmul(half, zz);
2801        let cos_value = self.fsub(cos_tail, half_zz);
2802        let one_f = self.c_f32(1.0);
2803        let cos_value = self.fadd(cos_value, one_f);
2804        // sin polynomial: z + z · zz · (s0 zz² + s1 zz + s2)
2805        let sin_poly = self.horner(zz, &[-1.951_529_6e-4, 8.332_161e-3, -1.666_665_5e-1]);
2806        let sin_tail = self.fmul(sin_poly, zz);
2807        let sin_tail = self.fmul(sin_tail, z);
2808        let sin_value = self.fadd(sin_tail, z);
2809        let one = self.c_u32(1);
2810        let two = self.c_u32(2);
2811        let octant_is_1 = self.ieq(octant, one);
2812        let octant_is_2 = self.ieq(octant, two);
2813        let middle = self.lor(octant_is_1, octant_is_2);
2814        // sin uses the cos polynomial in octants 1..2, cos uses the sin polynomial there.
2815        let value = if cosine {
2816            self.select_f32(middle, sin_value, cos_value)
2817        } else {
2818            self.select_f32(middle, cos_value, sin_value)
2819        };
2820        let mut negate = reflect;
2821        if cosine {
2822            let upper_half = self.ult(one, octant);
2823            negate = self.lxor(negate, upper_half);
2824        } else {
2825            // The sign bit, not an ordered comparison: `sin(-0)` is `-0`.
2826            let negative_input = self.signbit(x);
2827            negate = self.lxor(negate, negative_input);
2828        }
2829        let negated = self.fneg(value);
2830        let value = self.select_f32(negate, negated, value);
2831        // `sin` and `cos` of a non-finite argument are NaN, and no reduction defines one.
2832        let infinity = self.c_f32(f32::INFINITY);
2833        let finite = self.folt(magnitude, infinity);
2834        let nan = self.c_f32(f32::NAN);
2835        self.select_f32(finite, value, nan)
2836    }
2837
2838    /// Cephes `tanhf`: odd polynomial below 0.625, `1 - 2 / (exp(2|x|) + 1)` above.
2839    fn tanh(&mut self, x: Id) -> Id {
2840        let magnitude = self.fabs(x);
2841        let square = self.fmul(x, x);
2842        let poly = self.horner(
2843            square,
2844            &[
2845                -5.704_988_7e-3,
2846                2.063_909e-2,
2847                -5.373_971_6e-2,
2848                1.333_144_2e-1,
2849                -3.333_328e-1,
2850            ],
2851        );
2852        let small = self.fmul(poly, square);
2853        let small = self.fma_free(small, x, x);
2854        let two = self.c_f32(2.0);
2855        let doubled = self.fmul(magnitude, two);
2856        let exp = self.ext_f32(GLSL_EXP, &[doubled]);
2857        let one = self.c_f32(1.0);
2858        let denominator = self.fadd(exp, one);
2859        let ratio = self.fdiv(two, denominator);
2860        let large = self.fsub(one, ratio);
2861        let negative = self.signbit(x);
2862        let negated = self.fneg(large);
2863        let large = self.select_f32(negative, negated, large);
2864        let threshold = self.c_f32(0.625);
2865        let use_large = self.foge(magnitude, threshold);
2866        let value = self.select_f32(use_large, large, small);
2867        // `tanh(±0) = ±0`: the polynomial's `0 + (-0)` would round the sign away.
2868        let zero = self.c_f32(0.0);
2869        let is_zero = self.foeq(x, zero);
2870        self.select_f32(is_zero, x, value)
2871    }
2872
2873    /// `erf`: the alternating Maclaurin series through `x^21` below `|x| = 1` (relative accuracy
2874    /// near zero), and `1 - erfc` with the Chebyshev-fitted `erfc` rational form above it
2875    /// (fractional error below 1.2e-7 everywhere).
2876    fn erf(&mut self, x: Id) -> Id {
2877        let magnitude = self.fabs(x);
2878        let square = self.fmul(x, x);
2879        let series = self.horner(
2880            square,
2881            &[
2882                1.0 / 76_204_800.0,
2883                -1.0 / 6_894_720.0,
2884                1.0 / 685_440.0,
2885                -1.0 / 75_600.0,
2886                1.0 / 9_360.0,
2887                -1.0 / 1_320.0,
2888                1.0 / 216.0,
2889                -1.0 / 42.0,
2890                1.0 / 10.0,
2891                -1.0 / 3.0,
2892                1.0,
2893            ],
2894        );
2895        let two_over_sqrt_pi = self.c_f32(core::f32::consts::FRAC_2_SQRT_PI);
2896        let series = self.fmul(series, two_over_sqrt_pi);
2897        let series = self.fmul(series, x);
2898        let half = self.c_f32(0.5);
2899        let one = self.c_f32(1.0);
2900        let half_magnitude = self.fmul(magnitude, half);
2901        let denominator = self.fadd(one, half_magnitude);
2902        let t = self.fdiv(one, denominator);
2903        let poly = self.horner(
2904            t,
2905            &[
2906                0.170_872_77,
2907                -0.822_152_23,
2908                1.488_515_9,
2909                -1.135_204,
2910                0.278_868_07,
2911                -0.186_288_06,
2912                0.096_784_18,
2913                0.374_091_96,
2914                1.000_023_7,
2915                -1.265_512_2,
2916            ],
2917        );
2918        let exponent = self.fsub(poly, square);
2919        let exp = self.ext_f32(GLSL_EXP, &[exponent]);
2920        let erfc = self.fmul(t, exp);
2921        let tail = self.fsub(one, erfc);
2922        let zero = self.c_f32(0.0);
2923        let negative = self.folt(x, zero);
2924        let negated = self.fneg(tail);
2925        let tail = self.select_f32(negative, negated, tail);
2926        let use_series = self.folt(magnitude, one);
2927        let value = self.select_f32(use_series, series, tail);
2928        // `erf(±0) = ±0` exactly, whatever the series rounds to.
2929        let is_zero = self.foeq(x, zero);
2930        self.select_f32(is_zero, x, value)
2931    }
2932
2933    /// IEEE-style `pow` on top of the built-in: negative bases with integral exponents keep the
2934    /// parity sign, negative bases with fractional exponents are NaN, `pow(x, 0) = 1`, and
2935    /// `pow(1, y) = 1`.
2936    fn pow(&mut self, x: Id, y: Id) -> Id {
2937        let magnitude = self.fabs(x);
2938        let raw = self.ext_f32(GLSL_POW, &[magnitude, y]);
2939        let zero = self.c_f32(0.0);
2940        let one = self.c_f32(1.0);
2941        let half = self.c_f32(0.5);
2942        let floor_y = self.ext_f32(GLSL_FLOOR, &[y]);
2943        let y_integral = self.foeq(floor_y, y);
2944        let half_y = self.fmul(y, half);
2945        let floor_half = self.ext_f32(GLSL_FLOOR, &[half_y]);
2946        let y_even = self.foeq(floor_half, half_y);
2947        let y_odd = self.lnot(y_even);
2948        let y_odd = self.land(y_integral, y_odd);
2949        let negated = self.fneg(raw);
2950        let signed = self.select_f32(y_odd, negated, raw);
2951        let nan = self.c_f32(f32::NAN);
2952        let negative_base = self.select_f32(y_integral, signed, nan);
2953        let x_negative = self.folt(x, zero);
2954        // A zero base: the built-in is undefined for y <= 0, so spell the limits out.
2955        let x_zero = self.foeq(x, zero);
2956        let y_positive = self.fogt(y, zero);
2957        let inf = self.c_f32(f32::INFINITY);
2958        let zero_base = self.select_f32(y_positive, zero, inf);
2959        let zero_base_neg = self.fneg(zero_base);
2960        let x_sign = self.bitcast_f32_sign(x);
2961        let x_sign_negative = self.folt(x_sign, zero);
2962        let zero_base_signed = self.land(x_sign_negative, y_odd);
2963        let zero_base = self.select_f32(zero_base_signed, zero_base_neg, zero_base);
2964        let value = self.select_f32(x_negative, negative_base, raw);
2965        let value = self.select_f32(x_zero, zero_base, value);
2966        let y_zero = self.foeq(y, zero);
2967        let value = self.select_f32(y_zero, one, value);
2968        let x_one = self.foeq(x, one);
2969        self.select_f32(x_one, one, value)
2970    }
2971
2972    /// Whether `x` carries the IEEE sign bit; true for `-0.0`, unlike an ordered `x < 0`.
2973    fn signbit(&mut self, x: Id) -> Id {
2974        let bits = self.bitcast_u32(x);
2975        let sign_bit = self.c_u32(0x8000_0000);
2976        let masked = self.band(bits, sign_bit);
2977        let zero = self.c_u32(0);
2978        self.ine(masked, zero)
2979    }
2980
2981    /// `-1.0` when `x` carries the sign bit (including `-0.0`), else `+1.0`.
2982    fn bitcast_f32_sign(&mut self, x: Id) -> Id {
2983        let bits = self.bitcast_u32(x);
2984        let sign_bit = self.c_u32(0x8000_0000);
2985        let masked = self.band(bits, sign_bit);
2986        let zero = self.c_u32(0);
2987        let negative = self.ine(masked, zero);
2988        let minus_one = self.c_f32(-1.0);
2989        let plus_one = self.c_f32(1.0);
2990        self.select_f32(negative, minus_one, plus_one)
2991    }
2992
2993    /// The scalar lane of `op` over already-loaded inputs (`f32` ids for float inputs — binary16
2994    /// tensors are widened on load by [`Self::widen_f16`] and narrowed on store by
2995    /// [`Self::narrow_f16`], so every float lane is binary32 — and `bool` ids for byte inputs),
2996    /// yielding an `f32` or `bool` id per [`ElementwiseOp::output`].
2997    fn elementwise_lane(
2998        &mut self,
2999        op: ElementwiseOp,
3000        inputs: &[Id],
3001        clamp: Option<(Id, Id)>,
3002    ) -> Id {
3003        let x = inputs[0];
3004        match op {
3005            ElementwiseOp::Abs => self.fabs(x),
3006            ElementwiseOp::Ceil => self.ext_f32(GLSL_CEIL, &[x]),
3007            ElementwiseOp::Floor => self.ext_f32(GLSL_FLOOR, &[x]),
3008            ElementwiseOp::Cos => self.sincos(x, true),
3009            ElementwiseOp::Sin => self.sincos(x, false),
3010            ElementwiseOp::Erf => self.erf(x),
3011            ElementwiseOp::Exp => self.ext_f32(GLSL_EXP, &[x]),
3012            ElementwiseOp::Log => self.ext_f32(GLSL_LOG, &[x]),
3013            ElementwiseOp::Negate => self.fneg(x),
3014            ElementwiseOp::Reciprocal => {
3015                let one = self.c_f32(1.0);
3016                self.fdiv(one, x)
3017            }
3018            ElementwiseOp::Rsqrt => self.ext_f32(GLSL_INVERSE_SQRT, &[x]),
3019            ElementwiseOp::Sigmoid => {
3020                let negated = self.fneg(x);
3021                let exp = self.ext_f32(GLSL_EXP, &[negated]);
3022                let one = self.c_f32(1.0);
3023                let denominator = self.fadd(one, exp);
3024                self.fdiv(one, denominator)
3025            }
3026            ElementwiseOp::Tanh => self.tanh(x),
3027            ElementwiseOp::Clamp(nan_mode) => {
3028                // The bounds arrive as binary32 bit patterns (binary16 bounds are widened
3029                // host-side, exactly).
3030                let (lo_bits, hi_bits) = clamp.expect("clamp bounds");
3031                let lo = self.bitcast_f32(lo_bits);
3032                let hi = self.bitcast_f32(hi_bits);
3033                let floored = self.apply_max(x, lo, nan_mode);
3034                self.apply_min(floored, hi, nan_mode)
3035            }
3036            ElementwiseOp::Add => self.fadd(x, inputs[1]),
3037            ElementwiseOp::Sub => self.fsub(x, inputs[1]),
3038            ElementwiseOp::Mul => self.fmul(x, inputs[1]),
3039            ElementwiseOp::Pow => self.pow(x, inputs[1]),
3040            ElementwiseOp::Maximum(nan_mode) => self.apply_max(x, inputs[1], nan_mode),
3041            ElementwiseOp::Minimum(nan_mode) => self.apply_min(x, inputs[1], nan_mode),
3042            ElementwiseOp::Equal => self.foeq(x, inputs[1]),
3043            ElementwiseOp::Greater => self.fogt(x, inputs[1]),
3044            ElementwiseOp::GreaterEqual => self.foge(x, inputs[1]),
3045            ElementwiseOp::LogicalAnd => self.land(x, inputs[1]),
3046            ElementwiseOp::LogicalOr => self.lor(x, inputs[1]),
3047            ElementwiseOp::LogicalXor => self.lxor(x, inputs[1]),
3048            ElementwiseOp::LogicalNot => self.lnot(x),
3049            ElementwiseOp::Select => self.select_f32(x, inputs[1], inputs[2]),
3050            ElementwiseOp::CopyBytes => x,
3051        }
3052    }
3053}
3054
3055#[derive(Clone, Copy)]
3056struct LoopScope {
3057    header: Id,
3058    body: Id,
3059    cont: Id,
3060    merge: Id,
3061}
3062
3063// ---------------------------------------------------------------------------------------------
3064// Kernels
3065// ---------------------------------------------------------------------------------------------
3066
3067/// Elementwise kernel: grid-stride over `count` output elements, with `float` the storage of
3068/// the operator's floating-point tensors (`Word` for FP32, `Half` for FP16).
3069///
3070/// Specialization order: `count`; per input `(buffer, base)`; output `(buffer, base)`; when
3071/// `broadcast`, output `dims[MAX_RANK]` then per input `strides[MAX_RANK]`; then the operator's
3072/// trailing constants (`CLAMP`: `lo`, `hi` as binary32 bit patterns — binary16 bounds are
3073/// widened host-side).
3074fn assemble_elementwise(
3075    op: ElementwiseOp,
3076    float: Storage,
3077    broadcast: bool,
3078    workgroup: u32,
3079    buffers: u32,
3080) -> Vec<u32> {
3081    let mut b = Builder::new();
3082    let array = b.buffer_array(buffers);
3083    let count = b.spec_u32(1);
3084    let input_storage = op.inputs();
3085    // Lane tables mark floating-point operands `Word`; resolve those to this kernel's float
3086    // storage, keeping `BOOL` lanes on byte storage.
3087    let lane = |storage: Storage| match storage {
3088        Storage::Word => float,
3089        other => other,
3090    };
3091    let inputs: Vec<(Id, Id)> = input_storage.iter().map(|_| b.spec_operand()).collect();
3092    let output = b.spec_operand();
3093    let shape = broadcast.then(|| {
3094        let dims = b.spec_dims();
3095        let strides: Vec<[Id; MAX_RANK]> = input_storage.iter().map(|_| b.spec_strides()).collect();
3096        (dims, strides)
3097    });
3098    let clamp = matches!(op, ElementwiseOp::Clamp(_)).then(|| {
3099        let lo = b.spec_u32(0);
3100        let hi = b.spec_u32(0);
3101        (lo, hi)
3102    });
3103
3104    let (counter, stride) = b.grid_stride(workgroup);
3105    let (scope, i) = b.begin_loop(counter, count);
3106    let indices = match &shape {
3107        Some((dims, strides)) => b.strided_indices(i, dims, strides),
3108        None => vec![i; inputs.len()],
3109    };
3110    // IEEE negate and abs are sign-bit operations, not arithmetic: on binary16 they are integer
3111    // masks on the packed lane, exact for every pattern — subnormals and NaN payloads included
3112    // (ADR 0008). Every other binary16 lane widens exactly, evaluates at binary32, and narrows
3113    // once, all in crate-owned integer/binary32 code no driver can demote.
3114    let bits_lane =
3115        float == Storage::Half && matches!(op, ElementwiseOp::Negate | ElementwiseOp::Abs);
3116    let mut values = Vec::with_capacity(inputs.len());
3117    for (k, operand) in inputs.iter().enumerate() {
3118        let value = match lane(input_storage[k]) {
3119            Storage::Word => b.load_f32(array, *operand, indices[k]),
3120            Storage::Half if bits_lane => b.load_half_bits(array, *operand, indices[k]),
3121            Storage::Half => {
3122                let bits = b.load_half_bits(array, *operand, indices[k]);
3123                b.widen_f16(bits)
3124            }
3125            Storage::Byte => b.load_bool(array, *operand, indices[k]),
3126            // TOSA admits no FP8 elementwise operator, so the tier never selects this lane.
3127            Storage::Quarter(_) => unreachable!("elementwise FP8 operand"),
3128        };
3129        values.push(value);
3130    }
3131    let result = if bits_lane {
3132        let bits = values[0];
3133        match op {
3134            ElementwiseOp::Negate => {
3135                let sign = b.c_u32(0x8000);
3136                b.bxor(bits, sign)
3137            }
3138            ElementwiseOp::Abs => {
3139                let magnitude = b.c_u32(0x7fff);
3140                b.band(bits, magnitude)
3141            }
3142            _ => unreachable!("bits lane is negate/abs only"),
3143        }
3144    } else {
3145        b.elementwise_lane(op, &values, clamp)
3146    };
3147    match lane(op.output()) {
3148        Storage::Word => b.store_f32(array, output, i, result),
3149        Storage::Half if bits_lane => b.store_half_bits(array, output, i, result),
3150        Storage::Half => {
3151            let bits = b.narrow_f16(result);
3152            b.store_half_bits(array, output, i, bits);
3153        }
3154        Storage::Byte => b.store_bool(array, output, i, result),
3155        Storage::Quarter(_) => unreachable!("elementwise FP8 result"),
3156    }
3157    b.end_loop(scope, counter, stride);
3158    b.end_main();
3159    b.finish([workgroup, 1, 1])
3160}
3161
3162/// Reduction kernel: one invocation per `(outer, inner)` output element folds the axis in
3163/// ascending order. Inputs are read at `float` storage and widened to binary32 for the fold —
3164/// the accumulator width TOSA assigns FP16 sums and products, and exact for the max/min/argmax
3165/// selections — then narrowed back to `float` storage on store.
3166///
3167/// Specialization order: input `(buffer, base)`, output `(buffer, base)`, `outer`, `axis`,
3168/// `inner`.
3169fn assemble_reduce(op: ReduceOp, float: Storage, workgroup: u32, buffers: u32) -> Vec<u32> {
3170    let mut b = Builder::new();
3171    let array = b.buffer_array(buffers);
3172    let input = b.spec_operand();
3173    let output = b.spec_operand();
3174    let outer = b.spec_u32(1);
3175    let axis = b.spec_u32(1);
3176    let inner = b.spec_u32(1);
3177
3178    let (counter, stride) = b.grid_stride(workgroup);
3179    let u32_ty = b.u32_ty();
3180    let f32_ty = b.f32_ty();
3181    let acc_var = b.local(f32_ty);
3182    let index_var = b.local(u32_ty);
3183    let done_var = {
3184        let bool_ty = b.bool_ty();
3185        b.local(bool_ty)
3186    };
3187    let a_var = b.local(u32_ty);
3188    let count = b.imul(outer, inner);
3189    let (scope, o) = b.begin_loop(counter, count);
3190    let outer_index = b.udiv(o, inner);
3191    let inner_index = b.umod(o, inner);
3192    let row = b.imul(outer_index, axis);
3193    let row = b.imul(row, inner);
3194    let base = b.iadd(row, inner_index);
3195    let init = match op {
3196        ReduceOp::Sum => b.c_f32(0.0),
3197        ReduceOp::Product => b.c_f32(1.0),
3198        ReduceOp::Max(_) | ReduceOp::ArgMax(_) => b.c_f32(f32::NEG_INFINITY),
3199        ReduceOp::Min(_) => b.c_f32(f32::INFINITY),
3200    };
3201    b.store(acc_var, init);
3202    let zero = b.c_u32(0);
3203    let one = b.c_u32(1);
3204    b.store(index_var, zero);
3205    let false_id = b.c_false();
3206    b.store(done_var, false_id);
3207    b.store(a_var, zero);
3208    let (inner_scope, a) = b.begin_loop(a_var, axis);
3209    let offset = b.imul(a, inner);
3210    let element = b.iadd(base, offset);
3211    let value = b.load_float(float, array, input, element);
3212    let acc = b.load(f32_ty, acc_var);
3213    match op {
3214        ReduceOp::Sum => {
3215            let next = b.fadd(acc, value);
3216            b.store(acc_var, next);
3217        }
3218        ReduceOp::Product => {
3219            let next = b.fmul(acc, value);
3220            b.store(acc_var, next);
3221        }
3222        ReduceOp::Max(nan_mode) => {
3223            let next = b.apply_max(acc, value, nan_mode);
3224            b.store(acc_var, next);
3225        }
3226        ReduceOp::Min(nan_mode) => {
3227            let next = b.apply_min(acc, value, nan_mode);
3228            b.store(acc_var, next);
3229        }
3230        ReduceOp::ArgMax(nan_mode) => {
3231            let bool_ty = b.bool_ty();
3232            let done = b.load(bool_ty, done_var);
3233            let index = b.load(u32_ty, index_var);
3234            let greater = b.fogt(value, acc);
3235            let value_nan = b.is_nan(value);
3236            let take = match nan_mode {
3237                NanMode::Propagate => {
3238                    // The first NaN wins and freezes the result.
3239                    let candidate = b.lor(greater, value_nan);
3240                    let not_done = b.lnot(done);
3241                    let take = b.land(candidate, not_done);
3242                    let next_done = b.lor(done, value_nan);
3243                    b.store(done_var, next_done);
3244                    take
3245                }
3246                NanMode::Ignore => {
3247                    let ordered = b.lnot(value_nan);
3248                    b.land(greater, ordered)
3249                }
3250            };
3251            let next_acc = b.select_f32(take, value, acc);
3252            let next_index = b.select_u32(take, a, index);
3253            b.store(acc_var, next_acc);
3254            b.store(index_var, next_index);
3255        }
3256    }
3257    b.end_loop(inner_scope, a_var, one);
3258    match op {
3259        ReduceOp::ArgMax(_) => {
3260            let index = b.load(u32_ty, index_var);
3261            b.store_word(array, output, o, index);
3262        }
3263        _ => {
3264            let acc = b.load(f32_ty, acc_var);
3265            b.store_float(float, array, output, o, acc);
3266        }
3267    }
3268    b.end_loop(scope, counter, stride);
3269    b.end_main();
3270    b.finish([workgroup, 1, 1])
3271}
3272
3273/// One operand slab of the MATMUL kernel: `rows × cols` elements whose element `(r, c)` lives at
3274/// `base + (row0 + r) · stride + (col0 + c)` and is in range when `row0 + r < row_limit` and
3275/// `col0 + c < col_limit`; staged to `tile` at slot `c · rows + r` when `transposed`, else
3276/// `r · cols + c`.
3277struct Slab {
3278    operand: (Id, Id),
3279    rows: u32,
3280    cols: u32,
3281    row0: Id,
3282    col0: Id,
3283    stride: Id,
3284    base: Id,
3285    row_limit: Id,
3286    col_limit: Id,
3287    transposed: bool,
3288    tile: Id,
3289}
3290
3291/// Stage a slab cooperatively: the `invocations` of the workgroup take slab elements (or, for
3292/// sub-word storage, slab *words*) `q · invocations + lid` in turn, so consecutive invocations
3293/// read consecutive addresses. Out-of-range elements read element zero and stage `0.0`.
3294///
3295/// Sub-word storage is staged a storage word per invocation — four FP8 or two FP16 elements
3296/// from one `OpLoad` — whenever `stride` is a multiple of the lanes per word, which makes every
3297/// slab row word-aligned (`col0` and `base` are multiples of the row length or of `stride`).
3298/// The check is on a specialization constant, so the driver folds it and only one path
3299/// survives pipeline creation. This is what makes FP8 cheaper than FP32 rather than merely
3300/// smaller: on Intel Xe3 the per-element path issued one load instruction per element whatever
3301/// the width, and a GEMV over 16 MiB of FP8 weights ran no faster than one over 64 MiB of FP32,
3302/// because the memory pipeline was bound by load instructions rather than bytes.
3303fn stage_slab(b: &mut Builder, input: Storage, array: Id, invocations: u32, lid: Id, slab: &Slab) {
3304    let f32_ty = b.f32_ty();
3305    let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3306    let zero = b.c_u32(0);
3307    let zero_f = b.c_f32(0.0);
3308    let rows_c = b.c_u32(slab.rows);
3309    let cols_c = b.c_u32(slab.cols);
3310    let elements = slab.rows * slab.cols;
3311    let slot = |b: &mut Builder, r: Id, c: Id| {
3312        if slab.transposed {
3313            let slot = b.imul(c, rows_c);
3314            b.iadd(slot, r)
3315        } else {
3316            let slot = b.imul(r, cols_c);
3317            b.iadd(slot, c)
3318        }
3319    };
3320    // Element `(r, c)`'s global index and validity.
3321    let locate = |b: &mut Builder, r: Id, c: Id| {
3322        let global_row = b.iadd(slab.row0, r);
3323        let global_col = b.iadd(slab.col0, c);
3324        let row_ok = b.ult(global_row, slab.row_limit);
3325        let col_ok = b.ult(global_col, slab.col_limit);
3326        let ok = b.land(row_ok, col_ok);
3327        let element = b.imul(global_row, slab.stride);
3328        let element = b.iadd(element, slab.base);
3329        let element = b.iadd(element, global_col);
3330        let element = b.select_u32(ok, element, zero);
3331        (element, ok)
3332    };
3333    // `q · invocations + lid` over `count` items, guarded only when the count does not divide.
3334    let each = |b: &mut Builder, count: u32, body: &dyn Fn(&mut Builder, Id)| {
3335        for q in 0..count.div_ceil(invocations) {
3336            let offset = b.c_u32(q * invocations);
3337            let index = b.iadd(offset, lid);
3338            if count % invocations == 0 {
3339                body(b, index);
3340            } else {
3341                let limit = b.c_u32(count);
3342                let in_range = b.ult(index, limit);
3343                b.if_then(in_range, |b| body(b, index));
3344            }
3345        }
3346    };
3347    let scalar = |b: &mut Builder| {
3348        each(b, elements, &|b, index| {
3349            let r = b.udiv(index, cols_c);
3350            let c = b.umod(index, cols_c);
3351            let (element, ok) = locate(b, r, c);
3352            let loaded = b.load_float(input, array, slab.operand, element);
3353            let value = b.select_f32(ok, loaded, zero_f);
3354            let slot = slot(b, r, c);
3355            let pointer = b.access_chain(workgroup_ptr, slab.tile, &[slot]);
3356            b.store(pointer, value);
3357        });
3358    };
3359    let lanes = input.lanes();
3360    if lanes == 1 || slab.cols % lanes != 0 {
3361        scalar(b);
3362        return;
3363    }
3364    let lanes_c = b.c_u32(lanes);
3365    let words_per_row = slab.cols / lanes;
3366    let words_per_row_c = b.c_u32(words_per_row);
3367    let remainder = b.umod(slab.stride, lanes_c);
3368    let aligned = b.ieq(remainder, zero);
3369    b.if_then(aligned, |b| {
3370        each(b, elements / lanes, &|b, word| {
3371            let r = b.udiv(word, words_per_row_c);
3372            let word_in_row = b.umod(word, words_per_row_c);
3373            let c = b.imul(word_in_row, lanes_c);
3374            // `stride`, `base` and `col0` are all multiples of `lanes` here and `col_limit`
3375            // need not be, but a word is either wholly inside the tensor or wholly past it
3376            // only when `col_limit` is a multiple of `lanes`; check the last lane too.
3377            let (element, ok) = locate(b, r, c);
3378            let last = b.c_u32(lanes - 1);
3379            let last_col = b.iadd(c, last);
3380            let (_, last_ok) = locate(b, r, last_col);
3381            let ok = b.land(ok, last_ok);
3382            let element = b.select_u32(ok, element, zero);
3383            let bits = b.load_lane_group(input, array, slab.operand, element, lanes);
3384            for (lane, bits) in bits.into_iter().enumerate() {
3385                let widened = b.widen_bits(input, bits);
3386                let value = b.select_f32(ok, widened, zero_f);
3387                let lane_c = b.c_u32(lane as u32);
3388                let col = b.iadd(c, lane_c);
3389                let slot = slot(b, r, col);
3390                let pointer = b.access_chain(workgroup_ptr, slab.tile, &[slot]);
3391                b.store(pointer, value);
3392            }
3393        });
3394    });
3395    let unaligned = b.lnot(aligned);
3396    b.if_then(unaligned, scalar);
3397}
3398
3399/// Batched, register-tiled MATMUL reading `input` storage (FP32 words, packed FP16, or packed
3400/// FP8) and accumulating in binary32 — the accumulator width TOSA assigns FP16 and FP8 MATMUL,
3401/// and the FP32 tier's own width.
3402///
3403/// `out[b, m, n] = Σ_k lhs[b, m, k] · rhs[b, k, n]`, accumulated in ascending `k` with separately
3404/// rounded multiply and add, so every element is bit-identical to the untiled sequential loop.
3405///
3406/// Geometry ([`MatmulGeometry`]): a `tile_x × tile_y` workgroup computes a `block_m × block_n`
3407/// output block, each invocation a `micro_m × micro_n` register block whose rows are
3408/// `i · tile_y + ty` and columns `j · tile_x + tx` — interleaved rather than contiguous, so the
3409/// invocations of a row of the workgroup read consecutive columns and write consecutive
3410/// outputs. Per `depth`-deep step of `k` the workgroup stages a slab of each operand in shared
3411/// memory cooperatively (consecutive invocations loading consecutive addresses, the lhs slab
3412/// stored transposed so its inner-loop reads are conflict-free), then every invocation issues
3413/// `micro_m + micro_n` shared loads for `micro_m · micro_n` multiply-adds. The
3414/// one-output-per-invocation kernel this replaces issued two shared loads per multiply-add and
3415/// staged one element per invocation; on Intel Xe3 it ran a 1024³ FP8 MATMUL at 570 GFLOP/s,
3416/// below its own FP32 rate, because the per-element FP8 widening was paid once per multiply-add
3417/// rather than once per `MATMUL_MICRO` of them. (A 256-entry shared-memory widening table,
3418/// built per workgroup, was measured against the inline integer expansion on Intel Xe3 and made
3419/// no difference within noise; the expansion stays, having no table to build or barrier to wait
3420/// on.)
3421///
3422/// Out-of-range slab loads read element zero and contribute nothing: the inner loop bound is
3423/// `min(depth, k - k0)`, never a padded zero product, so signed zeros survive. Out-of-range
3424/// outputs are computed and discarded.
3425///
3426/// Specialization order: `lhs`, `rhs`, output `(buffer, base)`; `m`, `n`, `k`, `batch`.
3427fn assemble_matmul(
3428    input: Storage,
3429    output_storage: Storage,
3430    geometry: MatmulGeometry,
3431    buffers: u32,
3432) -> Vec<u32> {
3433    let MatmulGeometry {
3434        tile_x,
3435        tile_y,
3436        micro_m,
3437        micro_n,
3438        depth,
3439    } = geometry;
3440    let invocations = geometry.invocations();
3441    let block_m = geometry.block_m();
3442    let block_n = geometry.block_n();
3443    debug_assert_eq!((block_m * depth) % invocations, 0, "lhs slab stages evenly");
3444    debug_assert_eq!((block_n * depth) % invocations, 0, "rhs slab stages evenly");
3445    let mut b = Builder::new();
3446    let array = b.buffer_array(buffers);
3447    let lhs = b.spec_operand();
3448    let rhs = b.spec_operand();
3449    let output = b.spec_operand();
3450    let m = b.spec_u32(1);
3451    let n = b.spec_u32(1);
3452    let k = b.spec_u32(1);
3453    let _batch = b.spec_u32(1);
3454    // lhs slab transposed: `[kk][row]`, `depth × block_m`; rhs slab as is: `[kk][col]`,
3455    // `depth × block_n`.
3456    let lhs_tile = b.shared_f32_array(block_m * depth);
3457    let rhs_tile = b.shared_f32_array(block_n * depth);
3458    let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
3459    let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
3460
3461    b.begin_main();
3462    let u32_ty = b.u32_ty();
3463    let f32_ty = b.f32_ty();
3464    let accumulators: Vec<Id> = (0..micro_m * micro_n).map(|_| b.local(f32_ty)).collect();
3465    let t_var = b.local(u32_ty);
3466    let kk_var = b.local(u32_ty);
3467    let tx = b.builtin_component(local_id, 0);
3468    let ty = b.builtin_component(local_id, 1);
3469    let gx = b.builtin_component(group_id, 0);
3470    let gy = b.builtin_component(group_id, 1);
3471    let z = b.builtin_component(group_id, 2);
3472    let tile_x_c = b.c_u32(tile_x);
3473    let depth_c = b.c_u32(depth);
3474    let block_m_c = b.c_u32(block_m);
3475    let block_n_c = b.c_u32(block_n);
3476    let zero = b.c_u32(0);
3477    let one = b.c_u32(1);
3478    let zero_f = b.c_f32(0.0);
3479    for accumulator in &accumulators {
3480        b.store(*accumulator, zero_f);
3481    }
3482    b.store(t_var, zero);
3483    let row0 = b.imul(gy, block_m_c);
3484    let col0 = b.imul(gx, block_n_c);
3485    let lid = b.imul(ty, tile_x_c);
3486    let lid = b.iadd(lid, tx);
3487    // lhs batch base: z * m * k; rhs batch base: z * k * n.
3488    let lhs_batch = b.imul(z, m);
3489    let lhs_batch = b.imul(lhs_batch, k);
3490    let rhs_batch = b.imul(z, k);
3491    let rhs_batch = b.imul(rhs_batch, n);
3492    let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3493    let steps = {
3494        let k_plus = b.iadd(k, depth_c);
3495        let k_plus = b.isub(k_plus, one);
3496        b.udiv(k_plus, depth_c)
3497    };
3498    let (outer, t) = b.begin_loop(t_var, steps);
3499    let k0 = b.imul(t, depth_c);
3500    // Stage both slabs. lhs: `block_m` rows of `depth` consecutive `k`, stored transposed;
3501    // rhs: `depth` rows of `block_n` consecutive `n`.
3502    stage_slab(
3503        &mut b,
3504        input,
3505        array,
3506        invocations,
3507        lid,
3508        &Slab {
3509            operand: lhs,
3510            rows: block_m,
3511            cols: depth,
3512            row0,
3513            col0: k0,
3514            stride: k,
3515            base: lhs_batch,
3516            row_limit: m,
3517            col_limit: k,
3518            transposed: true,
3519            tile: lhs_tile,
3520        },
3521    );
3522    stage_slab(
3523        &mut b,
3524        input,
3525        array,
3526        invocations,
3527        lid,
3528        &Slab {
3529            operand: rhs,
3530            rows: depth,
3531            cols: block_n,
3532            row0: k0,
3533            col0,
3534            stride: n,
3535            base: rhs_batch,
3536            row_limit: k,
3537            col_limit: n,
3538            transposed: false,
3539            tile: rhs_tile,
3540        },
3541    );
3542    b.workgroup_barrier();
3543    let remaining = b.isub(k, k0);
3544    let k_max = b.umin(remaining, depth_c);
3545    // One `kk` of the inner product for every register-block element.
3546    let step = |b: &mut Builder, kk: Id| {
3547        let lhs_row = b.imul(kk, block_m_c);
3548        let rhs_row = b.imul(kk, block_n_c);
3549        let a: Vec<Id> = (0..micro_m)
3550            .map(|i| {
3551                let offset = b.c_u32(i * tile_y);
3552                let local_row = b.iadd(offset, ty);
3553                let slot = b.iadd(lhs_row, local_row);
3554                let pointer = b.access_chain(workgroup_ptr, lhs_tile, &[slot]);
3555                b.load(f32_ty, pointer)
3556            })
3557            .collect();
3558        let bv: Vec<Id> = (0..micro_n)
3559            .map(|j| {
3560                let offset = b.c_u32(j * tile_x);
3561                let local_col = b.iadd(offset, tx);
3562                let slot = b.iadd(rhs_row, local_col);
3563                let pointer = b.access_chain(workgroup_ptr, rhs_tile, &[slot]);
3564                b.load(f32_ty, pointer)
3565            })
3566            .collect();
3567        for i in 0..micro_m {
3568            for j in 0..micro_n {
3569                let accumulator = accumulators[(i * micro_n + j) as usize];
3570                let acc = b.load(f32_ty, accumulator);
3571                let next = b.fma(a[i as usize], bv[j as usize], acc);
3572                b.store(accumulator, next);
3573            }
3574        }
3575    };
3576    b.store(kk_var, zero);
3577    let (inner, kk) = b.begin_loop(kk_var, k_max);
3578    step(&mut b, kk);
3579    b.end_loop(inner, kk_var, one);
3580    b.workgroup_barrier();
3581    b.end_loop(outer, t_var, one);
3582    let out_batch = b.imul(z, m);
3583    let out_batch = b.imul(out_batch, n);
3584    for i in 0..micro_m {
3585        let offset = b.c_u32(i * tile_y);
3586        let row = b.iadd(row0, offset);
3587        let row = b.iadd(row, ty);
3588        let row_ok = b.ult(row, m);
3589        let out_row = b.imul(row, n);
3590        let out_row = b.iadd(out_batch, out_row);
3591        for j in 0..micro_n {
3592            let offset = b.c_u32(j * tile_x);
3593            let col = b.iadd(col0, offset);
3594            let col = b.iadd(col, tx);
3595            let col_ok = b.ult(col, n);
3596            let in_range = b.land(row_ok, col_ok);
3597            let accumulator = accumulators[(i * micro_n + j) as usize];
3598            b.if_then(in_range, |b| {
3599                let out_index = b.iadd(out_row, col);
3600                let acc = b.load(f32_ty, accumulator);
3601                match output_storage {
3602                    Storage::Half => {
3603                        let bits = b.narrow_f16(acc);
3604                        b.store_half_bits(array, output, out_index, bits);
3605                    }
3606                    // TOSA defines no MATMUL whose result is FP8 or BOOL; the tier never
3607                    // selects one.
3608                    Storage::Byte | Storage::Quarter(_) => unreachable!("MATMUL result storage"),
3609                    Storage::Word => b.store_f32(array, output, out_index, acc),
3610                }
3611            });
3612        }
3613    }
3614    b.end_main();
3615    b.finish(geometry.local_size())
3616}
3617
3618/// Split-`k` streaming MATMUL for `m ≤` [`STREAM_ROWS`]: the decode shape, a few activation rows
3619/// against a wide weight matrix, where the kernel's only job is to stream `rhs` once at the
3620/// memory system's rate.
3621///
3622/// A 1-D workgroup of [`STREAM_WORKGROUP`] invocations covers [`STREAM_COLUMNS`] columns.
3623/// Invocation `lid` owns weight word `lid % words` — `lanes` adjacent columns, four FP8, two
3624/// FP16 or one FP32 — and `k` slice `lid / words` of `splits = STREAM_WORKGROUP / words` equal
3625/// slices, and carries all [`STREAM_ROWS`] rows of those columns in registers. Its loop is one
3626/// `rhs` word load, one uniform binary32 lhs load per row, and `rows · lanes` fused
3627/// multiply-adds — no shared memory, no barrier, and no widening of the lhs, which lowering
3628/// pre-widens to binary32 (an `m × k` intermediate, negligible next to `rhs`) when the tier's
3629/// lhs is narrower. Rows at or past `m` are compiled out: the row test is on a specialization
3630/// constant. At the end every invocation writes its partials to shared memory and each
3631/// invocation reduces two outputs over the `splits` slices in ascending order, then stores.
3632///
3633/// Numerics (ADR 0011): binary32 accumulation with fused multiply-add, `k` split into `splits`
3634/// contiguous ascending slices summed in fixed order. Not bit-identical to the sequential sum,
3635/// but deterministic: the same device always produces the same bits, and every device whose
3636/// `Fma` is correctly rounded produces the same bits as every other.
3637///
3638/// The `rhs` word path needs every row of `rhs` word-aligned, i.e. `n` a multiple of `lanes`;
3639/// the check is on the specialization constant `n`, so only one path survives pipeline
3640/// creation, and an unaligned `n` falls back to per-element loads.
3641///
3642/// Specialization order: `lhs`, `rhs`, output `(buffer, base)`; `m`, `n`, `k`, `batch` — the
3643/// same payload as [`assemble_matmul`].
3644fn assemble_matmul_stream(rhs_storage: Storage, output_storage: Storage, buffers: u32) -> Vec<u32> {
3645    let lanes = rhs_storage.lanes();
3646    let words = STREAM_COLUMNS / lanes;
3647    let splits = STREAM_WORKGROUP / words;
3648    let rows = STREAM_ROWS;
3649    let mut b = Builder::new();
3650    let array = b.buffer_array(buffers);
3651    let lhs = b.spec_operand();
3652    let rhs = b.spec_operand();
3653    let output = b.spec_operand();
3654    let m = b.spec_u32(1);
3655    let n = b.spec_u32(1);
3656    let k = b.spec_u32(1);
3657    let _batch = b.spec_u32(1);
3658    let partials = b.shared_f32_array(STREAM_WORKGROUP * rows * lanes);
3659    let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
3660    let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
3661
3662    b.begin_main();
3663    let u32_ty = b.u32_ty();
3664    let f32_ty = b.f32_ty();
3665    let accumulators: Vec<Id> = (0..rows * lanes).map(|_| b.local(f32_ty)).collect();
3666    let kk_var = b.local(u32_ty);
3667    let lid = b.builtin_component(local_id, 0);
3668    let gx = b.builtin_component(group_id, 0);
3669    let z = b.builtin_component(group_id, 2);
3670    let zero = b.c_u32(0);
3671    let one = b.c_u32(1);
3672    let zero_f = b.c_f32(0.0);
3673    let words_c = b.c_u32(words);
3674    let lanes_c = b.c_u32(lanes);
3675    let splits_c = b.c_u32(splits);
3676    let columns_c = b.c_u32(STREAM_COLUMNS);
3677    for accumulator in &accumulators {
3678        b.store(*accumulator, zero_f);
3679    }
3680    let w = b.umod(lid, words_c);
3681    let slice = b.udiv(lid, words_c);
3682    let col0 = b.imul(gx, columns_c);
3683    let word_col = b.imul(w, lanes_c);
3684    let col_base = b.iadd(col0, word_col);
3685    // This invocation's `k` slice: `[slice · chunk, min(k, (slice + 1) · chunk))`.
3686    let splits_less_one = b.c_u32(splits - 1);
3687    let k_plus = b.iadd(k, splits_less_one);
3688    let chunk = b.udiv(k_plus, splits_c);
3689    let k_begin = b.imul(slice, chunk);
3690    let k_end = b.iadd(k_begin, chunk);
3691    let k_end = b.umin(k_end, k);
3692    let lhs_batch = b.imul(z, m);
3693    let lhs_batch = b.imul(lhs_batch, k);
3694    let rhs_batch = b.imul(z, k);
3695    let rhs_batch = b.imul(rhs_batch, n);
3696    let rhs_col = b.iadd(rhs_batch, col_base);
3697    let row_ok: Vec<Id> = (0..rows)
3698        .map(|i| {
3699            let row = b.c_u32(i);
3700            b.ult(row, m)
3701        })
3702        .collect();
3703    let lhs_rows: Vec<Id> = (0..rows)
3704        .map(|i| {
3705            let row = b.c_u32(i);
3706            let offset = b.imul(row, k);
3707            b.iadd(lhs_batch, offset)
3708        })
3709        .collect();
3710    // Element index of `rhs[kk, col_base]`, clamped to zero when the word lies past `n`.
3711    let rhs_word_element = |b: &mut Builder, kk: Id| {
3712        let col_ok = b.ult(col_base, n);
3713        let element = b.imul(kk, n);
3714        let element = b.iadd(element, rhs_col);
3715        b.select_u32(col_ok, element, zero)
3716    };
3717    // The inner product over the slice, `load_rhs` yielding the `lanes` weights at `kk`.
3718    let accumulate = |b: &mut Builder, load_rhs: &dyn Fn(&mut Builder, Id) -> Vec<Id>| {
3719        b.store(kk_var, k_begin);
3720        let (scope, kk) = b.begin_loop(kk_var, k_end);
3721        let weights = load_rhs(b, kk);
3722        for i in 0..rows as usize {
3723            b.if_then(row_ok[i], |b| {
3724                let element = b.iadd(lhs_rows[i], kk);
3725                let a = b.load_f32(array, lhs, element);
3726                for (l, weight) in weights.iter().enumerate() {
3727                    let accumulator = accumulators[i * lanes as usize + l];
3728                    let acc = b.load(f32_ty, accumulator);
3729                    let next = b.fma(a, *weight, acc);
3730                    b.store(accumulator, next);
3731                }
3732            });
3733        }
3734        b.end_loop(scope, kk_var, one);
3735    };
3736    if lanes == 1 {
3737        accumulate(&mut b, &|b, kk| {
3738            let element = rhs_word_element(b, kk);
3739            vec![b.load_f32(array, rhs, element)]
3740        });
3741    } else {
3742        let remainder = b.umod(n, lanes_c);
3743        let aligned = b.ieq(remainder, zero);
3744        b.if_then(aligned, |b| {
3745            accumulate(b, &|b, kk| {
3746                let element = rhs_word_element(b, kk);
3747                let bits = b.load_lane_group(rhs_storage, array, rhs, element, lanes);
3748                bits.into_iter()
3749                    .map(|bits| b.widen_bits(rhs_storage, bits))
3750                    .collect()
3751            });
3752        });
3753        let unaligned = b.lnot(aligned);
3754        b.if_then(unaligned, |b| {
3755            accumulate(b, &|b, kk| {
3756                (0..lanes)
3757                    .map(|l| {
3758                        let l = b.c_u32(l);
3759                        let col = b.iadd(col_base, l);
3760                        let col_ok = b.ult(col, n);
3761                        let element = b.imul(kk, n);
3762                        let element = b.iadd(element, rhs_batch);
3763                        let element = b.iadd(element, col);
3764                        let element = b.select_u32(col_ok, element, zero);
3765                        b.load_float(rhs_storage, array, rhs, element)
3766                    })
3767                    .collect()
3768            });
3769        });
3770    }
3771    // Partials to shared memory: invocation `lid`'s block of `rows · lanes` values.
3772    let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3773    let block = b.c_u32(rows * lanes);
3774    let base = b.imul(lid, block);
3775    for (index, accumulator) in accumulators.iter().enumerate() {
3776        let offset = b.c_u32(index as u32);
3777        let slot = b.iadd(base, offset);
3778        let value = b.load(f32_ty, *accumulator);
3779        let pointer = b.access_chain(workgroup_ptr, partials, &[slot]);
3780        b.store(pointer, value);
3781    }
3782    b.workgroup_barrier();
3783    // Reduce: output `o = lid · per_invocation + j` is row `o / STREAM_COLUMNS`, column
3784    // `o % STREAM_COLUMNS`, summed over the slices in ascending order from slice zero.
3785    let per_invocation = rows * STREAM_COLUMNS / STREAM_WORKGROUP;
3786    let per_invocation_c = b.c_u32(per_invocation);
3787    let first = b.imul(lid, per_invocation_c);
3788    let out_batch = b.imul(z, m);
3789    let out_batch = b.imul(out_batch, n);
3790    for j in 0..per_invocation {
3791        let j = b.c_u32(j);
3792        let o = b.iadd(first, j);
3793        let row = b.udiv(o, columns_c);
3794        let col = b.umod(o, columns_c);
3795        let word = b.udiv(col, lanes_c);
3796        let lane = b.umod(col, lanes_c);
3797        let within = b.imul(row, lanes_c);
3798        let within = b.iadd(within, lane);
3799        let mut sum = None;
3800        for split in 0..splits {
3801            let offset = b.c_u32(split * words);
3802            let owner = b.iadd(offset, word);
3803            let slot = b.imul(owner, block);
3804            let slot = b.iadd(slot, within);
3805            let pointer = b.access_chain(workgroup_ptr, partials, &[slot]);
3806            let partial = b.load(f32_ty, pointer);
3807            sum = Some(match sum {
3808                None => partial,
3809                Some(sum) => b.fadd(sum, partial),
3810            });
3811        }
3812        let sum = sum.expect("at least one slice");
3813        let global_col = b.iadd(col0, col);
3814        let row_in = b.ult(row, m);
3815        let col_in = b.ult(global_col, n);
3816        let in_range = b.land(row_in, col_in);
3817        b.if_then(in_range, |b| {
3818            let out = b.imul(row, n);
3819            let out = b.iadd(out, out_batch);
3820            let out = b.iadd(out, global_col);
3821            b.store_float(output_storage, array, output, out, sum);
3822        });
3823    }
3824    b.end_main();
3825    b.finish([STREAM_WORKGROUP, 1, 1])
3826}
3827
3828/// Cooperative-matrix NVFP4 projection for devices that advertise subgroup FP16 8x16x16 with
3829/// FP32 accumulation. One subgroup owns an 8-token by 16-output tile. Packed FP4 weights are
3830/// decoded once into a binary16 workgroup tile, then the driver lowers
3831/// `OpCooperativeMatrixMulAddKHR` to the architecture's matrix engine (XMX on Intel Xe2/Xe3).
3832fn assemble_nvfp4_matmul_cooperative(buffers: u32) -> Vec<u32> {
3833    let mut b = Builder::new();
3834    b.enable_cooperative_matrix();
3835    let array = b.buffer_array(buffers);
3836    let activation = b.spec_operand();
3837    let packed = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
3838    let block_scales = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
3839    let tensor_scale = b.spec_operand();
3840    let output = b.spec_operand();
3841    let m = b.spec_u32(1);
3842    let n = b.spec_u32(1);
3843    let k = b.spec_u32(16);
3844    let epilogue = b.spec_u32(0);
3845    let weight_mode = b.spec_u32(0);
3846
3847    let a_tile = b.shared_f16_array(8 * 16);
3848    let b_tile = b.shared_f16_array(16 * 16);
3849    let c_tile = b.shared_f32_array(8 * 16);
3850    let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
3851    let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
3852
3853    b.begin_main();
3854    let u32_ty = b.u32_ty();
3855    let f32_ty = b.f32_ty();
3856    let f16_ty = b.f16_ty();
3857    let matrix_a_ty = b.cooperative_matrix_ty(f16_ty, 8, 16, 0);
3858    let matrix_b_ty = b.cooperative_matrix_ty(f16_ty, 16, 16, 1);
3859    let matrix_c_ty = b.cooperative_matrix_ty(f32_ty, 8, 16, 2);
3860    let accumulator = b.local(matrix_c_ty);
3861    let block_var = b.local(u32_ty);
3862    let fp4 = b.private_u32_array(&[
3863        0.0f32.to_bits(),
3864        0.5f32.to_bits(),
3865        1.0f32.to_bits(),
3866        1.5f32.to_bits(),
3867        2.0f32.to_bits(),
3868        3.0f32.to_bits(),
3869        4.0f32.to_bits(),
3870        6.0f32.to_bits(),
3871        (-0.0f32).to_bits(),
3872        (-0.5f32).to_bits(),
3873        (-1.0f32).to_bits(),
3874        (-1.5f32).to_bits(),
3875        (-2.0f32).to_bits(),
3876        (-3.0f32).to_bits(),
3877        (-4.0f32).to_bits(),
3878        (-6.0f32).to_bits(),
3879    ]);
3880    let lid = b.builtin_component(local_id, 0);
3881    let output_tile = b.builtin_component(group_id, 0);
3882    let token_tile = b.builtin_component(group_id, 1);
3883    let zero = b.c_u32(0);
3884    let one = b.c_u32(1);
3885    let two = b.c_u32(2);
3886    let four = b.c_u32(4);
3887    let eight = b.c_u32(8);
3888    let sixteen = b.c_u32(16);
3889    let zero_f = b.c_f32(0.0);
3890    let zero_h = b.f32_to_f16(zero_f);
3891    let blocks = b.udiv(k, sixteen);
3892    let output_base = b.imul(output_tile, sixteen);
3893    let token_base = b.imul(token_tile, eight);
3894    // Cooperative lowering is selected only for shared weights. Keep the specialization operand
3895    // live so an artifact compiled with the wrong mode cannot silently use this shader.
3896    let shared_mode = b.ieq(weight_mode, zero);
3897    let row_bytes = b.udiv(k, two);
3898
3899    let f32_workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
3900    let f16_workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f16_ty);
3901    // Initialize the accumulator tile once. Each lane owns four consecutive-stride elements.
3902    for wave in 0..4 {
3903        let offset = b.c_u32(wave * 32);
3904        let element = b.iadd(lid, offset);
3905        let pointer = b.access_chain(f32_workgroup_ptr, c_tile, &[element]);
3906        b.store(pointer, zero_f);
3907    }
3908    b.workgroup_barrier();
3909    let c_pointer = b.access_chain(f32_workgroup_ptr, c_tile, &[zero]);
3910    let initial = b.cooperative_load(matrix_c_ty, c_pointer, sixteen);
3911    b.store(accumulator, initial);
3912
3913    b.store(block_var, zero);
3914    let (scope, block) = b.begin_loop(block_var, blocks);
3915    let activation_block = b.imul(block, sixteen);
3916    // A is a real 8x16 token tile. The final partial tile is zero padded, allowing every native
3917    // cooperative-matrix row to perform useful prompt work whenever eight tokens remain.
3918    for wave in 0..4 {
3919        let offset = b.c_u32(wave * 32);
3920        let element = b.iadd(lid, offset);
3921        let tile_row = b.udiv(element, sixteen);
3922        let tile_column = b.umod(element, sixteen);
3923        let token = b.iadd(token_base, tile_row);
3924        let token_in_range = b.ult(token, m);
3925        let activation_row = b.imul(token, k);
3926        let source = b.iadd(activation_row, activation_block);
3927        let source = b.iadd(source, tile_column);
3928        let source = b.select_u32(token_in_range, source, zero);
3929        let value = b.load_f32(array, activation, source);
3930        let value = b.select_f32(token_in_range, value, zero_f);
3931        let value = b.f32_to_f16(value);
3932        let pointer = b.access_chain(f16_workgroup_ptr, a_tile, &[element]);
3933        b.store(pointer, value);
3934    }
3935
3936    // B is KxN in row-major cooperative-matrix order, while the checkpoint is N rows of packed
3937    // K. Decode and transpose one 16x16 tile cooperatively.
3938    for wave in 0..8 {
3939        let offset = b.c_u32(wave * 32);
3940        let element = b.iadd(lid, offset);
3941        let k_lane = b.udiv(element, sixteen);
3942        let column = b.umod(element, sixteen);
3943        let output_row = b.iadd(output_base, column);
3944        let in_range = b.ult(output_row, n);
3945        let pointer = b.access_chain(f16_workgroup_ptr, b_tile, &[element]);
3946        b.store(pointer, zero_h);
3947        b.if_then(in_range, |b| {
3948            let weight_row = output_row;
3949            let packed_row = b.imul(weight_row, row_bytes);
3950            let packed_block = b.imul(block, eight);
3951            let packed_base = b.iadd(packed_row, packed_block);
3952            let packed_lane = b.udiv(k_lane, two);
3953            let packed_index = b.iadd(packed_base, packed_lane);
3954            let codes = b.load_byte_bits(array, packed[0], packed_index);
3955            let nibble_mask = b.c_u32(15);
3956            let low = b.band(codes, nibble_mask);
3957            let high = b.shr(codes, four);
3958            let parity = b.band(k_lane, one);
3959            let odd = b.ine(parity, zero);
3960            let code = b.select_u32(odd, high, low);
3961            let private_u32 = b.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
3962            let fp4_pointer = b.access_chain(private_u32, fp4, &[code]);
3963            let weight_bits = b.load(u32_ty, fp4_pointer);
3964            let weight = b.bitcast_f32(weight_bits);
3965            let scale_row = b.imul(weight_row, blocks);
3966            let scale_index = b.iadd(scale_row, block);
3967            let scale_bits = b.load_byte_bits(array, block_scales[0], scale_index);
3968            let scale = b.widen_fp8(Fp8Format::E4M3, scale_bits);
3969            let weight = b.fmul(weight, scale);
3970            let weight = b.f32_to_f16(weight);
3971            b.store(pointer, weight);
3972        });
3973    }
3974    b.workgroup_barrier();
3975    let a_pointer = b.access_chain(f16_workgroup_ptr, a_tile, &[zero]);
3976    let b_pointer = b.access_chain(f16_workgroup_ptr, b_tile, &[zero]);
3977    let matrix_a = b.cooperative_load(matrix_a_ty, a_pointer, sixteen);
3978    let matrix_b = b.cooperative_load(matrix_b_ty, b_pointer, sixteen);
3979    let matrix_c = b.load(matrix_c_ty, accumulator);
3980    let matrix_c = b.cooperative_mul_add(matrix_c_ty, matrix_a, matrix_b, matrix_c);
3981    b.store(accumulator, matrix_c);
3982    b.workgroup_barrier();
3983    b.end_loop(scope, block_var, one);
3984
3985    let matrix_c = b.load(matrix_c_ty, accumulator);
3986    b.cooperative_store(c_pointer, matrix_c, sixteen);
3987    b.workgroup_barrier();
3988    for wave in 0..4 {
3989        let offset = b.c_u32(wave * 32);
3990        let element = b.iadd(lid, offset);
3991        let tile_row = b.udiv(element, sixteen);
3992        let tile_column = b.umod(element, sixteen);
3993        let token = b.iadd(token_base, tile_row);
3994        let output_row = b.iadd(output_base, tile_column);
3995        let token_in_range = b.ult(token, m);
3996        let row_in_tensor = b.ult(output_row, n);
3997        let write = b.land(token_in_range, row_in_tensor);
3998        let write = b.land(write, shared_mode);
3999        b.if_then(write, |b| {
4000            let pointer = b.access_chain(f32_workgroup_ptr, c_tile, &[element]);
4001            let sum = b.load(f32_ty, pointer);
4002            let scale = b.load_f32(array, tensor_scale, zero);
4003            let sum = b.fmul(sum, scale);
4004            let negated = b.fneg(sum);
4005            let exp = b.ext_f32(GLSL_EXP, &[negated]);
4006            let one_f = b.c_f32(1.0);
4007            let denominator = b.fadd(one_f, exp);
4008            let sigmoid = b.fdiv(one_f, denominator);
4009            let silu = b.fmul(sum, sigmoid);
4010            let is_silu = b.ieq(epilogue, one);
4011            let is_sigmoid = b.ieq(epilogue, two);
4012            let three = b.c_u32(3);
4013            let zero_f = b.c_f32(0.0);
4014            let squared_input = sum;
4015            let positive = b.fogt(squared_input, zero_f);
4016            let rectified = b.select_f32(positive, squared_input, zero_f);
4017            let squared = b.fmul(rectified, rectified);
4018            let is_squared = b.ieq(epilogue, three);
4019            let sum = b.select_f32(is_silu, silu, sum);
4020            let sum = b.select_f32(is_sigmoid, sigmoid, sum);
4021            let sum = b.select_f32(is_squared, squared, sum);
4022            let destination = b.imul(token, n);
4023            let destination = b.iadd(destination, output_row);
4024            b.store_f32(array, output, destination, sum);
4025        });
4026    }
4027    b.end_main();
4028    b.finish([32, 1, 1])
4029}
4030
4031/// Subgroup-reduced scalar NVFP4 projection. A fixed 32-lane subgroup folds four output rows,
4032/// reusing every activation load across them and cutting dispatch count by four. Devices without
4033/// subgroup arithmetic retain the portable 64-lane shared-memory reduction below.
4034fn assemble_nvfp4_matmul_subgroup(buffers: u32) -> Vec<u32> {
4035    let mut b = Builder::new();
4036    b.enable_subgroup_arithmetic();
4037    let array = b.buffer_array(buffers);
4038    let activation = b.spec_operand();
4039    let packed = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4040    let block_scales = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4041    let tensor_scale = b.spec_operand();
4042    let output = b.spec_operand();
4043    let _m = b.spec_u32(1);
4044    let n = b.spec_u32(1);
4045    let k = b.spec_u32(16);
4046    let epilogue = b.spec_u32(0);
4047    let weight_mode = b.spec_u32(0);
4048    let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
4049    let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
4050
4051    b.begin_main();
4052    let u32_ty = b.u32_ty();
4053    let f32_ty = b.f32_ty();
4054    let block_var = b.local(u32_ty);
4055    let accumulators = core::array::from_fn::<_, 4, _>(|_| b.local(f32_ty));
4056    let fp4 = b.private_u32_array(&[
4057        0.0f32.to_bits(),
4058        0.5f32.to_bits(),
4059        1.0f32.to_bits(),
4060        1.5f32.to_bits(),
4061        2.0f32.to_bits(),
4062        3.0f32.to_bits(),
4063        4.0f32.to_bits(),
4064        6.0f32.to_bits(),
4065        (-0.0f32).to_bits(),
4066        (-0.5f32).to_bits(),
4067        (-1.0f32).to_bits(),
4068        (-1.5f32).to_bits(),
4069        (-2.0f32).to_bits(),
4070        (-3.0f32).to_bits(),
4071        (-4.0f32).to_bits(),
4072        (-6.0f32).to_bits(),
4073    ]);
4074    let lid = b.builtin_component(local_id, 0);
4075    let output_group = b.builtin_component(group_id, 0);
4076    let token = b.builtin_component(group_id, 1);
4077    let zero = b.c_u32(0);
4078    let zero_f = b.c_f32(0.0);
4079    let four = b.c_u32(4);
4080    let eight = b.c_u32(8);
4081    let sixteen = b.c_u32(16);
4082    let thirty_two = b.c_u32(32);
4083    let nibble_mask = b.c_u32(15);
4084    let blocks = b.udiv(k, sixteen);
4085    let output_base = b.imul(output_group, four);
4086    let mut packed_operand = packed[0];
4087    let mut scale_operand = block_scales[0];
4088    for index in 1..6 {
4089        let index_id = b.c_u32(index as u32);
4090        let selected = b.ieq(token, index_id);
4091        packed_operand.0 = b.select_u32(selected, packed[index].0, packed_operand.0);
4092        packed_operand.1 = b.select_u32(selected, packed[index].1, packed_operand.1);
4093        scale_operand.0 = b.select_u32(selected, block_scales[index].0, scale_operand.0);
4094        scale_operand.1 = b.select_u32(selected, block_scales[index].1, scale_operand.1);
4095    }
4096    let contiguous_mode = b.c_u32(1);
4097    let contiguous = b.ieq(weight_mode, contiguous_mode);
4098    let weight_batch = b.select_u32(contiguous, token, zero);
4099    let batch_rows = b.imul(weight_batch, n);
4100    let two = b.c_u32(2);
4101    let row_bytes = b.udiv(k, two);
4102    let activation_row = b.imul(token, k);
4103    let output_rows = core::array::from_fn::<_, 4, _>(|index| {
4104        let offset = b.c_u32(index as u32);
4105        b.iadd(output_base, offset)
4106    });
4107    let row_valid = output_rows.map(|row| b.ult(row, n));
4108    let weight_rows = core::array::from_fn::<_, 4, _>(|index| {
4109        let safe_row = b.select_u32(row_valid[index], output_rows[index], zero);
4110        b.iadd(batch_rows, safe_row)
4111    });
4112    let row_blocks = weight_rows.map(|row| b.imul(row, blocks));
4113    let packed_rows = weight_rows.map(|row| b.imul(row, row_bytes));
4114    for accumulator in accumulators {
4115        b.store(accumulator, zero_f);
4116    }
4117    b.store(block_var, lid);
4118    let (scope, block) = b.begin_loop(block_var, blocks);
4119    let scales = core::array::from_fn::<_, 4, _>(|index| {
4120        let scale_element = b.iadd(row_blocks[index], block);
4121        let scale_bits = b.load_byte_bits(array, scale_operand, scale_element);
4122        b.widen_fp8(Fp8Format::E4M3, scale_bits)
4123    });
4124    let packed_block = b.imul(block, eight);
4125    let packed_bases = packed_rows.map(|row| b.iadd(row, packed_block));
4126    let activation_block = b.imul(block, sixteen);
4127    let activation_base = b.iadd(activation_row, activation_block);
4128    for byte in 0..8 {
4129        let byte_offset = b.c_u32(byte);
4130        let codes = core::array::from_fn::<_, 4, _>(|index| {
4131            let byte_index = b.iadd(packed_bases[index], byte_offset);
4132            b.load_byte_bits(array, packed_operand, byte_index)
4133        });
4134        for lane in 0..2 {
4135            let inner = b.c_u32(byte * 2 + lane as u32);
4136            let element = b.iadd(activation_base, inner);
4137            let value = b.load_f32(array, activation, element);
4138            for index in 0..4 {
4139                let code = if lane == 0 {
4140                    b.band(codes[index], nibble_mask)
4141                } else {
4142                    b.shr(codes[index], four)
4143                };
4144                let pointer_ty = b.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
4145                let pointer = b.access_chain(pointer_ty, fp4, &[code]);
4146                let weight_bits = b.load(u32_ty, pointer);
4147                let weight = b.bitcast_f32(weight_bits);
4148                let weight = b.fmul(weight, scales[index]);
4149                let acc = b.load(f32_ty, accumulators[index]);
4150                let next = b.fma(value, weight, acc);
4151                b.store(accumulators[index], next);
4152            }
4153        }
4154    }
4155    b.end_loop(scope, block_var, thirty_two);
4156
4157    let sums = accumulators.map(|accumulator| {
4158        let partial = b.load(f32_ty, accumulator);
4159        b.subgroup_sum_f32(partial)
4160    });
4161
4162    let first = b.ieq(lid, zero);
4163    b.if_then(first, |b| {
4164        let shared_mode = b.c_u32(0);
4165        let batched = b.ine(weight_mode, shared_mode);
4166        let tensor_scale_index = b.select_u32(batched, token, zero);
4167        let tensor_scale = b.load_f32(array, tensor_scale, tensor_scale_index);
4168        for index in 0..4 {
4169            b.if_then(row_valid[index], |b| {
4170                let sum = b.fmul(sums[index], tensor_scale);
4171                let negated = b.fneg(sum);
4172                let exp = b.ext_f32(GLSL_EXP, &[negated]);
4173                let one_f = b.c_f32(1.0);
4174                let denominator = b.fadd(one_f, exp);
4175                let sigmoid = b.fdiv(one_f, denominator);
4176                let silu = b.fmul(sum, sigmoid);
4177                let silu_mode = b.c_u32(1);
4178                let sigmoid_mode = b.c_u32(2);
4179                let squared_mode = b.c_u32(3);
4180                let is_silu = b.ieq(epilogue, silu_mode);
4181                let is_sigmoid = b.ieq(epilogue, sigmoid_mode);
4182                let zero_f = b.c_f32(0.0);
4183                let squared_input = sum;
4184                let positive = b.fogt(squared_input, zero_f);
4185                let rectified = b.select_f32(positive, squared_input, zero_f);
4186                let squared = b.fmul(rectified, rectified);
4187                let is_squared = b.ieq(epilogue, squared_mode);
4188                let sum = b.select_f32(is_silu, silu, sum);
4189                let sum = b.select_f32(is_sigmoid, sigmoid, sum);
4190                let sum = b.select_f32(is_squared, squared, sum);
4191                let element = b.imul(token, n);
4192                let element = b.iadd(element, output_rows[index]);
4193                b.store_f32(array, output, element, sum);
4194            });
4195        }
4196    });
4197    b.end_main();
4198    b.finish([32, 1, 1])
4199}
4200
4201/// Scalar baseline for native NVFP4 projections. One workgroup owns one
4202/// `[token, output-row]` dot product. Its 64 invocations divide whole
4203/// 16-weight scale blocks, so every packed byte and scale is read exactly once
4204/// and only 64 partial sums cross workgroup memory. The same artifact can later
4205/// select a cooperative-matrix/XMX kernel without changing its storage ABI.
4206fn assemble_nvfp4_matmul(buffers: u32) -> Vec<u32> {
4207    let mut b = Builder::new();
4208    let array = b.buffer_array(buffers);
4209    let activation = b.spec_operand();
4210    let packed = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4211    let block_scales = core::array::from_fn::<_, 6, _>(|_| b.spec_operand());
4212    let tensor_scale = b.spec_operand();
4213    let output = b.spec_operand();
4214    let _m = b.spec_u32(1);
4215    let n = b.spec_u32(1);
4216    let k = b.spec_u32(16);
4217    let epilogue = b.spec_u32(0);
4218    let weight_mode = b.spec_u32(0);
4219    let partials = b.shared_f32_array(STREAM_WORKGROUP);
4220    let local_id = b.builtin_uvec3(BUILT_IN_LOCAL_INVOCATION_ID);
4221    let group_id = b.builtin_uvec3(BUILT_IN_WORKGROUP_ID);
4222
4223    b.begin_main();
4224    let u32_ty = b.u32_ty();
4225    let f32_ty = b.f32_ty();
4226    let block_var = b.local(u32_ty);
4227    let accumulator = b.local(f32_ty);
4228    let fp4 = b.private_u32_array(&[
4229        0.0f32.to_bits(),
4230        0.5f32.to_bits(),
4231        1.0f32.to_bits(),
4232        1.5f32.to_bits(),
4233        2.0f32.to_bits(),
4234        3.0f32.to_bits(),
4235        4.0f32.to_bits(),
4236        6.0f32.to_bits(),
4237        (-0.0f32).to_bits(),
4238        (-0.5f32).to_bits(),
4239        (-1.0f32).to_bits(),
4240        (-1.5f32).to_bits(),
4241        (-2.0f32).to_bits(),
4242        (-3.0f32).to_bits(),
4243        (-4.0f32).to_bits(),
4244        (-6.0f32).to_bits(),
4245    ]);
4246    let lid = b.builtin_component(local_id, 0);
4247    let out_row = b.builtin_component(group_id, 0);
4248    let token = b.builtin_component(group_id, 1);
4249    let zero = b.c_u32(0);
4250    let zero_f = b.c_f32(0.0);
4251    let one = b.c_u32(1);
4252    let four = b.c_u32(4);
4253    let eight = b.c_u32(8);
4254    let sixteen = b.c_u32(16);
4255    let sixty_four = b.c_u32(STREAM_WORKGROUP);
4256    let nibble_mask = b.c_u32(15);
4257    let blocks = b.udiv(k, sixteen);
4258    let mut packed_operand = packed[0];
4259    let mut scale_operand = block_scales[0];
4260    for index in 1..6 {
4261        let index_id = b.c_u32(index as u32);
4262        let selected = b.ieq(token, index_id);
4263        packed_operand.0 = b.select_u32(selected, packed[index].0, packed_operand.0);
4264        packed_operand.1 = b.select_u32(selected, packed[index].1, packed_operand.1);
4265        scale_operand.0 = b.select_u32(selected, block_scales[index].0, scale_operand.0);
4266        scale_operand.1 = b.select_u32(selected, block_scales[index].1, scale_operand.1);
4267    }
4268    let contiguous_mode = b.c_u32(1);
4269    let contiguous = b.ieq(weight_mode, contiguous_mode);
4270    let weight_batch = b.select_u32(contiguous, token, zero);
4271    let batch_rows = b.imul(weight_batch, n);
4272    let weight_row = b.iadd(batch_rows, out_row);
4273    let row_blocks = b.imul(weight_row, blocks);
4274    let two = b.c_u32(2);
4275    let row_bytes = b.udiv(k, two);
4276    let packed_row = b.imul(weight_row, row_bytes);
4277    let activation_row = b.imul(token, k);
4278    b.store(accumulator, zero_f);
4279    b.store(block_var, lid);
4280    let (scope, block) = b.begin_loop(block_var, blocks);
4281    let scale_element = b.iadd(row_blocks, block);
4282    let scale_bits = b.load_byte_bits(array, scale_operand, scale_element);
4283    let scale = b.widen_fp8(Fp8Format::E4M3, scale_bits);
4284    let packed_block = b.imul(block, eight);
4285    let packed_base = b.iadd(packed_row, packed_block);
4286    let activation_block = b.imul(block, sixteen);
4287    let activation_base = b.iadd(activation_row, activation_block);
4288    for byte in 0..8 {
4289        let byte_offset = b.c_u32(byte);
4290        let byte_index = b.iadd(packed_base, byte_offset);
4291        let codes = b.load_byte_bits(array, packed_operand, byte_index);
4292        let low = b.band(codes, nibble_mask);
4293        let high = b.shr(codes, four);
4294        for (lane, code) in [low, high].into_iter().enumerate() {
4295            let pointer_ty = b.pointer(STORAGE_CLASS_PRIVATE, u32_ty);
4296            let pointer = b.access_chain(pointer_ty, fp4, &[code]);
4297            let weight_bits = b.load(u32_ty, pointer);
4298            let weight = b.bitcast_f32(weight_bits);
4299            let weight = b.fmul(weight, scale);
4300            let inner = b.c_u32(byte * 2 + lane as u32);
4301            let element = b.iadd(activation_base, inner);
4302            let value = b.load_f32(array, activation, element);
4303            let acc = b.load(f32_ty, accumulator);
4304            let next = b.fma(value, weight, acc);
4305            b.store(accumulator, next);
4306        }
4307    }
4308    b.end_loop(scope, block_var, sixty_four);
4309
4310    let workgroup_ptr = b.pointer(STORAGE_CLASS_WORKGROUP, f32_ty);
4311    let partial_ptr = b.access_chain(workgroup_ptr, partials, &[lid]);
4312    let partial = b.load(f32_ty, accumulator);
4313    b.store(partial_ptr, partial);
4314    b.workgroup_barrier();
4315
4316    let first = b.ieq(lid, zero);
4317    b.if_then(first, |b| {
4318        let sum_var = b.local(f32_ty);
4319        let index_var = b.local(u32_ty);
4320        b.store(sum_var, zero_f);
4321        b.store(index_var, zero);
4322        let (sum_scope, index) = b.begin_loop(index_var, sixty_four);
4323        let pointer = b.access_chain(workgroup_ptr, partials, &[index]);
4324        let value = b.load(f32_ty, pointer);
4325        let sum = b.load(f32_ty, sum_var);
4326        let next = b.fadd(sum, value);
4327        b.store(sum_var, next);
4328        b.end_loop(sum_scope, index_var, one);
4329        let sum = b.load(f32_ty, sum_var);
4330        let shared_mode = b.c_u32(0);
4331        let batched = b.ine(weight_mode, shared_mode);
4332        let tensor_scale_index = b.select_u32(batched, token, zero);
4333        let tensor_scale = b.load_f32(array, tensor_scale, tensor_scale_index);
4334        let sum = b.fmul(sum, tensor_scale);
4335        let negated = b.fneg(sum);
4336        let exp = b.ext_f32(GLSL_EXP, &[negated]);
4337        let one_f = b.c_f32(1.0);
4338        let denominator = b.fadd(one_f, exp);
4339        let sigmoid = b.fdiv(one_f, denominator);
4340        let silu = b.fmul(sum, sigmoid);
4341        let silu_mode = b.c_u32(1);
4342        let sigmoid_mode = b.c_u32(2);
4343        let squared_mode = b.c_u32(3);
4344        let is_silu = b.ieq(epilogue, silu_mode);
4345        let is_sigmoid = b.ieq(epilogue, sigmoid_mode);
4346        let zero_f = b.c_f32(0.0);
4347        let squared_input = sum;
4348        let positive = b.fogt(squared_input, zero_f);
4349        let rectified = b.select_f32(positive, squared_input, zero_f);
4350        let squared = b.fmul(rectified, rectified);
4351        let is_squared = b.ieq(epilogue, squared_mode);
4352        let sum = b.select_f32(is_silu, silu, sum);
4353        let sum = b.select_f32(is_sigmoid, sigmoid, sum);
4354        let sum = b.select_f32(is_squared, squared, sum);
4355        let element = b.imul(token, n);
4356        let element = b.iadd(element, out_row);
4357        b.store_f32(array, output, element, sum);
4358    });
4359    b.end_main();
4360    b.finish([STREAM_WORKGROUP, 1, 1])
4361}
4362
4363/// NHWC MAX_POOL2D at `float` storage: one invocation per output element folds its window with
4364/// `apply_max_s` from `-inf`; padded positions are skipped, never substituted. FP16 inputs are
4365/// widened for the fold — an exact selection, so the narrowed store is exact.
4366///
4367/// Specialization order: input, output `(buffer, base)`; `batch`, `height`, `width`, `channels`,
4368/// `out_height`, `out_width`, `kernel_h`, `kernel_w`, `stride_h`, `stride_w`, `pad_top`,
4369/// `pad_left`.
4370fn assemble_max_pool(nan_mode: NanMode, float: Storage, workgroup: u32, buffers: u32) -> Vec<u32> {
4371    let mut b = Builder::new();
4372    let array = b.buffer_array(buffers);
4373    let input = b.spec_operand();
4374    let output = b.spec_operand();
4375    let batch = b.spec_u32(1);
4376    let height = b.spec_u32(1);
4377    let width = b.spec_u32(1);
4378    let channels = b.spec_u32(1);
4379    let out_height = b.spec_u32(1);
4380    let out_width = b.spec_u32(1);
4381    let kernel_h = b.spec_u32(1);
4382    let kernel_w = b.spec_u32(1);
4383    let stride_h = b.spec_u32(1);
4384    let stride_w = b.spec_u32(1);
4385    let pad_top = b.spec_u32(0);
4386    let pad_left = b.spec_u32(0);
4387
4388    let (counter, stride) = b.grid_stride(workgroup);
4389    let u32_ty = b.u32_ty();
4390    let f32_ty = b.f32_ty();
4391    let acc_var = b.local(f32_ty);
4392    let kh_var = b.local(u32_ty);
4393    let kw_var = b.local(u32_ty);
4394    let count = b.imul(batch, out_height);
4395    let count = b.imul(count, out_width);
4396    let count = b.imul(count, channels);
4397    let zero = b.c_u32(0);
4398    let one = b.c_u32(1);
4399    let (scope, o) = b.begin_loop(counter, count);
4400    let c = b.umod(o, channels);
4401    let t = b.udiv(o, channels);
4402    let ow = b.umod(t, out_width);
4403    let t = b.udiv(t, out_width);
4404    let oh = b.umod(t, out_height);
4405    let nb = b.udiv(t, out_height);
4406    let neg_inf = b.c_f32(f32::NEG_INFINITY);
4407    b.store(acc_var, neg_inf);
4408    b.store(kh_var, zero);
4409    let row_origin = b.imul(oh, stride_h);
4410    let col_origin = b.imul(ow, stride_w);
4411    let (rows, kh) = b.begin_loop(kh_var, kernel_h);
4412    let padded_row = b.iadd(row_origin, kh);
4413    let row_in_low = b.uge(padded_row, pad_top);
4414    let ih = b.isub(padded_row, pad_top);
4415    let row_in_high = b.ult(ih, height);
4416    let row_ok = b.land(row_in_low, row_in_high);
4417    b.store(kw_var, zero);
4418    let (cols, kw) = b.begin_loop(kw_var, kernel_w);
4419    let padded_col = b.iadd(col_origin, kw);
4420    let col_in_low = b.uge(padded_col, pad_left);
4421    let iw = b.isub(padded_col, pad_left);
4422    let col_in_high = b.ult(iw, width);
4423    let col_ok = b.land(col_in_low, col_in_high);
4424    let ok = b.land(row_ok, col_ok);
4425    // ((n * H + ih) * W + iw) * C + c, clamped to element zero when the tap is padding.
4426    let index = b.imul(nb, height);
4427    let index = b.iadd(index, ih);
4428    let index = b.imul(index, width);
4429    let index = b.iadd(index, iw);
4430    let index = b.imul(index, channels);
4431    let index = b.iadd(index, c);
4432    let index = b.select_u32(ok, index, zero);
4433    let value = b.load_float(float, array, input, index);
4434    let acc = b.load(f32_ty, acc_var);
4435    let folded = b.apply_max(acc, value, nan_mode);
4436    let next = b.select_f32(ok, folded, acc);
4437    b.store(acc_var, next);
4438    b.end_loop(cols, kw_var, one);
4439    b.end_loop(rows, kh_var, one);
4440    let acc = b.load(f32_ty, acc_var);
4441    b.store_float(float, array, output, o, acc);
4442    b.end_loop(scope, counter, stride);
4443    b.end_main();
4444    b.finish([workgroup, 1, 1])
4445}
4446
4447/// Elementwise float conversion over `count` elements: load at `input` storage as binary32,
4448/// store at `output` storage. Every narrowing is crate-owned integer code, so a `CAST` produces
4449/// the same bits on every device. Sub-word outputs are written a whole word per invocation
4450/// ([`assemble_contiguous_lanes`]).
4451///
4452/// Specialization order: input, output `(buffer, base)`; `count`.
4453fn assemble_cast(
4454    input: Storage,
4455    output_storage: Storage,
4456    workgroup: u32,
4457    buffers: u32,
4458) -> Vec<u32> {
4459    assemble_contiguous_lanes(input, output_storage, workgroup, buffers, |b, bits| {
4460        let value = b.widen_bits(input, bits);
4461        b.narrow_bits(output_storage, value)
4462    })
4463}
4464
4465/// A contiguous lane kernel: `out[i] = convert(in[i])` over `count` elements, where `convert`
4466/// maps one element's raw source bits to its raw destination bits.
4467///
4468/// Word-storage outputs run one element per invocation. Sub-word outputs (`Half`, `Byte`, FP8
4469/// `Quarter`) run one *destination word* per invocation: the invocation loads the source lanes
4470/// of that word's elements (each source word once), converts them, packs them, and stores the
4471/// word with a plain `OpStore`. That is the difference between a copy and a read-modify-write:
4472/// the per-element path clears and sets every lane through two device-scope atomics on a word
4473/// three other invocations are also atomically updating, which on Intel Xe3 ran FP8 `IDENTITY`
4474/// at 6 GB/s against 109 GB/s for FP32. Only a tensor's final partial word — whose remaining
4475/// bytes are not the tensor's to write — still goes through the neighbour-safe atomic sequence,
4476/// lane by lane, so the documented guarantee holds unchanged.
4477///
4478/// Specialization order: input, output `(buffer, base)`; `count` (elements).
4479fn assemble_contiguous_lanes(
4480    input: Storage,
4481    output: Storage,
4482    workgroup: u32,
4483    buffers: u32,
4484    convert: impl Fn(&mut Builder, Id) -> Id,
4485) -> Vec<u32> {
4486    let mut b = Builder::new();
4487    let array = b.buffer_array(buffers);
4488    let source = b.spec_operand();
4489    let destination = b.spec_operand();
4490    let count = b.spec_u32(1);
4491    let (counter, stride) = b.grid_stride(workgroup);
4492    let lanes = output.lanes();
4493    if lanes == 1 {
4494        let (scope, i) = b.begin_loop(counter, count);
4495        let bits = b.load_lane_group(input, array, source, i, 1)[0];
4496        let word = convert(&mut b, bits);
4497        b.store_word(array, destination, i, word);
4498        b.end_loop(scope, counter, stride);
4499    } else {
4500        let lanes_c = b.c_u32(lanes);
4501        let lanes_less_one = b.c_u32(lanes - 1);
4502        let words = b.iadd(count, lanes_less_one);
4503        let words = b.udiv(words, lanes_c);
4504        let (scope, w) = b.begin_loop(counter, words);
4505        let base = b.imul(w, lanes_c);
4506        let sources = b.load_lane_group(input, array, source, base, lanes);
4507        let bits: Vec<Id> = sources
4508            .iter()
4509            .map(|source_bits| convert(&mut b, *source_bits))
4510            .collect();
4511        let packed = b.pack_lanes(&bits);
4512        let end = b.iadd(base, lanes_c);
4513        let full = b.uge(count, end);
4514        b.if_then(full, |b| b.store_word(array, destination, w, packed));
4515        let partial = b.lnot(full);
4516        b.if_then(partial, |b| {
4517            for (lane, bits) in bits.iter().enumerate() {
4518                let lane = b.c_u32(lane as u32);
4519                let element = b.iadd(base, lane);
4520                let in_range = b.ult(element, count);
4521                b.if_then(in_range, |b| {
4522                    b.store_lane_bits(output, array, destination, element, *bits);
4523                });
4524            }
4525        });
4526        b.end_loop(scope, counter, stride);
4527    }
4528    b.end_main();
4529    b.finish([workgroup, 1, 1])
4530}
4531
4532/// Strided copy: `out[out_offset + Σ c_d · out_stride_d] = in[in_offset + Σ c_d · in_stride_d]`
4533/// over the iteration space `dims`; `contiguous` degenerates to `out[i] = in[i]`, which runs
4534/// through [`assemble_contiguous_lanes`] — a whole-word copy for sub-word storage, with `BOOL`
4535/// canonicalized to `0`/`1` lane by lane and FP8/FP16 bit patterns moved untouched.
4536///
4537/// Specialization order: input, output `(buffer, base)`; `count`; then (strided only)
4538/// `dims[MAX_RANK]`, `in_strides[MAX_RANK]`, `in_offset`, `out_strides[MAX_RANK]`, `out_offset`.
4539fn assemble_move(storage: Storage, contiguous: bool, workgroup: u32, buffers: u32) -> Vec<u32> {
4540    if contiguous {
4541        return assemble_contiguous_lanes(storage, storage, workgroup, buffers, |b, bits| {
4542            match storage {
4543                // Canonical `0`/`1`: any nonzero byte is true.
4544                Storage::Byte => {
4545                    let zero = b.c_u32(0);
4546                    let one = b.c_u32(1);
4547                    let set = b.ine(bits, zero);
4548                    b.select_u32(set, one, zero)
4549                }
4550                // Raw lanes: an FP8 or binary16 pattern — NaN payload, subnormal, signed
4551                // zero — moves exactly. The `Byte` arm above must not be reused for these.
4552                Storage::Word | Storage::Half | Storage::Quarter(_) => bits,
4553            }
4554        });
4555    }
4556    let mut b = Builder::new();
4557    let array = b.buffer_array(buffers);
4558    let input = b.spec_operand();
4559    let output = b.spec_operand();
4560    let count = b.spec_u32(1);
4561    let dims = b.spec_dims();
4562    let in_strides = b.spec_strides();
4563    let in_offset = b.spec_u32(0);
4564    let out_strides = b.spec_strides();
4565    let out_offset = b.spec_u32(0);
4566
4567    let (counter, stride) = b.grid_stride(workgroup);
4568    let (scope, i) = b.begin_loop(counter, count);
4569    let indices = b.strided_indices(i, &dims, &[in_strides, out_strides]);
4570    let source = b.iadd(indices[0], in_offset);
4571    let destination = b.iadd(indices[1], out_offset);
4572    match storage {
4573        Storage::Word => {
4574            let word = b.load_word(array, input, source);
4575            b.store_word(array, output, destination, word);
4576        }
4577        Storage::Byte => {
4578            let value = b.load_bool(array, input, source);
4579            b.store_bool(array, output, destination, value);
4580        }
4581        // A raw 8-bit lane copy: FP8 bit patterns — NaNs, subnormals, signed zeros — move
4582        // exactly. The `Byte` arm above must not be reused: it canonicalizes to `0`/`1`.
4583        Storage::Quarter(_) => {
4584            let bits = b.load_byte_bits(array, input, source);
4585            b.store_byte_bits(array, output, destination, bits);
4586        }
4587        // A raw 16-bit lane copy: no float conversion, so every binary16 bit pattern — NaN
4588        // payloads and subnormals included — moves exactly.
4589        Storage::Half => {
4590            let bits = b.load_half_bits(array, input, source);
4591            b.store_half_bits(array, output, destination, bits);
4592        }
4593    }
4594    b.end_loop(scope, counter, stride);
4595    b.end_main();
4596    b.finish([workgroup, 1, 1])
4597}
4598
4599#[cfg(test)]
4600mod tests {
4601    use super::*;
4602
4603    /// Every kernel variant the lowering can select, at one representative tuning.
4604    fn all_keys() -> Vec<KernelKey> {
4605        KernelKey::every_variant()
4606    }
4607
4608    /// Walk a module instruction by instruction: word counts tile the body exactly, every id is
4609    /// below the bound, the section order is legal, and exactly one function exists.
4610    fn check_well_formed(key: KernelKey, words: &[u32]) {
4611        assert_eq!(words[0], SPIRV_MAGIC, "{key:?}");
4612        assert_eq!(words[1], SPIRV_VERSION_1_3, "{key:?}");
4613        let bound = words[3];
4614        let mut cursor = 5;
4615        let mut opcodes = Vec::new();
4616        let mut spec_ids = Vec::new();
4617        let mut defined = std::collections::HashSet::new();
4618        while cursor < words.len() {
4619            let word_count = (words[cursor] >> 16) as usize;
4620            assert!(
4621                word_count >= 1,
4622                "{key:?}: zero-length instruction at {cursor}"
4623            );
4624            assert!(
4625                cursor + word_count <= words.len(),
4626                "{key:?}: instruction overruns"
4627            );
4628            let opcode = (words[cursor] & 0xffff) as u16;
4629            opcodes.push(opcode);
4630            if opcode == OP_DECORATE && words[cursor + 2] == DECORATION_SPEC_ID {
4631                spec_ids.push(words[cursor + 3]);
4632            }
4633            // Result-id-bearing instructions used here: check the bound and uniqueness.
4634            let result_id = match opcode {
4635                OP_TYPE_VOID
4636                | OP_TYPE_BOOL
4637                | OP_TYPE_INT
4638                | OP_TYPE_FLOAT
4639                | OP_TYPE_VECTOR
4640                | OP_TYPE_ARRAY
4641                | OP_TYPE_RUNTIME_ARRAY
4642                | OP_TYPE_STRUCT
4643                | OP_TYPE_POINTER
4644                | OP_TYPE_FUNCTION
4645                | OP_LABEL
4646                | OP_EXT_INST_IMPORT => Some(words[cursor + 1]),
4647                OP_CONSTANT | OP_SPEC_CONSTANT | OP_CONSTANT_FALSE | OP_VARIABLE | OP_LOAD
4648                | OP_ACCESS_CHAIN | OP_FUNCTION | OP_EXT_INST | OP_SELECT | OP_BITCAST => {
4649                    Some(words[cursor + 2])
4650                }
4651                _ => None,
4652            };
4653            if let Some(id) = result_id {
4654                assert!(id < bound, "{key:?}: id {id} exceeds bound {bound}");
4655                assert!(defined.insert(id), "{key:?}: id {id} defined twice");
4656            }
4657            cursor += word_count;
4658        }
4659        assert_eq!(cursor, words.len(), "{key:?}");
4660        assert_eq!(opcodes[0], OP_CAPABILITY, "{key:?}");
4661        let import = opcodes
4662            .iter()
4663            .position(|op| *op == OP_EXT_INST_IMPORT)
4664            .unwrap();
4665        let memory = opcodes
4666            .iter()
4667            .position(|op| *op == OP_MEMORY_MODEL)
4668            .unwrap();
4669        let entry = opcodes.iter().position(|op| *op == OP_ENTRY_POINT).unwrap();
4670        let mode = opcodes
4671            .iter()
4672            .position(|op| *op == OP_EXECUTION_MODE)
4673            .unwrap();
4674        assert!(import < memory && memory < entry && entry < mode, "{key:?}");
4675        assert_eq!(*opcodes.last().unwrap(), OP_FUNCTION_END, "{key:?}");
4676        assert_eq!(
4677            opcodes.iter().filter(|op| **op == OP_FUNCTION).count(),
4678            1,
4679            "{key:?}"
4680        );
4681        // Annotations precede every type declaration.
4682        let last_decorate = opcodes
4683            .iter()
4684            .rposition(|op| *op == OP_DECORATE || *op == OP_MEMBER_DECORATE)
4685            .unwrap();
4686        let first_type = opcodes.iter().position(|op| *op == OP_TYPE_VOID).unwrap();
4687        assert!(
4688            last_decorate < first_type,
4689            "{key:?}: annotation after types"
4690        );
4691        // Specialization ids are exactly 0..count.
4692        spec_ids.sort_unstable();
4693        let expected: Vec<u32> = (0..key.spec_constant_count()).collect();
4694        assert_eq!(spec_ids, expected, "{key:?}: specialization ids");
4695        // Every function-local variable sits at the top of the entry block.
4696        let function_at = opcodes.iter().position(|op| *op == OP_FUNCTION).unwrap();
4697        let body = &opcodes[function_at + 2..];
4698        let first_non_variable = body.iter().position(|op| *op != OP_VARIABLE).unwrap();
4699        assert!(
4700            !body[first_non_variable..].contains(&OP_VARIABLE),
4701            "{key:?}: OpVariable after the entry block prologue"
4702        );
4703    }
4704
4705    #[test]
4706    fn every_kernel_variant_is_well_formed() {
4707        for key in all_keys() {
4708            let words = key.assemble();
4709            check_well_formed(key, &words);
4710        }
4711    }
4712
4713    #[test]
4714    fn assembly_is_deterministic() {
4715        for key in all_keys() {
4716            assert_eq!(key.assemble(), key.assemble(), "{key:?}");
4717        }
4718    }
4719
4720    #[test]
4721    fn spec_encoders_match_declared_counts() {
4722        let operand = Operand { buffer: 1, base: 2 };
4723        let strides = [[1; MAX_RANK]; 3];
4724        for key in all_keys() {
4725            let words = match key {
4726                KernelKey::Elementwise { op, broadcast, .. } => ElementwiseSpec {
4727                    count: 8,
4728                    inputs: &vec![operand; op.inputs().len()],
4729                    output: operand,
4730                    dims: [1; MAX_RANK],
4731                    strides: &strides[..op.inputs().len()],
4732                    clamp: matches!(op, ElementwiseOp::Clamp(_)).then_some([0, 0x3f80_0000]),
4733                }
4734                .words(broadcast),
4735                KernelKey::Reduce { .. } => reduce_spec(operand, operand, 2, 3, 1),
4736                KernelKey::Matmul { .. } | KernelKey::MatmulStream { .. } => {
4737                    matmul_spec(operand, operand, operand, 1, 2, 3, 1)
4738                }
4739                KernelKey::Nvfp4Matmul { .. } => Nvfp4MatmulSpec {
4740                    activation: operand,
4741                    packed: &[operand],
4742                    block_scales: &[operand],
4743                    tensor_scale: operand,
4744                    output: operand,
4745                    m: 1,
4746                    n: 2,
4747                    k: 16,
4748                    epilogue: 0,
4749                    weight_mode: 0,
4750                }
4751                .words(),
4752                KernelKey::MaxPool { .. } => max_pool_spec(
4753                    operand,
4754                    operand,
4755                    PoolGeometry {
4756                        batch: 1,
4757                        height: 4,
4758                        width: 4,
4759                        channels: 2,
4760                        out_height: 2,
4761                        out_width: 2,
4762                        kernel: [2, 2],
4763                        stride: [2, 2],
4764                        pad_top: 0,
4765                        pad_left: 0,
4766                    },
4767                ),
4768                // A CAST's specialization payload is a contiguous move's: two operands and a
4769                // count.
4770                KernelKey::Cast { .. } => move_spec(
4771                    operand,
4772                    operand,
4773                    MoveGeometry {
4774                        count: 8,
4775                        dims: [1; MAX_RANK],
4776                        in_strides: [1; MAX_RANK],
4777                        in_offset: 0,
4778                        out_strides: [1; MAX_RANK],
4779                        out_offset: 0,
4780                    },
4781                    true,
4782                ),
4783                KernelKey::Move { contiguous, .. } => move_spec(
4784                    operand,
4785                    operand,
4786                    MoveGeometry {
4787                        count: 6,
4788                        dims: [1; MAX_RANK],
4789                        in_strides: [1; MAX_RANK],
4790                        in_offset: 0,
4791                        out_strides: [1; MAX_RANK],
4792                        out_offset: 0,
4793                    },
4794                    contiguous,
4795                ),
4796            };
4797            assert_eq!(
4798                words.len() as u32,
4799                key.spec_constant_count(),
4800                "{key:?}: encoder length"
4801            );
4802        }
4803    }
4804
4805    #[test]
4806    fn literal_strings_are_nul_terminated_and_word_padded() {
4807        assert_eq!(literal_string("main"), vec![0x6e69_616d, 0]);
4808        assert_eq!(literal_string("abc"), vec![0x0063_6261]);
4809    }
4810
4811    #[test]
4812    fn linear_workgroups_cover_and_cap() {
4813        assert_eq!(linear_workgroups(1, 64, 65_535), 1);
4814        assert_eq!(linear_workgroups(64, 64, 65_535), 1);
4815        assert_eq!(linear_workgroups(65, 64, 65_535), 2);
4816        assert_eq!(linear_workgroups(u32::MAX, 64, 65_535), 65_535);
4817        assert_eq!(linear_workgroups(0, 64, 65_535), 1);
4818    }
4819
4820    #[test]
4821    fn matmul_workgroups_cover_all_dimensions() {
4822        assert_eq!(matmul_block(16), 64);
4823        assert_eq!(matmul_block(8), 32);
4824        assert_eq!(MatmulGeometry::wide(16).shared_bytes(), 8192);
4825        assert_eq!(stream_matmul_shared_bytes(), 8192);
4826        assert_eq!(matmul_shared_bytes(8), 8192);
4827        for tile in [8, 16] {
4828            let geometry = MatmulGeometry::wide(tile);
4829            assert_eq!(geometry.invocations(), tile * tile, "{geometry:?}");
4830            assert_eq!(
4831                (geometry.block_m() * geometry.depth) % geometry.invocations(),
4832                0,
4833                "{geometry:?}"
4834            );
4835            assert_eq!(
4836                (geometry.block_n() * geometry.depth) % geometry.invocations(),
4837                0,
4838                "{geometry:?}"
4839            );
4840        }
4841        // Every storage's word count divides the workgroup, and the outputs divide evenly.
4842        for lanes in [1, 2, 4] {
4843            let words = STREAM_COLUMNS / lanes;
4844            assert_eq!(STREAM_WORKGROUP % words, 0);
4845            assert_eq!((STREAM_ROWS * STREAM_COLUMNS) % STREAM_WORKGROUP, 0);
4846        }
4847        assert_eq!(stream_matmul_workgroups(4096, 2), [256, 1, 2]);
4848        assert_eq!(stream_matmul_workgroups(17, 1), [2, 1, 1]);
4849        assert_eq!(matmul_workgroups(1, 1, 1, 16), [1, 1, 1]);
4850        assert_eq!(matmul_workgroups(64, 64, 1, 16), [1, 1, 1]);
4851        assert_eq!(matmul_workgroups(65, 64, 1, 16), [1, 2, 1]);
4852        assert_eq!(matmul_workgroups(64, 65, 1, 16), [2, 1, 1]);
4853        assert_eq!(matmul_workgroups(32, 32, 3, 8), [1, 1, 3]);
4854        assert_eq!(matmul_workgroups(33, 8, 3, 8), [1, 2, 3]);
4855    }
4856
4857    #[test]
4858    fn elementwise_storage_tables_are_consistent() {
4859        for key in all_keys() {
4860            if let KernelKey::Elementwise { op, .. } = key {
4861                assert!(!op.inputs().is_empty());
4862                assert!(op.inputs().len() <= MAX_ELEMENTWISE_INPUTS);
4863            }
4864        }
4865    }
4866
4867    /// Independent binary32 → binary16 reference: decompose the (exact) binary64 value into its
4868    /// integer significand and round to the binary16 significand with round-to-nearest-even —
4869    /// deliberately a different formulation than [`f32_to_f16_bits`].
4870    fn reference_f32_to_f16(value: f32) -> u16 {
4871        let value = f64::from(value);
4872        if value.is_nan() {
4873            return if value.is_sign_negative() {
4874                0xfe00
4875            } else {
4876                0x7e00
4877            };
4878        }
4879        let sign: u16 = if value.is_sign_negative() { 0x8000 } else { 0 };
4880        let magnitude = value.abs();
4881        if magnitude >= 65520.0 {
4882            return sign | 0x7c00;
4883        }
4884        if magnitude == 0.0 {
4885            return sign;
4886        }
4887        let bits = magnitude.to_bits();
4888        let biased = ((bits >> 52) & 0x7ff) as i32;
4889        // `f32` inputs are never binary64-subnormal, so the implicit leading bit is present.
4890        let mantissa = (bits & ((1_u64 << 52) - 1)) | (1_u64 << 52);
4891        let exponent = biased - 1023;
4892        // value = mantissa · 2^(exponent - 52)
4893        if exponent >= -14 {
4894            // Binary16 normal: keep the top 11 significand bits, round the rest.
4895            let kept = mantissa >> 42;
4896            let dropped = mantissa & ((1_u64 << 42) - 1);
4897            let half = 1_u64 << 41;
4898            let rounded = if dropped > half || (dropped == half && kept & 1 == 1) {
4899                kept + 1
4900            } else {
4901                kept
4902            };
4903            let (exponent, mantissa) = if rounded == 2048 {
4904                (exponent + 1, 1024)
4905            } else {
4906                (exponent, rounded)
4907            };
4908            let biased = exponent + 15;
4909            if biased >= 31 {
4910                return sign | 0x7c00;
4911            }
4912            return sign | ((biased as u16) << 10) | (mantissa - 1024) as u16;
4913        }
4914        // Binary16 subnormal (step 2^-24): scale into the integer domain and round.
4915        let drop = (28 - exponent) as u32;
4916        if drop >= 64 {
4917            return sign;
4918        }
4919        let kept = mantissa >> drop;
4920        let dropped = mantissa & ((1_u64 << drop) - 1);
4921        let half = 1_u64 << (drop - 1);
4922        let rounded = if dropped > half || (dropped == half && kept & 1 == 1) {
4923            kept + 1
4924        } else {
4925            kept
4926        };
4927        // At most 1024: the smallest normal encoding, correctly reached by the rounding carry.
4928        sign | rounded as u16
4929    }
4930
4931    #[test]
4932    fn binary16_widening_is_exact_for_every_pattern() {
4933        for pattern in 0_u32..=0xffff {
4934            let bits = pattern as u16;
4935            let value = f16_to_f32(bits);
4936            let exponent = (bits >> 10) & 0x1f;
4937            let mantissa = bits & 0x3ff;
4938            match (exponent, mantissa) {
4939                (0, 0) => assert_eq!(value.to_bits(), (u32::from(bits) & 0x8000) << 16),
4940                (31, 0) => assert!(value.is_infinite(), "{bits:#06x}"),
4941                (31, _) => assert!(value.is_nan(), "{bits:#06x}"),
4942                _ => {
4943                    // Round trip: every binary16 value is exactly binary32-representable.
4944                    assert_eq!(f32_to_f16_bits(value), bits, "{bits:#06x}");
4945                }
4946            }
4947        }
4948    }
4949
4950    #[test]
4951    fn binary16_narrowing_matches_the_reference_everywhere() {
4952        // Every binary16 pattern (round trip), the rounding boundaries between them, and a
4953        // pseudo-random sweep across the whole binary32 range.
4954        let mut cases: Vec<f32> = (0_u32..=0xffff)
4955            .map(|pattern| f16_to_f32(pattern as u16))
4956            .collect();
4957        for pattern in 0_u32..=0xfffe {
4958            let lo = f16_to_f32(pattern as u16);
4959            let hi = f16_to_f32((pattern + 1) as u16);
4960            if lo.is_finite() && hi.is_finite() {
4961                cases.push(f32::from_bits(
4962                    lo.to_bits() + hi.to_bits().abs_diff(lo.to_bits()) / 2,
4963                ));
4964            }
4965        }
4966        let mut state = 0x243f_6a88_85a3_08d3_u64;
4967        for _ in 0..1_000_000 {
4968            state = state
4969                .wrapping_mul(6_364_136_223_846_793_005)
4970                .wrapping_add(1);
4971            cases.push(f32::from_bits((state >> 32) as u32));
4972        }
4973        for case in cases {
4974            assert_eq!(
4975                f32_to_f16_bits(case),
4976                reference_f32_to_f16(case),
4977                "{case:e} ({:#010x})",
4978                case.to_bits()
4979            );
4980        }
4981    }
4982}
4983
4984#[cfg(test)]
4985mod fp8_narrowing_tests {
4986    use super::*;
4987    use virtio_accel_tosa::{fp8e4m3_to_f32, fp8e5m2_to_f32};
4988
4989    /// An independent oracle: scan every non-negative encoding for the nearest magnitude, ties
4990    /// to even, then apply the sign separately as IEEE does — so a value that rounds to zero
4991    /// keeps its sign. No bit arithmetic in common with [`f32_to_fp8_bits`].
4992    fn nearest_by_search(format: Fp8Format, value: f32) -> u8 {
4993        let sign = if value.is_sign_negative() { 0x80 } else { 0x00 };
4994        let magnitude = value.abs();
4995        let mut best: Option<(f32, u8)> = None;
4996        for bits in 0..=0x7f_u8 {
4997            let candidate = fp8_decode(format, bits);
4998            if !candidate.is_finite() {
4999                continue;
5000            }
5001            let distance = (candidate - magnitude).abs();
5002            best = match best {
5003                None => Some((distance, bits)),
5004                Some((best_distance, _)) if distance < best_distance => Some((distance, bits)),
5005                Some((best_distance, best_bits))
5006                    if distance == best_distance && bits & 1 == 0 && best_bits & 1 == 1 =>
5007                {
5008                    Some((distance, bits))
5009                }
5010                other => other,
5011            };
5012        }
5013        sign | best.expect("a finite encoding exists").1
5014    }
5015
5016    #[test]
5017    fn narrowing_round_trips_every_encoding_exactly() {
5018        for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
5019            for bits in 0..=u8::MAX {
5020                let value = match format {
5021                    Fp8Format::E4M3 => fp8e4m3_to_f32(bits),
5022                    Fp8Format::E5M2 => fp8e5m2_to_f32(bits),
5023                };
5024                if !value.is_finite() {
5025                    continue;
5026                }
5027                assert_eq!(
5028                    f32_to_fp8_bits(format, value),
5029                    bits,
5030                    "{format:?}: {value} did not return to {bits:#04x}"
5031                );
5032            }
5033        }
5034    }
5035
5036    #[test]
5037    fn narrowing_matches_an_independent_nearest_search() {
5038        for format in [Fp8Format::E4M3, Fp8Format::E5M2] {
5039            let max_finite = match format {
5040                Fp8Format::E4M3 => 448.0_f32,
5041                Fp8Format::E5M2 => 57344.0_f32,
5042            };
5043            let mut state = 0x1234_5678_u32;
5044            for index in 0..200_000 {
5045                // A mix of exact midpoints, subnormal-range values, and pseudo-random draws.
5046                state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
5047                let value = if index % 3 == 0 {
5048                    let a = fp8_decode(format, (index % 256) as u8);
5049                    let b = fp8_decode(format, ((index + 1) % 256) as u8);
5050                    (a + b) * 0.5
5051                } else {
5052                    let scaled = (state >> 8) as f32 / (1_u32 << 24) as f32;
5053                    (scaled * 2.0 - 1.0) * max_finite * 1.2
5054                };
5055                if !value.is_finite() {
5056                    continue;
5057                }
5058                let actual = f32_to_fp8_bits(format, value);
5059                if value.abs() > max_finite {
5060                    continue; // overflow is this crate's policy, not a nearest-value question
5061                }
5062                assert_eq!(
5063                    actual,
5064                    nearest_by_search(format, value),
5065                    "{format:?}: {value}"
5066                );
5067            }
5068        }
5069    }
5070
5071    fn fp8_decode(format: Fp8Format, bits: u8) -> f32 {
5072        match format {
5073            Fp8Format::E4M3 => fp8e4m3_to_f32(bits),
5074            Fp8Format::E5M2 => fp8e5m2_to_f32(bits),
5075        }
5076    }
5077
5078    #[test]
5079    fn overflow_policy_is_nan_for_e4m3_and_infinity_for_e5m2() {
5080        // 464 is the midpoint above 448 and ties to even, so it stays finite; anything beyond
5081        // it is unrepresentable.
5082        assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, 464.0), 0x7e);
5083        assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, 464.001), 0x7f);
5084        assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, -464.001), 0xff);
5085        assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, f32::INFINITY), 0x7f);
5086        assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, f32::NAN), 0x7f);
5087        assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, 1.0e9), 0x7c);
5088        assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, f32::NEG_INFINITY), 0xfc);
5089        assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, f32::NAN), 0x7e);
5090        // Signed zero survives.
5091        assert_eq!(f32_to_fp8_bits(Fp8Format::E4M3, -0.0), 0x80);
5092        assert_eq!(f32_to_fp8_bits(Fp8Format::E5M2, 0.0), 0x00);
5093    }
5094}