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}