Skip to main content

rlx_ir/
quant.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3//
4// This program is free software: you can redistribute it and/or modify
5// it under the terms of the GNU General Public License as published by
6// the Free Software Foundation, version 3.
7//
8// This program is distributed in the hope that it will be useful,
9// but WITHOUT ANY WARRANTY; without even the implied warranty of
10// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
11// GNU General Public License for more details.
12//
13// You should have received a copy of the GNU General Public License
14// along with this program. If not, see <https://www.gnu.org/licenses/>.
15
16//! Quantization metadata as graph annotations (plan #57).
17//!
18//! Schemes attach per-tensor via [`QuantMap`] on [`crate::Graph`], not inside
19//! [`crate::Node`], so the common case stays lean.
20//!
21//! # GGUF schemes and backends
22//!
23//! All `Gguf*` variants use two-input `Op::DequantMatMul` (`x`, packed weights).
24//! GPU backends share integer **scheme ids** (0–23) in `dequant_gguf` kernels;
25//! legacy tail: Q4_0 = 19, Q8_0 = 20, Q4_1 = 21, Q5_0 = 22, Q5_1 = 23.
26//!
27//! Per-backend dispatch (GPU dequant, fused GEMV, ANE constexpr, TPU compile-time
28//! bake, env toggles): [`docs/gguf-backend-paths.md`](../../docs/gguf-backend-paths.md).
29
30use crate::NodeId;
31use std::collections::HashMap;
32
33/// How a tensor is quantized. Mirrors the schemes RLX needs for LLM
34/// inference on Apple Silicon: blockwise int8 (GPTQ-style),
35/// blockwise int4 (Q4_K), and per-tensor fp8 (e4m3 / e5m2).
36///
37/// Each variant carries the parameters the dequantizer needs to read
38/// at runtime — scale, zero-point, block size. Where these live in
39/// the actual weight tensor is up to the loader (#56).
40#[cfg_attr(feature = "serialize", derive(serde::Serialize, serde::Deserialize))]
41#[derive(Debug, Clone, Copy, PartialEq)]
42pub enum QuantScheme {
43    /// Symmetric int8 with one scale per `block_size` elements.
44    Int8Block { block_size: u32 },
45    /// Asymmetric int8 with scale + zero-point per `block_size` elements.
46    Int8BlockAsym { block_size: u32 },
47    /// Int4 packed two-per-byte, scale per `block_size` elements
48    /// (Q4_K-ish; matches GGUF block layout).
49    Int4Block { block_size: u32 },
50    /// FP8 e4m3 (no scale; same domain as half).
51    Fp8E4m3,
52    /// FP8 e5m2 (no scale; wider range than e4m3).
53    Fp8E5m2,
54    /// GGUF / llama.cpp Q4_K super-block (256 elements / 144 bytes).
55    /// Packs an f16 super-scale + f16 super-min + 8 sub-block 6-bit
56    /// scales + 8 sub-block 6-bit mins + 128 nibbles. Block layout is
57    /// fixed by the format — there's no `block_size` knob.
58    GgufQ4K,
59    /// GGUF Q5_K (256 / 176 bytes). Adds a 32-byte high-bit plane on
60    /// top of Q4_K.
61    GgufQ5K,
62    /// GGUF Q6_K (256 / 210 bytes). Per-sub-block signed scales,
63    /// no min term.
64    GgufQ6K,
65    /// GGUF Q8_K (256 / 276 bytes). Per-super-block f32 scale plus
66    /// i8 quants and a 32-byte sum-of-blocks table that's only used
67    /// by Q8_K × Q8_K matmul accumulation paths.
68    GgufQ8K,
69    /// GGUF Q2_K (256 / 84 bytes). 2-bit quants with per-sub-block scale/min.
70    GgufQ2K,
71    /// GGUF Q3_K (256 / 110 bytes). 3-bit quants with hmask high bit plane.
72    GgufQ3K,
73    /// GGUF Q4_0 (32 / 18 bytes). Legacy llama.cpp block: f16 scale + signed nibbles (−8..7).
74    GgufQ4_0,
75    /// GGUF Q4_1 (32 / 20 bytes). Legacy block: f16 scale + f16 min + unsigned nibbles (0..15).
76    /// GPU kernel scheme id **21** (shared with Metal/CUDA/ROCm/WGPU).
77    GgufQ4_1,
78    /// GGUF Q5_0 (32 / 22 bytes). Legacy block: f16 scale + 32-bit high plane + signed 5-bit quants.
79    /// GPU kernel scheme id **22**.
80    GgufQ5_0,
81    /// GGUF Q5_1 (32 / 24 bytes). Legacy block: f16 scale + f16 min + high plane + unsigned 5-bit quants.
82    /// GPU kernel scheme id **23**.
83    GgufQ5_1,
84    /// GGUF Q8_0 (32 / 34 bytes). Legacy block: f16 scale + 32×i8 quants.
85    GgufQ8_0,
86    /// NVIDIA FP4 (E2M1) block — fixed 16-element groups, FP8 E4M3 block
87    /// scales, optional f32 global scale on input 3 (legacy `zp` slot).
88    /// Used by FLUX.2 / MLX `nvfp4` checkpoints.
89    Nvfp4Block,
90    // ── GGUF IQ-family (sub-byte LUT-coded) ─────────────────────
91    /// IQ4_NL: 4.5 bpw non-linear. 32-element block (18 bytes).
92    GgufIQ4NL,
93    /// IQ4_XS: 4.25 bpw. 256-element super-block (136 bytes).
94    GgufIQ4XS,
95    /// IQ2_XXS: 2.0625 bpw. 256 / 66.
96    GgufIQ2XXS,
97    /// IQ2_XS: 2.3125 bpw. 256 / 74.
98    GgufIQ2XS,
99    /// IQ2_S: 2.5625 bpw. 256 / 82.
100    GgufIQ2S,
101    /// IQ3_XXS: 3.0625 bpw. 256 / 98.
102    GgufIQ3XXS,
103    /// IQ3_S: 3.4375 bpw. 256 / 110.
104    GgufIQ3S,
105    /// IQ1_S: 1.5625 bpw. 256 / 50.
106    GgufIQ1S,
107    /// IQ1_M: 1.75 bpw. 256 / 56.
108    GgufIQ1M,
109    /// TQ1_0: 1.6875 bpw ternary. 256 / 54.
110    GgufTQ1_0,
111    /// TQ2_0: 2.0625 bpw ternary. 256 / 66.
112    GgufTQ2_0,
113    /// MXFP4: OCP microscaling FP4 with E8M0 scale. 32 / 17.
114    GgufMXFP4,
115    /// NVFP4 GGUF variant: E4M3 scale + E2M1 nibbles. 16 / 9.
116    GgufNVFP4,
117    /// Custom 1-bit format (PrismML `Bonsai-27B`, ggml type 41): one
118    /// f16 group scale + 128 sign bits selecting `±d`. 128 / 18 bytes
119    /// (1.125 bpw). See [`crate`] consumers `rlx_gguf::q1_dequant`.
120    GgufQ1_0,
121    /// PrismML / upstream `Q2_0` (ggml type 42): f16 group scale + 128×2-bit
122    /// codes → `(q−1)·d`. 128 / 34 bytes (2.125 bpw). Used by Ternary Bonsai
123    /// (`prism-ml/Ternary-Bonsai-*-gguf`). See `rlx_gguf::q2_dequant`.
124    GgufQ2_0,
125}
126
127/// Single source of truth for GGUF → GPU `dequant_gguf` scheme ids.
128///
129/// Expanding this macro inside `impl QuantScheme` defines:
130/// - [`QuantScheme::gpu_dequant_scheme_id`]
131/// - [`QuantScheme::from_gpu_dequant_scheme_id`]
132/// - test helper [`QuantScheme::GPU_DEQUANT_SCHEME_ID_PAIRS`]
133///
134/// GPU backends (`gguf_host::gguf_scheme_id` / `scheme_from_id`) must
135/// delegate here — do not re-list variants by hand.
136macro_rules! define_gguf_gpu_dequant_ids {
137    ($(($variant:ident, $id:literal)),+ $(,)?) => {
138        /// Canonical numeric id for the shared on-device GGUF dequant kernel
139        /// (`dequant_gguf` in MSL / CU / WGSL / GLSL).
140        ///
141        /// Generated by `define_gguf_gpu_dequant_ids!` so encode and
142        /// [`Self::from_gpu_dequant_scheme_id`] share one table — backends
143        /// must only call these helpers (no hand-rolled `match` in
144        /// `gguf_host`). Adding a scheme is one macro row + the matching
145        /// kernel branch. Guarded by `gpu_dequant_scheme_id_is_stable`.
146        pub const fn gpu_dequant_scheme_id(self) -> Option<u32> {
147            match self {
148                $(Self::$variant => Some($id),)+
149                _ => None,
150            }
151        }
152
153        /// Inverse of [`Self::gpu_dequant_scheme_id`].
154        pub const fn from_gpu_dequant_scheme_id(id: u32) -> Option<Self> {
155            match id {
156                $($id => Some(Self::$variant),)+
157                _ => None,
158            }
159        }
160
161        /// `(scheme, id)` pairs for tests / docs — keep in sync via the macro.
162        #[cfg(test)]
163        pub(crate) const GPU_DEQUANT_SCHEME_ID_PAIRS: &'static [(Self, u32)] = &[
164            $((Self::$variant, $id),)+
165        ];
166    };
167}
168
169impl QuantScheme {
170    /// Bits per element after packing (×10 for K-quants since they
171    /// have fractional bit budgets — divide by 10 when comparing).
172    pub const fn bits_per_element_x10(self) -> u32 {
173        match self {
174            Self::Int8Block { .. } | Self::Int8BlockAsym { .. } => 80,
175            Self::Int4Block { .. } => 40,
176            Self::Fp8E4m3 | Self::Fp8E5m2 => 80,
177            // GGUF K-quants: header + per-element bits over a 256-element block.
178            Self::GgufQ4K => 45,  // 144 bytes / 256 elems × 8 = 4.5 bpe
179            Self::GgufQ5K => 55,  // 176 / 256 × 8 ≈ 5.5
180            Self::GgufQ6K => 66,  // 210 / 256 × 8 ≈ 6.5625 → 66 (rounded)
181            Self::GgufQ8K => 91,  // 292 / 256 × 8 ≈ 9.125 → 91
182            Self::GgufQ2K => 26,  // 84 / 256 × 8 ≈ 2.625 → 26
183            Self::GgufQ3K => 34,  // 110 / 256 × 8 ≈ 3.4375 → 34
184            Self::GgufQ4_0 => 45, // 18 / 32 × 8 = 4.5 bpe
185            Self::GgufQ4_1 => 50, // 20 / 32 × 8 = 5.0 bpe
186            Self::GgufQ5_0 => 55, // 22 / 32 × 8 = 5.5 bpe
187            Self::GgufQ5_1 => 60, // 24 / 32 × 8 = 6.0 bpe
188            Self::GgufQ8_0 => 85, // 34 / 32 × 8 = 8.5 bpe
189            Self::Nvfp4Block => 40,
190            Self::GgufIQ4NL => 45,
191            Self::GgufIQ4XS => 42, // 136/256 × 8 = 4.25
192            Self::GgufIQ2XXS => 20,
193            Self::GgufIQ2XS => 23,
194            Self::GgufIQ2S => 25,
195            Self::GgufIQ3XXS => 30,
196            Self::GgufIQ3S => 34,
197            Self::GgufIQ1S => 15,
198            Self::GgufIQ1M => 17,
199            Self::GgufTQ1_0 => 16, // 54/256 × 8 = 1.6875 → 16
200            Self::GgufTQ2_0 => 20,
201            Self::GgufMXFP4 => 42,
202            Self::GgufNVFP4 => 45,
203            Self::GgufQ1_0 => 11, // 18 bytes / 128 elems × 8 = 1.125 bpe
204            Self::GgufQ2_0 => 21, // 34 bytes / 128 elems × 8 = 2.125 bpe
205        }
206    }
207
208    /// Bits per element after packing (rounded down). Use
209    /// `bits_per_element_x10` for the K-quant fractional values.
210    pub const fn bits_per_element(self) -> u32 {
211        self.bits_per_element_x10() / 10
212    }
213
214    // Table defined once via `define_gguf_gpu_dequant_ids!` above.
215    define_gguf_gpu_dequant_ids! {
216        (GgufQ4K, 0),
217        (GgufQ5K, 1),
218        (GgufQ6K, 2),
219        (GgufQ8K, 3),
220        (GgufQ2K, 4),
221        (GgufQ3K, 5),
222        (GgufIQ4NL, 6),
223        (GgufIQ4XS, 7),
224        (GgufTQ1_0, 8),
225        (GgufTQ2_0, 9),
226        (GgufMXFP4, 10),
227        (GgufNVFP4, 11),
228        (GgufIQ2XXS, 12),
229        (GgufIQ2XS, 13),
230        (GgufIQ2S, 14),
231        (GgufIQ3XXS, 15),
232        (GgufIQ3S, 16),
233        (GgufIQ1S, 17),
234        (GgufIQ1M, 18),
235        (GgufQ4_0, 19),
236        (GgufQ8_0, 20),
237        (GgufQ4_1, 21),
238        (GgufQ5_0, 22),
239        (GgufQ5_1, 23),
240        (GgufQ1_0, 24),
241        (GgufQ2_0, 25),
242    }
243
244    /// True if this scheme requires a per-block scale tensor on the side.
245    pub const fn has_scale(self) -> bool {
246        matches!(
247            self,
248            Self::Int8Block { .. }
249                | Self::Int8BlockAsym { .. }
250                | Self::Int4Block { .. }
251                | Self::Nvfp4Block
252        )
253    }
254
255    /// True for NVFP4 block scales stored as FP8 E4M3 bytes (not f32).
256    pub const fn scale_is_fp8(self) -> bool {
257        matches!(self, Self::Nvfp4Block)
258    }
259
260    /// Fixed NVFP4 group size along K (0 for other schemes).
261    pub const fn nvfp4_group_size(self) -> u32 {
262        match self {
263            Self::Nvfp4Block => crate::nvfp4::NVFP4_GROUP_SIZE as u32,
264            _ => 0,
265        }
266    }
267
268    /// True if this scheme requires a per-block zero-point.
269    pub const fn has_zero_point(self) -> bool {
270        matches!(self, Self::Int8BlockAsym { .. })
271    }
272
273    /// GGUF K-quant block size (256 elements) — meaningless for the
274    /// non-GGUF schemes (returns 0).
275    pub const fn gguf_block_size(self) -> u32 {
276        match self {
277            Self::GgufQ4K
278            | Self::GgufQ5K
279            | Self::GgufQ6K
280            | Self::GgufQ8K
281            | Self::GgufQ2K
282            | Self::GgufQ3K
283            | Self::GgufIQ4XS
284            | Self::GgufIQ2XXS
285            | Self::GgufIQ2XS
286            | Self::GgufIQ2S
287            | Self::GgufIQ3XXS
288            | Self::GgufIQ3S
289            | Self::GgufIQ1S
290            | Self::GgufIQ1M
291            | Self::GgufTQ1_0
292            | Self::GgufTQ2_0 => 256,
293            Self::GgufQ4_0
294            | Self::GgufQ4_1
295            | Self::GgufQ5_0
296            | Self::GgufQ5_1
297            | Self::GgufQ8_0
298            | Self::GgufIQ4NL
299            | Self::GgufMXFP4 => 32,
300            Self::GgufNVFP4 => 16,
301            Self::GgufQ1_0 => 128,
302            Self::GgufQ2_0 => 128,
303            _ => 0,
304        }
305    }
306
307    /// Bytes per GGUF super-block. 0 for non-GGUF schemes.
308    pub const fn gguf_block_bytes(self) -> u32 {
309        match self {
310            Self::GgufQ4K => 144, // f16 d + f16 dmin + 12 packed scales + 128 nibbles
311            Self::GgufQ5K => 176, // + 32-byte high-bit plane
312            Self::GgufQ6K => 210, // 128 ql + 64 qh + 16 i8 scales + f16 d
313            Self::GgufQ8K => 292, // f32 d + 256 i8 + 16 i16 bsums = 4 + 256 + 32
314            Self::GgufQ2K => 84,  // f16 d + f16 dmin + 16 scales + 64 qs
315            Self::GgufQ3K => 110, // f16 d + 12 scales + 32 hmask + 64 qs
316            Self::GgufQ4_0 => 18, // f16 d + 16 packed nibbles
317            Self::GgufQ4_1 => 20, // f16 d + f16 min + 16 packed nibbles
318            Self::GgufQ5_0 => 22, // f16 d + 32-bit qh + 16 packed nibbles
319            Self::GgufQ5_1 => 24, // f16 d + f16 min + 32-bit qh + 16 packed nibbles
320            Self::GgufQ8_0 => 34, // f16 d + 32 i8 quants
321            Self::GgufIQ4NL => 18,
322            Self::GgufIQ4XS => 136,
323            Self::GgufIQ2XXS => 66,
324            Self::GgufIQ2XS => 74,
325            Self::GgufIQ2S => 82,
326            Self::GgufIQ3XXS => 98,
327            Self::GgufIQ3S => 110,
328            Self::GgufIQ1S => 50,
329            Self::GgufIQ1M => 56,
330            Self::GgufTQ1_0 => 54,
331            Self::GgufTQ2_0 => 66,
332            Self::GgufMXFP4 => 17,
333            Self::GgufNVFP4 => 9,
334            Self::GgufQ1_0 => 18, // f16 scale + 128 sign bits
335            Self::GgufQ2_0 => 34, // Metal fused Q2_0 block (128 elems)
336            _ => 0,
337        }
338    }
339
340    /// True for any GGUF-format block scheme. GGUF schemes carry
341    /// their scales / mins / sub-block metadata *inside* the packed
342    /// weight bytes — they don't need separate `scale` / `zp`
343    /// tensors fed alongside as the legacy `Int8Block` paths do.
344    pub const fn is_gguf(self) -> bool {
345        matches!(
346            self,
347            Self::GgufQ4K
348                | Self::GgufQ5K
349                | Self::GgufQ6K
350                | Self::GgufQ8K
351                | Self::GgufQ2K
352                | Self::GgufQ3K
353                | Self::GgufQ4_0
354                | Self::GgufQ4_1
355                | Self::GgufQ5_0
356                | Self::GgufQ5_1
357                | Self::GgufQ8_0
358                | Self::GgufIQ4NL
359                | Self::GgufIQ4XS
360                | Self::GgufIQ2XXS
361                | Self::GgufIQ2XS
362                | Self::GgufIQ2S
363                | Self::GgufIQ3XXS
364                | Self::GgufIQ3S
365                | Self::GgufIQ1S
366                | Self::GgufIQ1M
367                | Self::GgufTQ1_0
368                | Self::GgufTQ2_0
369                | Self::GgufMXFP4
370                | Self::GgufNVFP4
371                | Self::GgufQ1_0
372                | Self::GgufQ2_0
373        )
374    }
375}
376
377impl std::fmt::Display for QuantScheme {
378    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
379        match self {
380            Self::Int8Block { block_size } => write!(f, "int8/{block_size}"),
381            Self::Int8BlockAsym { block_size } => write!(f, "int8a/{block_size}"),
382            Self::Int4Block { block_size } => write!(f, "int4/{block_size}"),
383            Self::Fp8E4m3 => write!(f, "fp8e4m3"),
384            Self::Fp8E5m2 => write!(f, "fp8e5m2"),
385            Self::GgufQ4K => write!(f, "gguf_q4k"),
386            Self::GgufQ5K => write!(f, "gguf_q5k"),
387            Self::GgufQ6K => write!(f, "gguf_q6k"),
388            Self::GgufQ8K => write!(f, "gguf_q8k"),
389            Self::GgufQ2K => write!(f, "gguf_q2k"),
390            Self::GgufQ3K => write!(f, "gguf_q3k"),
391            Self::GgufQ4_0 => write!(f, "gguf_q4_0"),
392            Self::GgufQ4_1 => write!(f, "gguf_q4_1"),
393            Self::GgufQ5_0 => write!(f, "gguf_q5_0"),
394            Self::GgufQ5_1 => write!(f, "gguf_q5_1"),
395            Self::GgufQ8_0 => write!(f, "gguf_q8_0"),
396            Self::Nvfp4Block => write!(f, "nvfp4/16"),
397            Self::GgufIQ4NL => write!(f, "gguf_iq4_nl"),
398            Self::GgufIQ4XS => write!(f, "gguf_iq4_xs"),
399            Self::GgufIQ2XXS => write!(f, "gguf_iq2_xxs"),
400            Self::GgufIQ2XS => write!(f, "gguf_iq2_xs"),
401            Self::GgufIQ2S => write!(f, "gguf_iq2_s"),
402            Self::GgufIQ3XXS => write!(f, "gguf_iq3_xxs"),
403            Self::GgufIQ3S => write!(f, "gguf_iq3_s"),
404            Self::GgufIQ1S => write!(f, "gguf_iq1_s"),
405            Self::GgufIQ1M => write!(f, "gguf_iq1_m"),
406            Self::GgufTQ1_0 => write!(f, "gguf_tq1_0"),
407            Self::GgufTQ2_0 => write!(f, "gguf_tq2_0"),
408            Self::GgufMXFP4 => write!(f, "gguf_mxfp4"),
409            Self::GgufNVFP4 => write!(f, "gguf_nvfp4"),
410            Self::GgufQ1_0 => write!(f, "gguf_q1_0"),
411            Self::GgufQ2_0 => write!(f, "gguf_q2_0"),
412        }
413    }
414}
415
416/// Element format for a **native low-precision tensor-core GEMM** operand
417/// (see [`crate::op::Op::ScaledMatMul`]).
418///
419/// Distinct from [`QuantScheme`]: a `QuantScheme` is block-scaled *storage*
420/// that the CPU/GPU decodes to f32 *before* a normal sgemm. A `ScaledFormat`
421/// is the raw element encoding hardware tensor cores consume **directly**, with
422/// f32 accumulation — the whole point of FP8/FP6/FP4 on Hopper / Ada /
423/// Blackwell / CDNA3 / CDNA4. The per-block/per-tensor scale layout is
424/// orthogonal and lives in [`ScaleLayout`].
425///
426/// Operands flow through the graph as `DType::U8` byte buffers (one code per
427/// byte on the CPU oracle; packed on GPU); the format is carried on the op,
428/// not the dtype, so no `DType` variant is needed.
429#[cfg_attr(feature = "serialize", derive(serde::Serialize, serde::Deserialize))]
430#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
431pub enum ScaledFormat {
432    /// OCP FP8 E4M3 — bias 7, no infinities, only `S.1111.111` is NaN.
433    /// Max ±448. Hopper / Ada / Blackwell; weights + forward activations.
434    F8E4M3,
435    /// OCP FP8 E5M2 — bias 15, IEEE-like inf/NaN at max exponent.
436    /// Max ±57344. Wider range — the gradient format.
437    F8E5M2,
438    /// AMD "FNUZ" FP8 E4M3 — bias 8, single NaN at `0x80`, no inf, no −0.
439    /// Max ±240. CDNA3 / MI300.
440    F8E4M3Fnuz,
441    /// AMD "FNUZ" FP8 E5M2 — bias 16, single NaN at `0x80`, no inf, no −0.
442    /// Max ±57344. CDNA3 / MI300.
443    F8E5M2Fnuz,
444    /// MX FP6 E2M3 — bias 1, all 64 codes finite. Max ±7.5. Blackwell / CDNA4.
445    F6E2M3,
446    /// MX FP6 E3M2 — bias 3, all 64 codes finite. Max ±28. Blackwell / CDNA4.
447    F6E3M2,
448    /// FP4 E2M1 — all 16 codes finite. Max ±6. Blackwell / CDNA4.
449    /// Shared by NVFP4 and MXFP4 — the [`ScaleLayout`] tells them apart.
450    F4E2M1,
451    /// Arbitrary sub-byte minifloat parameterized by exponent / mantissa width
452    /// — the `fNeXmY` family (`N = 1 + exp_bits + mant_bits` total bits).
453    ///
454    /// All-finite (no infinities, no NaN, no AMD FNUZ), IEEE-style gradual
455    /// underflow (subnormals when `exp == 0`), sign in bit `exp_bits +
456    /// mant_bits`, one code per byte. The whole code must fit in a byte, so
457    /// `1 + exp_bits + mant_bits <= 8` and `exp_bits >= 1`.
458    ///
459    /// This is a *research / experimental* encoding: it has no hardware
460    /// tensor-core GEMM, so it always runs the decode-and-accumulate path (CPU
461    /// oracle + generic GPU kernel). Build one with [`ScaledFormat::custom`]
462    /// (IEEE bias `2^(exp_bits-1) - 1`) or [`ScaledFormat::custom_with_bias`],
463    /// or parse a name via `"f4e3m0".parse()`.
464    ///
465    /// Example — `f4e3m0` (`exp_bits: 3, mant_bits: 0, bias: 3`): the 16 codes
466    /// are `±0` and `±{0.25, 0.5, 1, 2, 4, 8, 16}` — a signed power-of-two
467    /// (near-logarithmic) 4-bit format.
468    Custom {
469        exp_bits: u8,
470        mant_bits: u8,
471        bias: i8,
472    },
473}
474
475impl ScaledFormat {
476    /// Total bit width of one element code (4, 6, or 8).
477    pub const fn bit_width(self) -> u32 {
478        match self {
479            Self::F8E4M3 | Self::F8E5M2 | Self::F8E4M3Fnuz | Self::F8E5M2Fnuz => 8,
480            Self::F6E2M3 | Self::F6E3M2 => 6,
481            Self::F4E2M1 => 4,
482            Self::Custom {
483                exp_bits,
484                mant_bits,
485                ..
486            } => 1 + exp_bits as u32 + mant_bits as u32,
487        }
488    }
489
490    /// `(exponent_bits, mantissa_bits, exponent_bias)`.
491    pub const fn fields(self) -> (u32, u32, i32) {
492        match self {
493            Self::F8E4M3 => (4, 3, 7),
494            Self::F8E5M2 => (5, 2, 15),
495            Self::F8E4M3Fnuz => (4, 3, 8),
496            Self::F8E5M2Fnuz => (5, 2, 16),
497            Self::F6E2M3 => (2, 3, 1),
498            Self::F6E3M2 => (3, 2, 3),
499            Self::F4E2M1 => (2, 1, 1),
500            Self::Custom {
501                exp_bits,
502                mant_bits,
503                bias,
504            } => (exp_bits as u32, mant_bits as u32, bias as i32),
505        }
506    }
507
508    /// Build an arbitrary all-finite minifloat with `exp_bits` exponent and
509    /// `mant_bits` mantissa bits, using the IEEE-style bias `2^(exp_bits-1) - 1`
510    /// — the convention every named rlx format follows (E4M3 → 7, E5M2 → 15,
511    /// E2M1 → 1, …). The whole code must fit in a byte. `const` — usable in
512    /// `const` contexts and pattern-free construction.
513    ///
514    /// Panics if `exp_bits == 0` or `1 + exp_bits + mant_bits > 8`.
515    ///
516    /// ```
517    /// use rlx_ir::ScaledFormat;
518    /// const F4E3M0: ScaledFormat = ScaledFormat::custom(3, 0);
519    /// assert_eq!(F4E3M0.to_string(), "f4e3m0");
520    /// assert_eq!(F4E3M0.exp_bits(), 3);
521    /// assert_eq!(F4E3M0.mant_bits(), 0);
522    /// // Nearest representable value ("quantize" a single f32):
523    /// assert_eq!(F4E3M0.quantize(1.4), 1.0);
524    /// assert_eq!(F4E3M0.max_finite(), 16.0);
525    /// ```
526    pub const fn custom(exp_bits: u8, mant_bits: u8) -> Self {
527        assert!(exp_bits >= 1, "minifloat needs at least one exponent bit");
528        assert!(
529            1 + exp_bits + mant_bits <= 8,
530            "minifloat code must fit in a byte (1 + exp + mant <= 8)"
531        );
532        let bias = (1i32 << (exp_bits - 1)) - 1;
533        Self::Custom {
534            exp_bits,
535            mant_bits,
536            bias: bias as i8,
537        }
538    }
539
540    /// Like [`custom`](Self::custom) but with an explicit exponent bias
541    /// (for formats that don't use the IEEE `2^(e-1)-1` convention). `const`.
542    pub const fn custom_with_bias(exp_bits: u8, mant_bits: u8, bias: i8) -> Self {
543        assert!(exp_bits >= 1, "minifloat needs at least one exponent bit");
544        assert!(
545            1 + exp_bits + mant_bits <= 8,
546            "minifloat code must fit in a byte (1 + exp + mant <= 8)"
547        );
548        Self::Custom {
549            exp_bits,
550            mant_bits,
551            bias,
552        }
553    }
554
555    /// The seven named hardware element formats, for enumeration / format sweeps.
556    pub const NAMED: [ScaledFormat; 7] = [
557        Self::F8E4M3,
558        Self::F8E5M2,
559        Self::F8E4M3Fnuz,
560        Self::F8E5M2Fnuz,
561        Self::F6E2M3,
562        Self::F6E3M2,
563        Self::F4E2M1,
564    ];
565
566    /// Exponent-field bit count (`X` in `fNeXmY`).
567    pub const fn exp_bits(self) -> u32 {
568        self.fields().0
569    }
570    /// Mantissa-field bit count (`Y` in `fNeXmY`; `0` for a pure power-of-two
571    /// format such as `f4e3m0`).
572    pub const fn mant_bits(self) -> u32 {
573        self.fields().1
574    }
575    /// Exponent bias.
576    pub const fn bias(self) -> i32 {
577        self.fields().2
578    }
579    /// True for a parameterized [`Custom`](Self::Custom) format (no hardware
580    /// tensor core — always runs the decode path).
581    pub const fn is_custom(self) -> bool {
582        matches!(self, Self::Custom { .. })
583    }
584    /// True for one of the seven [`NAMED`](Self::NAMED) hardware formats.
585    pub const fn is_named(self) -> bool {
586        !self.is_custom()
587    }
588
589    /// AMD OCP-FNUZ encoding (single NaN at `0x80`, no inf, no −0).
590    pub const fn is_fnuz(self) -> bool {
591        matches!(self, Self::F8E4M3Fnuz | Self::F8E5M2Fnuz)
592    }
593
594    /// True for IEEE-like formats carrying inf / NaN at the max exponent
595    /// (OCP E5M2 only — E4M3 is finite-with-one-NaN, FP6/FP4 are all-finite).
596    pub const fn has_inf(self) -> bool {
597        matches!(self, Self::F8E5M2)
598    }
599
600    /// Largest finite magnitude representable (used as the amax divisor when
601    /// computing a per-tensor / per-block scale).
602    pub fn max_finite(self) -> f32 {
603        crate::lowp_codec::max_finite(self)
604    }
605
606    /// Decode one packed code to f32 (bit-exact; `±inf` / `NaN` for the formats
607    /// that encode them). Method form of [`crate::lowp_codec::decode`].
608    pub fn decode(self, code: u8) -> f32 {
609        crate::lowp_codec::decode(self, code)
610    }
611
612    /// Encode an f32 to its nearest representable code (round-half-to-even,
613    /// saturating overflow / `±inf` to `±max_finite`, `NaN → 0`). Method form of
614    /// [`crate::lowp_codec::encode`].
615    pub fn encode(self, x: f32) -> u8 {
616        crate::lowp_codec::encode(self, x)
617    }
618
619    /// Round-trip `x` through the format — the nearest value it can represent.
620    /// Handy for eyeballing a format's resolution without building a graph
621    /// (e.g. `ScaledFormat::custom(3, 0).quantize(1.4) == 1.0`).
622    pub fn quantize(self, x: f32) -> f32 {
623        self.decode(self.encode(x))
624    }
625
626    /// Every finite value the format represents, ascending and deduplicated —
627    /// the format's grid. Useful for inspecting a research minifloat, e.g.
628    /// `ScaledFormat::custom(3, 0).representable_values()` yields
629    /// `[-16, -8, -4, -2, -1, -0.5, -0.25, 0, 0.25, 0.5, 1, 2, 4, 8, 16]`.
630    pub fn representable_values(self) -> Vec<f32> {
631        let n = 1u16 << self.bit_width();
632        let mut v: Vec<f32> = (0..n)
633            .map(|c| self.decode(c as u8))
634            .filter(|x| x.is_finite())
635            .collect();
636        v.sort_by(|a, b| a.partial_cmp(b).unwrap());
637        v.dedup();
638        v
639    }
640
641    /// The GPU dispatch word passed to `scaled_lowp_general.cu` (`rlx_decode_lowp`).
642    ///
643    /// For the seven named formats this is a small stable id matching the
644    /// kernel's `switch`: e4m3=0, e5m2=1, e4m3fnuz=2, e5m2fnuz=3, e2m3=4,
645    /// e3m2=5, e2m1=6.
646    ///
647    /// For a [`Custom`](Self::Custom) format there is no fixed id, so this
648    /// returns a **packed field descriptor** with the top bit set as a
649    /// sentinel: `0x8000_0000 | exp_bits | mant_bits<<4 | (bias as u8)<<8`.
650    /// The kernel detects the sentinel bit and decodes generically from the
651    /// unpacked `(exp_bits, mant_bits, bias)` — so a new format needs no kernel
652    /// edit. The named ids stay in `0..=6` (top bit clear), so the existing
653    /// hardware `switch` path is unchanged.
654    pub const fn kernel_id(self) -> u32 {
655        match self {
656            Self::F8E4M3 => 0,
657            Self::F8E5M2 => 1,
658            Self::F8E4M3Fnuz => 2,
659            Self::F8E5M2Fnuz => 3,
660            Self::F6E2M3 => 4,
661            Self::F6E3M2 => 5,
662            Self::F4E2M1 => 6,
663            Self::Custom {
664                exp_bits,
665                mant_bits,
666                bias,
667            } => {
668                0x8000_0000
669                    | (exp_bits as u32 & 0xF)
670                    | ((mant_bits as u32 & 0xF) << 4)
671                    | (((bias as u8) as u32) << 8)
672            }
673        }
674    }
675
676    /// True for the OCP FP8 variants that the native cublasLt / hipBLASLt FP8
677    /// GEMM accepts (per-tensor only). Other formats use the decode fallback.
678    pub const fn is_native_fp8(self) -> bool {
679        matches!(self, Self::F8E4M3 | Self::F8E5M2)
680    }
681}
682
683impl std::fmt::Display for ScaledFormat {
684    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
685        let s = match self {
686            Self::F8E4M3 => "f8e4m3",
687            Self::F8E5M2 => "f8e5m2",
688            Self::F8E4M3Fnuz => "f8e4m3fnuz",
689            Self::F8E5M2Fnuz => "f8e5m2fnuz",
690            Self::F6E2M3 => "f6e2m3",
691            Self::F6E3M2 => "f6e3m2",
692            Self::F4E2M1 => "f4e2m1",
693            // `fNeXmY`, N = 1 + exp + mant. The bias is implied by the width
694            // (IEEE convention) and not rendered; round-trips through `FromStr`.
695            Self::Custom {
696                exp_bits,
697                mant_bits,
698                ..
699            } => {
700                return write!(
701                    f,
702                    "f{}e{}m{}",
703                    1 + exp_bits + mant_bits,
704                    exp_bits,
705                    mant_bits
706                );
707            }
708        };
709        f.write_str(s)
710    }
711}
712
713impl std::str::FromStr for ScaledFormat {
714    type Err = String;
715
716    /// Parse a format name. The seven named formats (`"f8e4m3"`, `"f4e2m1"`,
717    /// the `…fnuz` variants, …) keep their exact hardware semantics (FNUZ,
718    /// inf/NaN); any other `fNeXmY` string becomes an all-finite
719    /// [`Custom`](Self::Custom) with the IEEE bias. `N` must equal `1 + X + Y`
720    /// and the code must fit in a byte (`N <= 8`, `X >= 1`).
721    fn from_str(s: &str) -> Result<Self, Self::Err> {
722        match s {
723            "f8e4m3" => return Ok(Self::F8E4M3),
724            "f8e5m2" => return Ok(Self::F8E5M2),
725            "f8e4m3fnuz" => return Ok(Self::F8E4M3Fnuz),
726            "f8e5m2fnuz" => return Ok(Self::F8E5M2Fnuz),
727            "f6e2m3" => return Ok(Self::F6E2M3),
728            "f6e3m2" => return Ok(Self::F6E3M2),
729            "f4e2m1" => return Ok(Self::F4E2M1),
730            _ => {}
731        }
732        let err = || format!("invalid minifloat format {s:?} (expected e.g. \"f4e3m0\")");
733        let rest = s.strip_prefix('f').ok_or_else(err)?;
734        let (total, rest) = take_leading_u32(rest).ok_or_else(err)?;
735        let rest = rest.strip_prefix('e').ok_or_else(err)?;
736        let (exp, rest) = take_leading_u32(rest).ok_or_else(err)?;
737        let rest = rest.strip_prefix('m').ok_or_else(err)?;
738        let (mant, rest) = take_leading_u32(rest).ok_or_else(err)?;
739        if !rest.is_empty() || exp == 0 || total != 1 + exp + mant || 1 + exp + mant > 8 {
740            return Err(err());
741        }
742        Ok(Self::custom(exp as u8, mant as u8))
743    }
744}
745
746/// Split the leading run of ASCII digits off `s`, returning `(value, tail)`.
747/// `None` if `s` doesn't start with a digit.
748fn take_leading_u32(s: &str) -> Option<(u32, &str)> {
749    let end = s.find(|c: char| !c.is_ascii_digit()).unwrap_or(s.len());
750    if end == 0 {
751        return None;
752    }
753    Some((s[..end].parse().ok()?, &s[end..]))
754}
755
756/// How scale factors are laid out for an [`crate::op::Op::ScaledMatMul`]
757/// operand. The reconstructed value of element `i` is
758/// `decode(code[i]) * scale(block_of(i))`.
759#[cfg_attr(feature = "serialize", derive(serde::Serialize, serde::Deserialize))]
760#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
761pub enum ScaleLayout {
762    /// One f32 amax scale for the whole operand (classic per-tensor FP8 GEMM).
763    PerTensor,
764    /// OCP microscaling: one power-of-two **E8M0** scale per `block`
765    /// consecutive elements along K (MXFP8 / MXFP6 / MXFP4). Scale tensor
766    /// is `DType::U8` (E8M0 bytes).
767    BlockMxE8M0 { block: u32 },
768    /// NVFP4: one FP8 **E4M3** scale per `group` (16) elements along K, plus
769    /// an optional per-tensor f32 global. Scale tensor is `DType::U8`.
770    Nvfp4 { group: u32 },
771}
772
773impl ScaleLayout {
774    /// OCP microscaling default (32-element blocks).
775    pub const fn mx() -> Self {
776        Self::BlockMxE8M0 { block: 32 }
777    }
778    /// NVFP4 default (16-element groups).
779    pub fn nvfp4() -> Self {
780        Self::Nvfp4 {
781            group: crate::nvfp4::NVFP4_GROUP_SIZE as u32,
782        }
783    }
784    /// Element dtype of the scale tensor for this layout.
785    pub const fn scale_dtype(self) -> crate::DType {
786        match self {
787            Self::PerTensor => crate::DType::F32,
788            Self::BlockMxE8M0 { .. } | Self::Nvfp4 { .. } => crate::DType::U8,
789        }
790    }
791    /// Number of consecutive elements sharing one scale (1 for per-tensor).
792    pub const fn block(self) -> u32 {
793        match self {
794            Self::PerTensor => 1,
795            Self::BlockMxE8M0 { block } => block,
796            Self::Nvfp4 { group } => group,
797        }
798    }
799
800    /// `(scale_mode, block)` for GPU kernels (`scaled_lowp_general.cu`):
801    /// per-tensor=0, block-E8M0=1, NVFP4-E4M3=2.
802    pub const fn mode_block(self) -> (u32, u32) {
803        match self {
804            Self::PerTensor => (0, 1),
805            Self::BlockMxE8M0 { block } => (1, block),
806            Self::Nvfp4 { group } => (2, group),
807        }
808    }
809}
810
811impl std::fmt::Display for ScaleLayout {
812    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
813        match self {
814            Self::PerTensor => write!(f, "per_tensor"),
815            Self::BlockMxE8M0 { block } => write!(f, "mx_e8m0/{block}"),
816            Self::Nvfp4 { group } => write!(f, "nvfp4/{group}"),
817        }
818    }
819}
820
821impl std::str::FromStr for ScaleLayout {
822    type Err = String;
823
824    /// Parse a scale-layout name: `"per_tensor"`, `"mx"` (32-element E8M0
825    /// microscaling), `"nvfp4"` (16-element E4M3), or an explicit block size
826    /// `"mx/<block>"` / `"nvfp4/<group>"`. Mirrors [`Display`](Self) round-trip.
827    fn from_str(s: &str) -> Result<Self, Self::Err> {
828        match s {
829            "per_tensor" | "pertensor" => return Ok(Self::PerTensor),
830            "mx" | "mxfp8" | "block" => return Ok(Self::mx()),
831            "nvfp4" => return Ok(Self::nvfp4()),
832            _ => {}
833        }
834        if let Some(rest) = s.strip_prefix("mx/").or_else(|| s.strip_prefix("mx_e8m0/")) {
835            if let Ok(block) = rest.parse::<u32>() {
836                return Ok(Self::BlockMxE8M0 { block });
837            }
838        }
839        if let Some(rest) = s.strip_prefix("nvfp4/") {
840            if let Ok(group) = rest.parse::<u32>() {
841                return Ok(Self::Nvfp4 { group });
842            }
843        }
844        Err(format!(
845            "unknown scale layout {s:?} (expected \"per_tensor\", \"mx\", \"nvfp4\", or \"mx/<block>\")"
846        ))
847    }
848}
849
850/// Per-graph map of quantized tensors. Lookup is O(1).
851#[derive(Debug, Clone, Default)]
852pub struct QuantMap {
853    map: HashMap<NodeId, QuantScheme>,
854}
855
856impl QuantMap {
857    pub fn new() -> Self {
858        Self::default()
859    }
860    pub fn get(&self, id: NodeId) -> Option<QuantScheme> {
861        self.map.get(&id).copied()
862    }
863    pub fn insert(&mut self, id: NodeId, scheme: QuantScheme) -> Option<QuantScheme> {
864        self.map.insert(id, scheme)
865    }
866    pub fn is_empty(&self) -> bool {
867        self.map.is_empty()
868    }
869    pub fn len(&self) -> usize {
870        self.map.len()
871    }
872    pub fn iter(&self) -> impl Iterator<Item = (&NodeId, &QuantScheme)> {
873        self.map.iter()
874    }
875}
876
877#[cfg(test)]
878mod tests {
879    use super::*;
880
881    #[test]
882    fn scheme_traits() {
883        assert_eq!(
884            QuantScheme::Int4Block { block_size: 32 }.bits_per_element(),
885            4
886        );
887        assert!(QuantScheme::Int8BlockAsym { block_size: 64 }.has_zero_point());
888        assert!(!QuantScheme::Fp8E4m3.has_scale());
889    }
890
891    /// Pins the shared GGUF dequant-kernel ids so they can NEVER drift out
892    /// from under the compiled MSL/CU/WGSL/GLSL `switch (scheme_id)` in any
893    /// backend. Every GPU backend delegates to `gpu_dequant_scheme_id`; if
894    /// someone reorders a variant here this test fails before a mis-decode
895    /// ships. Non-GGUF schemes must report `None`.
896    #[test]
897    fn gpu_dequant_scheme_id_is_stable() {
898        use QuantScheme::*;
899        assert_eq!(QuantScheme::GPU_DEQUANT_SCHEME_ID_PAIRS.len(), 26);
900        for &(scheme, id) in QuantScheme::GPU_DEQUANT_SCHEME_ID_PAIRS {
901            assert_eq!(scheme.gpu_dequant_scheme_id(), Some(id), "{scheme:?}");
902            assert_eq!(
903                QuantScheme::from_gpu_dequant_scheme_id(id),
904                Some(scheme),
905                "id {id}"
906            );
907        }
908        for scheme in [
909            Int8Block { block_size: 32 },
910            Int8BlockAsym { block_size: 32 },
911            Int4Block { block_size: 32 },
912            Fp8E4m3,
913            Fp8E5m2,
914            Nvfp4Block,
915        ] {
916            assert_eq!(scheme.gpu_dequant_scheme_id(), None, "{scheme:?}");
917        }
918        assert_eq!(QuantScheme::from_gpu_dequant_scheme_id(999), None);
919    }
920
921    #[test]
922    fn quant_map_lookup() {
923        let mut q = QuantMap::new();
924        let id = NodeId(7);
925        q.insert(id, QuantScheme::Int8Block { block_size: 32 });
926        assert_eq!(q.get(id), Some(QuantScheme::Int8Block { block_size: 32 }));
927        assert_eq!(q.get(NodeId(99)), None);
928    }
929
930    #[test]
931    fn custom_format_ieee_bias_and_width() {
932        // IEEE bias 2^(e-1)-1, matching every named format.
933        assert_eq!(ScaledFormat::custom(3, 0).fields(), (3, 0, 3)); // f4e3m0
934        assert_eq!(ScaledFormat::custom(4, 3).fields(), (4, 3, 7)); // e4m3 bias
935        assert_eq!(ScaledFormat::custom(5, 2).fields(), (5, 2, 15)); // e5m2 bias
936        assert_eq!(ScaledFormat::custom(2, 1).fields(), (2, 1, 1)); // e2m1 bias
937        assert_eq!(ScaledFormat::custom(3, 0).bit_width(), 4);
938        assert_eq!(ScaledFormat::custom(3, 4).bit_width(), 8);
939    }
940
941    #[test]
942    fn custom_format_parse_display_round_trip() {
943        let f: ScaledFormat = "f4e3m0".parse().unwrap();
944        assert_eq!(f, ScaledFormat::custom(3, 0));
945        assert_eq!(f.to_string(), "f4e3m0");
946        // A few more splits.
947        assert_eq!(
948            "f5e2m2".parse::<ScaledFormat>().unwrap().to_string(),
949            "f5e2m2"
950        );
951        assert_eq!(
952            "f8e3m4".parse::<ScaledFormat>().unwrap().to_string(),
953            "f8e3m4"
954        );
955        // Named formats parse to their exact hardware variants, not Custom.
956        assert_eq!(
957            "f8e4m3".parse::<ScaledFormat>().unwrap(),
958            ScaledFormat::F8E4M3
959        );
960        assert_eq!(
961            "f4e2m1".parse::<ScaledFormat>().unwrap(),
962            ScaledFormat::F4E2M1
963        );
964        assert_eq!(
965            "f8e5m2fnuz".parse::<ScaledFormat>().unwrap(),
966            ScaledFormat::F8E5M2Fnuz
967        );
968    }
969
970    #[test]
971    fn custom_format_parse_rejects_invalid() {
972        assert!("f4e3m1".parse::<ScaledFormat>().is_err()); // 1+3+1 != 4
973        assert!("f9e4m4".parse::<ScaledFormat>().is_err()); // 9 bits > byte
974        assert!("f1e0m0".parse::<ScaledFormat>().is_err()); // 0 exponent bits
975        assert!("e3m0".parse::<ScaledFormat>().is_err()); // missing 'f'
976        assert!("f4e3m0x".parse::<ScaledFormat>().is_err()); // trailing junk
977        assert!("garbage".parse::<ScaledFormat>().is_err());
978    }
979
980    #[test]
981    fn custom_gpu_word_is_packed_descriptor() {
982        // Named formats keep their small ids (top bit clear).
983        assert_eq!(ScaledFormat::F8E4M3.kernel_id(), 0);
984        assert_eq!(ScaledFormat::F4E2M1.kernel_id(), 6);
985        assert!(ScaledFormat::F4E2M1.kernel_id() & 0x8000_0000 == 0);
986        // Custom packs (exp, mant, bias) with the sentinel bit set.
987        let w = ScaledFormat::custom(3, 0).kernel_id(); // e=3, m=0, bias=3
988        assert_eq!(w & 0x8000_0000, 0x8000_0000);
989        assert_eq!(w & 0xF, 3); // exp_bits
990        assert_eq!((w >> 4) & 0xF, 0); // mant_bits
991        assert_eq!((w >> 8) & 0xFF, 3); // bias
992        // Negative bias round-trips through the u8 reinterpretation.
993        let wn = ScaledFormat::custom_with_bias(3, 0, -2).kernel_id();
994        assert_eq!(((wn >> 8) & 0xFF) as u8 as i8, -2);
995    }
996
997    #[test]
998    fn format_is_const_and_introspectable() {
999        // Usable in const context.
1000        const F: ScaledFormat = ScaledFormat::custom(3, 0);
1001        assert!(F.is_custom() && !F.is_named());
1002        assert_eq!((F.exp_bits(), F.mant_bits(), F.bias()), (3, 0, 3));
1003        assert!(ScaledFormat::F8E4M3.is_named() && !ScaledFormat::F8E4M3.is_custom());
1004        assert_eq!(ScaledFormat::NAMED.len(), 7);
1005        assert!(ScaledFormat::NAMED.iter().all(|f| f.is_named()));
1006    }
1007
1008    #[test]
1009    fn format_codec_methods_and_quantize() {
1010        // Method forms match the free functions for every named format.
1011        for f in ScaledFormat::NAMED {
1012            for c in 0..(1u16 << f.bit_width()) {
1013                assert_eq!(
1014                    f.decode(c as u8).to_bits(),
1015                    crate::lowp_codec::decode(f, c as u8).to_bits()
1016                );
1017            }
1018        }
1019        let cf = ScaledFormat::custom(3, 0); // f4e3m0 grid: ±{0.25,..,16}, ±0
1020        assert_eq!(cf.quantize(1.4), 1.0);
1021        assert_eq!(cf.quantize(1.6), 2.0);
1022        assert_eq!(cf.quantize(-3.0), -2.0); // nearer 2 than 4 in magnitude
1023        assert_eq!(cf.decode(cf.encode(100.0)), 16.0); // saturates
1024        assert_eq!(cf.decode(cf.encode(f32::INFINITY)), 16.0);
1025    }
1026
1027    #[test]
1028    fn representable_values_grid() {
1029        assert_eq!(
1030            ScaledFormat::custom(3, 0).representable_values(),
1031            vec![
1032                -16.0, -8.0, -4.0, -2.0, -1.0, -0.5, -0.25, 0.0, 0.25, 0.5, 1.0, 2.0, 4.0, 8.0,
1033                16.0
1034            ]
1035        );
1036        // Named formats: the grid size never exceeds the finite code count.
1037        for f in ScaledFormat::NAMED {
1038            let g = f.representable_values();
1039            assert!(!g.is_empty() && g.windows(2).all(|w| w[0] < w[1]));
1040        }
1041    }
1042
1043    #[test]
1044    fn enumerate_all_fnexmy() {
1045        // Every valid fNeXmY: exp >= 1, mant >= 0, 1 + exp + mant <= 8.
1046        let mut all = vec![];
1047        for exp in 1u8..=7 {
1048            for mant in 0u8..=(7 - exp) {
1049                all.push(ScaledFormat::custom(exp, mant));
1050            }
1051        }
1052        assert_eq!(all.len(), 28, "there are exactly 28 fNeXmY formats");
1053        eprintln!(
1054            "{:<8} {:>4} {:>3} {:>3} {:>4} {:>12} {:>12} {:>5}",
1055            "name", "bits", "e", "m", "bias", "max_finite", "min_pos", "vals"
1056        );
1057        for f in &all {
1058            let vals = f.representable_values();
1059            let min_pos = vals.iter().copied().find(|&x| x > 0.0).unwrap();
1060            assert!(f.bit_width() <= 8 && f.exp_bits() >= 1 && !vals.is_empty());
1061            eprintln!(
1062                "{:<8} {:>4} {:>3} {:>3} {:>4} {:>12.3e} {:>12.3e} {:>5}",
1063                f.to_string(),
1064                f.bit_width(),
1065                f.exp_bits(),
1066                f.mant_bits(),
1067                f.bias(),
1068                f.max_finite(),
1069                min_pos,
1070                vals.len()
1071            );
1072        }
1073    }
1074
1075    #[test]
1076    fn scale_layout_parse_round_trip() {
1077        assert_eq!(
1078            "per_tensor".parse::<ScaleLayout>().unwrap(),
1079            ScaleLayout::PerTensor
1080        );
1081        assert_eq!("mx".parse::<ScaleLayout>().unwrap(), ScaleLayout::mx());
1082        assert_eq!(
1083            "nvfp4".parse::<ScaleLayout>().unwrap(),
1084            ScaleLayout::nvfp4()
1085        );
1086        assert_eq!(
1087            "mx/64".parse::<ScaleLayout>().unwrap(),
1088            ScaleLayout::BlockMxE8M0 { block: 64 }
1089        );
1090        assert!("bogus".parse::<ScaleLayout>().is_err());
1091        // Display round-trips through FromStr.
1092        let mx = ScaleLayout::mx();
1093        assert_eq!(mx.to_string().parse::<ScaleLayout>().unwrap(), mx);
1094    }
1095}