spark_model/
precision_schedule.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3#![allow(clippy::doc_lazy_continuation)]
4#![allow(clippy::doc_overindented_list_items)]
5
6//! Per-layer + per-tensor precision overrides (C.3, 2026-04-25).
7//!
8//! Reference: NVIDIA Transformer Engine 2.14 + EAQuant (arXiv:2506.13329)
9//! + community 2025 mixed-precision recipes. In MoE models, the
10//! quantization-sensitivity hierarchy holds across re-tested
11//! benchmarks:
12//!
13//!   1. **Router** (gate weights): hidden × num_experts, tiny in
14//!      memory, but routing accuracy collapses fast under quant.
15//!      Keep BF16 wherever feasible.
16//!   2. **LM head**: hidden × vocab, large but determines output
17//!      logit fidelity. BF16 closes the dominant chunk of
18//!      perplexity gap.
19//!   3. **First 1-2 transformer blocks** + **last 2-3 blocks**:
20//!      the embedding-adjacent layers carry sink-token outliers and
21//!      the output-adjacent layers shape final logits. Keep at FP8
22//!      (one tier above the bulk).
23//!   4. **Bulk MoE experts**: NVFP4 / FP8 — the model has the most
24//!      slack here.
25//!
26//! ## Scope
27//!
28//! This module ships:
29//!   - [`Role`] — semantic tag for each tensor the loader wants to
30//!     classify (router, lm_head, attention, expert, etc.).
31//!   - [`Dtype`] — target precision values the schedule emits.
32//!   - [`PrecisionSchedule`] — the per-(layer, role) → dtype
33//!     decision table, built from `[precision]` in MODEL.toml.
34//!
35//! The loader consults `schedule.dtype_for(layer_idx, role)` at
36//! tensor-load time and chooses the matching path. When the
37//! schedule is in its `default()` state (no `[precision]` block in
38//! MODEL.toml), every lookup returns `Dtype::Inherit` — meaning
39//! "use whatever the existing per-checkpoint logic decides." This
40//! keeps the pre-2026-04-25 behaviour bit-exact.
41
42use std::collections::BTreeSet;
43
44/// Semantic role of a tensor, used for precision lookups. The set is
45/// closed and minimal — adding a new role requires extending the
46/// `Dtype::for_role` match.
47#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
48pub enum Role {
49    /// MoE router gate (hidden × num_experts).
50    Router,
51    /// Final unembedding (hidden × vocab).
52    LmHead,
53    /// Token embedding (vocab × hidden).
54    Embedding,
55    /// Attention Q/K/V/O projection.
56    Attention,
57    /// Expert FFN (gate / up / down per expert).
58    Expert,
59    /// Shared-expert FFN, when present (DeepSeek-V3 / Qwen3.5 style).
60    SharedExpert,
61    /// Layer norm scales (RMSNorm `weight`).
62    Norm,
63}
64
65impl Role {
66    pub fn name(&self) -> &'static str {
67        match self {
68            Role::Router => "router",
69            Role::LmHead => "lm_head",
70            Role::Embedding => "embedding",
71            Role::Attention => "attention",
72            Role::Expert => "expert",
73            Role::SharedExpert => "shared_expert",
74            Role::Norm => "norm",
75        }
76    }
77
78    /// Parse a role tag from MODEL.toml. Returns `None` for unknown
79    /// names so the operator gets a load-time warning rather than a
80    /// silent miss.
81    #[allow(clippy::should_implement_trait)]
82    pub fn from_str(s: &str) -> Option<Self> {
83        match s {
84            "router" => Some(Role::Router),
85            "lm_head" => Some(Role::LmHead),
86            "embedding" => Some(Role::Embedding),
87            "attention" => Some(Role::Attention),
88            "expert" => Some(Role::Expert),
89            "shared_expert" => Some(Role::SharedExpert),
90            "norm" => Some(Role::Norm),
91            _ => None,
92        }
93    }
94}
95
96/// Target precision for a tensor. `Inherit` means "let the existing
97/// per-checkpoint logic decide" (preserves pre-C.3 behaviour); the
98/// other variants are hard requests.
99#[derive(Debug, Clone, Copy, PartialEq, Eq)]
100pub enum Dtype {
101    /// Honour the existing checkpoint-format detection and CLI flag.
102    /// Equivalent to "no override" — the loader falls through.
103    Inherit,
104    Bf16,
105    Fp8,
106    Nvfp4,
107}
108
109impl Dtype {
110    #[allow(clippy::should_implement_trait)]
111    pub fn from_str(s: &str) -> Option<Self> {
112        match s {
113            "inherit" => Some(Dtype::Inherit),
114            "bf16" => Some(Dtype::Bf16),
115            "fp8" => Some(Dtype::Fp8),
116            "nvfp4" => Some(Dtype::Nvfp4),
117            _ => None,
118        }
119    }
120
121    /// True iff the loader should override the inherited path.
122    pub fn is_override(&self) -> bool {
123        !matches!(self, Dtype::Inherit)
124    }
125}
126
127/// Per-(layer, role) precision schedule built from MODEL.toml's
128/// `[precision]` block. Lookups are O(1) for the common tables;
129/// per-layer overrides hit a small sorted set.
130#[derive(Debug, Clone)]
131pub struct PrecisionSchedule {
132    /// Default for any (layer, role) not specifically overridden.
133    /// Typically `Dtype::Inherit` so the existing path runs unchanged.
134    default: Dtype,
135    /// Role-specific defaults. `router_dtype = "bf16"` populates
136    /// `roles[Role::Router]`. Lookups fall back from per-layer
137    /// override → per-role default → global default.
138    router_dtype: Dtype,
139    lm_head_dtype: Dtype,
140    embedding_dtype: Dtype,
141    attention_dtype: Dtype,
142    expert_dtype: Dtype,
143    shared_expert_dtype: Dtype,
144    norm_dtype: Dtype,
145    /// Layer indices marked "sensitive" — typically the first 1-2
146    /// and last 2-3 transformer blocks. Tensors of role
147    /// `Attention`/`Expert` in these layers get `sensitive_dtype`.
148    sensitive_layers: BTreeSet<u16>,
149    sensitive_dtype: Dtype,
150}
151
152impl Default for PrecisionSchedule {
153    /// Empty schedule — every lookup returns `Inherit`. Bit-exact
154    /// equivalent to the pre-C.3 behaviour. MODEL.toml omits the
155    /// `[precision]` block to opt into this default.
156    fn default() -> Self {
157        Self {
158            default: Dtype::Inherit,
159            router_dtype: Dtype::Inherit,
160            lm_head_dtype: Dtype::Inherit,
161            embedding_dtype: Dtype::Inherit,
162            attention_dtype: Dtype::Inherit,
163            expert_dtype: Dtype::Inherit,
164            shared_expert_dtype: Dtype::Inherit,
165            norm_dtype: Dtype::Inherit,
166            sensitive_layers: BTreeSet::new(),
167            sensitive_dtype: Dtype::Inherit,
168        }
169    }
170}
171
172impl PrecisionSchedule {
173    /// Build from the four documented MODEL.toml fields:
174    ///   - `router_dtype`: dtype for the MoE gate
175    ///   - `lm_head_dtype`: dtype for the final unembedding
176    ///   - `sensitive_block_dtype` + `sensitive_block_indices`: the
177    ///     "extra precision" tier for the first/last few blocks
178    ///   - `default_dtype`: bulk fallback (typically Inherit)
179    ///
180    /// For now, the simpler `[precision]` schema only exposes these
181    /// four; per-tensor / per-layer YAML can extend later.
182    pub fn build(
183        router_dtype: Dtype,
184        lm_head_dtype: Dtype,
185        sensitive_block_indices: &[u16],
186        sensitive_block_dtype: Dtype,
187        default_dtype: Dtype,
188    ) -> Self {
189        Self {
190            default: default_dtype,
191            router_dtype,
192            lm_head_dtype,
193            embedding_dtype: Dtype::Inherit,
194            attention_dtype: Dtype::Inherit,
195            expert_dtype: Dtype::Inherit,
196            shared_expert_dtype: Dtype::Inherit,
197            norm_dtype: Dtype::Inherit,
198            sensitive_layers: sensitive_block_indices.iter().copied().collect(),
199            sensitive_dtype: sensitive_block_dtype,
200        }
201    }
202
203    /// Resolve the target dtype for a tensor. `layer_idx = None` is
204    /// used for non-layer tensors (embedding, lm_head, final norm).
205    /// Lookup order:
206    ///   1. Sensitive-layer override (only for Attention/Expert)
207    ///   2. Per-role default
208    ///   3. Global default
209    pub fn dtype_for(&self, layer_idx: Option<u16>, role: Role) -> Dtype {
210        // Sensitive-layer pass: applies to weight-bearing layer
211        // tensors only. Norms / embeddings / LM head are exempt
212        // (they have their own role-specific overrides).
213        if let Some(li) = layer_idx
214            && matches!(role, Role::Attention | Role::Expert | Role::SharedExpert)
215            && self.sensitive_layers.contains(&li)
216            && self.sensitive_dtype.is_override()
217        {
218            return self.sensitive_dtype;
219        }
220        let role_dtype = match role {
221            Role::Router => self.router_dtype,
222            Role::LmHead => self.lm_head_dtype,
223            Role::Embedding => self.embedding_dtype,
224            Role::Attention => self.attention_dtype,
225            Role::Expert => self.expert_dtype,
226            Role::SharedExpert => self.shared_expert_dtype,
227            Role::Norm => self.norm_dtype,
228        };
229        if role_dtype.is_override() {
230            role_dtype
231        } else {
232            self.default
233        }
234    }
235
236    /// True iff the schedule will produce any non-Inherit overrides.
237    /// Loaders can use this to skip the per-tensor lookups entirely
238    /// when no overrides are configured (default case).
239    pub fn has_any_override(&self) -> bool {
240        self.default.is_override()
241            || self.router_dtype.is_override()
242            || self.lm_head_dtype.is_override()
243            || self.embedding_dtype.is_override()
244            || self.attention_dtype.is_override()
245            || self.expert_dtype.is_override()
246            || self.shared_expert_dtype.is_override()
247            || self.norm_dtype.is_override()
248            || (self.sensitive_dtype.is_override() && !self.sensitive_layers.is_empty())
249    }
250}
251
252#[cfg(test)]
253mod tests {
254    use super::*;
255
256    #[test]
257    fn default_schedule_is_all_inherit() {
258        let s = PrecisionSchedule::default();
259        assert!(!s.has_any_override());
260        assert_eq!(s.dtype_for(None, Role::LmHead), Dtype::Inherit);
261        assert_eq!(s.dtype_for(Some(0), Role::Router), Dtype::Inherit);
262        assert_eq!(s.dtype_for(Some(38), Role::Expert), Dtype::Inherit);
263    }
264
265    #[test]
266    fn role_override_wins_over_default() {
267        let s = PrecisionSchedule::build(
268            Dtype::Bf16,    // router
269            Dtype::Bf16,    // lm head
270            &[],            // no sensitive layers
271            Dtype::Inherit, // sensitive dtype unused
272            Dtype::Nvfp4,   // bulk default
273        );
274        assert!(s.has_any_override());
275        assert_eq!(s.dtype_for(None, Role::Router), Dtype::Bf16);
276        assert_eq!(s.dtype_for(None, Role::LmHead), Dtype::Bf16);
277        assert_eq!(s.dtype_for(Some(15), Role::Expert), Dtype::Nvfp4);
278    }
279
280    #[test]
281    fn sensitive_layer_overrides_bulk_default_for_weights() {
282        let s = PrecisionSchedule::build(
283            Dtype::Bf16,
284            Dtype::Bf16,
285            &[0, 1, 38, 39],
286            Dtype::Fp8,
287            Dtype::Nvfp4,
288        );
289        // Layer 0 expert is sensitive → FP8
290        assert_eq!(s.dtype_for(Some(0), Role::Expert), Dtype::Fp8);
291        // Layer 38 attention is sensitive → FP8
292        assert_eq!(s.dtype_for(Some(38), Role::Attention), Dtype::Fp8);
293        // Shared experts are weight-bearing too.
294        assert_eq!(s.dtype_for(Some(1), Role::SharedExpert), Dtype::Fp8);
295        // Layer 5 expert is bulk → NVFP4
296        assert_eq!(s.dtype_for(Some(5), Role::Expert), Dtype::Nvfp4);
297
298        let sensitive_only = PrecisionSchedule::build(
299            Dtype::Inherit,
300            Dtype::Inherit,
301            &[7],
302            Dtype::Fp8,
303            Dtype::Inherit,
304        );
305        assert!(sensitive_only.has_any_override());
306        assert_eq!(
307            sensitive_only.dtype_for(Some(7), Role::Attention),
308            Dtype::Fp8
309        );
310    }
311
312    #[test]
313    fn sensitive_layer_does_not_override_router_or_lm_head() {
314        // Router is not Attention/Expert; sensitivity table never
315        // applies to it. Routing dtype is governed by router_dtype only.
316        let s = PrecisionSchedule::build(Dtype::Bf16, Dtype::Bf16, &[0], Dtype::Fp8, Dtype::Nvfp4);
317        assert_eq!(s.dtype_for(Some(0), Role::Router), Dtype::Bf16);
318        assert_eq!(s.dtype_for(Some(0), Role::LmHead), Dtype::Bf16);
319        assert_eq!(s.dtype_for(Some(0), Role::Embedding), Dtype::Nvfp4);
320        assert_eq!(s.dtype_for(Some(0), Role::Norm), Dtype::Nvfp4);
321    }
322
323    #[test]
324    fn role_str_round_trips() {
325        for r in [
326            Role::Router,
327            Role::LmHead,
328            Role::Embedding,
329            Role::Attention,
330            Role::Expert,
331            Role::SharedExpert,
332            Role::Norm,
333        ] {
334            assert_eq!(Role::from_str(r.name()), Some(r));
335        }
336        assert_eq!(Role::from_str("nonsense"), None);
337    }
338
339    #[test]
340    fn dtype_str_parsing() {
341        assert_eq!(Dtype::from_str("bf16"), Some(Dtype::Bf16));
342        assert_eq!(Dtype::from_str("fp8"), Some(Dtype::Fp8));
343        assert_eq!(Dtype::from_str("nvfp4"), Some(Dtype::Nvfp4));
344        assert_eq!(Dtype::from_str("inherit"), Some(Dtype::Inherit));
345        assert_eq!(Dtype::from_str("bogus"), None);
346        assert!(!Dtype::Inherit.is_override());
347        assert!(Dtype::Bf16.is_override());
348    }
349}