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}