1#![allow(unused_imports)]
6
7use anyhow::{Context, Result, bail, ensure};
8use spark_runtime::gpu::{DevicePtr, GpuBackend};
9use spark_runtime::weights::{WeightDtype, WeightStore};
10
11use super::*;
12
13pub fn detect_nvfp4_variant(
28 store: &WeightStore,
29 config: &atlas_core::config::ModelConfig,
30) -> Nvfp4Variant {
31 if let Some(qc) = &config.quantization_config {
35 match qc.quant_method.as_str() {
36 "modelopt" if qc.quant_algo.eq_ignore_ascii_case("NVFP4") => {
37 return Nvfp4Variant::Standard;
38 }
39 "modelopt" if qc.quant_algo.eq_ignore_ascii_case("FP8") => {
40 return Nvfp4Variant::Fp8Dequanted;
41 }
42 "compressed-tensors" => {
43 let fmt = qc.format.to_ascii_lowercase();
49 if fmt.contains("fp8") || fmt.contains("float-quant") {
50 return Nvfp4Variant::Fp8Dequanted;
51 }
52 return Nvfp4Variant::CompressedTensors;
53 }
54 "fp8" => {
55 return Nvfp4Variant::Fp8Dequanted;
56 }
57 _ => {
58 }
62 }
63 }
64
65 let lp = config.layer_prefix(0);
66
67 let local_expert = config.local_expert_range().0;
69 let moe_sehyo_key = format!("{lp}.mlp.experts.{local_expert}.gate_proj.weight_packed");
70 if store.contains(&moe_sehyo_key) {
71 return Nvfp4Variant::CompressedTensors;
72 }
73
74 let dense_sehyo_key = format!("{lp}.mlp.gate_proj.weight_packed");
76 if store.contains(&dense_sehyo_key) {
77 return Nvfp4Variant::CompressedTensors;
78 }
79
80 let mistral_key = format!("layers.0.experts.{local_expert}.w1.weight_packed");
82 if store.contains(&mistral_key) {
83 return Nvfp4Variant::CompressedTensors;
84 }
85
86 if store.names().any(|k| k.ends_with(".weight_packed")) {
89 return Nvfp4Variant::CompressedTensors;
90 }
91
92 const ALT_LAYER0_PREFIX: &str = "model.language_model.layers.0";
107 let prefixes_to_check = [lp.clone(), ALT_LAYER0_PREFIX.to_string()];
108 for pfx in &prefixes_to_check {
109 let fp8_key = format!("{pfx}.mlp.experts.{local_expert}.gate_proj.weight_scale_inv");
110 if store.contains(&fp8_key) {
111 return Nvfp4Variant::Fp8Dequanted;
112 }
113 let fp8_dense_key = format!("{pfx}.mlp.gate_proj.weight_scale_inv");
114 if store.contains(&fp8_dense_key) {
115 return Nvfp4Variant::Fp8Dequanted;
116 }
117 let fp8_attn_key = format!("{pfx}.self_attn.q_proj.weight_scale_inv");
118 if store.contains(&fp8_attn_key) {
119 return Nvfp4Variant::Fp8Dequanted;
120 }
121 for key in [
131 format!("{pfx}.mlp.experts.{local_expert}.gate_proj.weight"),
132 format!("{pfx}.mlp.gate_proj.weight"),
133 format!("{pfx}.self_attn.q_proj.weight"),
134 ] {
135 if store
136 .get(&key)
137 .map(|w| w.dtype == WeightDtype::FP8E4M3)
138 .unwrap_or(false)
139 {
140 return Nvfp4Variant::Fp8Dequanted;
141 }
142 }
143 }
144 if store.names().any(|k| k.ends_with(".weight_scale_inv")) {
147 return Nvfp4Variant::Fp8Dequanted;
148 }
149
150 let any_standard_scale = store.names().any(|k| k.ends_with(".weight_scale"));
156 if !any_standard_scale {
157 tracing::warn!(
158 "No NVFP4/FP8 quantization metadata found (no .weight_packed / .weight_scale_inv / .weight_scale). \
159 Falling back to runtime BF16→NVFP4 quantization. Quality will be inferior to a calibrated NVFP4 release."
160 );
161 return Nvfp4Variant::Bf16Raw;
162 }
163
164 let has_mlp_scale = {
172 let k_dense = format!("{lp}.mlp.gate_proj.weight_scale");
173 let k_moe = format!("{lp}.mlp.experts.{local_expert}.gate_proj.weight_scale");
174 store.contains(&k_dense) || store.contains(&k_moe)
175 };
176 if !has_mlp_scale {
177 tracing::warn!(
178 "Partial NVFP4 metadata: `.weight_scale` exists for some tensors (e.g. KV scales) \
179 but not for MLP/MoE projections. Falling back to runtime BF16→NVFP4 quantization. \
180 For best quality use a fully-quantized NVFP4 release (e.g. Sehyo/*-NVFP4)."
181 );
182 return Nvfp4Variant::Bf16Raw;
183 }
184
185 Nvfp4Variant::Standard
186}
187
188pub(crate) fn quantized_auto(
193 store: &WeightStore,
194 prefix: &str,
195 gpu: &dyn GpuBackend,
196 variant: Nvfp4Variant,
197) -> Result<QuantizedWeight> {
198 match variant {
199 Nvfp4Variant::Standard => quantized(store, prefix, gpu),
200 Nvfp4Variant::CompressedTensors => quantized_v2(store, prefix, gpu),
201 Nvfp4Variant::Fp8Dequanted => {
202 unreachable!("Fp8Dequanted must use quantized_auto_fp8 with quant context")
203 }
204 Nvfp4Variant::Bf16Raw => {
205 unreachable!("Bf16Raw must use quantized_any with quant context")
206 }
207 }
208}
209
210#[derive(Clone, Copy)]
212pub(crate) struct QuantizeCtx {
213 pub absmax_k: spark_runtime::gpu::KernelHandle,
214 pub quantize_k: spark_runtime::gpu::KernelHandle,
215 pub stream: u64,
216}
217
218pub(crate) fn quantized_any(
221 store: &WeightStore,
222 prefix: &str,
223 n: usize,
224 k: usize,
225 gpu: &dyn GpuBackend,
226 variant: Nvfp4Variant,
227 qctx: QuantizeCtx,
228) -> Result<QuantizedWeight> {
229 let _t_detect = std::time::Instant::now();
230 let has_packed = store.contains(&format!("{prefix}.weight_packed"));
236 let has_scale = store.contains(&format!("{prefix}.weight_scale"));
237 let has_scale_inv = store.contains(&format!("{prefix}.weight_scale_inv"));
238 let has_only_dense =
239 !has_packed && !has_scale && !has_scale_inv && store.contains(&format!("{prefix}.weight"));
240
241 let has_fp8_dense = !has_packed
256 && !store.contains(&format!("{prefix}.weight_global_scale"))
257 && !store.contains(&format!("{prefix}.weight_scale_2"))
258 && (has_scale || has_scale_inv)
259 && store
260 .get(&format!("{prefix}.weight"))
261 .map(|w| w.dtype == WeightDtype::FP8E4M3)
262 .unwrap_or(false);
263
264 let effective_variant = if has_only_dense && !matches!(variant, Nvfp4Variant::Bf16Raw) {
265 tracing::debug!("{prefix}: no quantization metadata; falling back to runtime BF16→NVFP4");
266 Nvfp4Variant::Bf16Raw
267 } else if has_fp8_dense
268 && !matches!(variant, Nvfp4Variant::Fp8Dequanted | Nvfp4Variant::Bf16Raw)
269 {
270 tracing::debug!("{prefix}: FP8 key in an NVFP4 checkpoint; dequant FP8→BF16→NVFP4");
271 Nvfp4Variant::Fp8Dequanted
272 } else {
273 variant
274 };
275
276 let _t_detect_ns = _t_detect.elapsed().as_nanos() as u64;
277 match effective_variant {
278 Nvfp4Variant::Standard => quantized(store, prefix, gpu),
279 Nvfp4Variant::CompressedTensors => quantized_v2(store, prefix, gpu),
280 Nvfp4Variant::Fp8Dequanted => quantized_from_fp8(
281 store,
282 prefix,
283 n,
284 k,
285 gpu,
286 qctx.absmax_k,
287 qctx.quantize_k,
288 qctx.stream,
289 ),
290 Nvfp4Variant::Bf16Raw => {
291 use std::sync::atomic::{AtomicU64, Ordering};
292 static T_DETECT: AtomicU64 = AtomicU64::new(0);
293 static T_GET: AtomicU64 = AtomicU64::new(0);
294 static T_QUANT: AtomicU64 = AtomicU64::new(0);
295 static T_FREE: AtomicU64 = AtomicU64::new(0);
296 static N: AtomicU64 = AtomicU64::new(0);
297 T_DETECT.fetch_add(_t_detect_ns, Ordering::Relaxed);
298 let _t = std::time::Instant::now();
300 let w = store.get(&format!("{prefix}.weight"))?;
301 let bf16 = DenseWeight { weight: w.ptr };
302 T_GET.fetch_add(_t.elapsed().as_nanos() as u64, Ordering::Relaxed);
303 let _t = std::time::Instant::now();
304 let q = quantize_to_nvfp4(
305 &bf16,
306 n,
307 k,
308 gpu,
309 qctx.absmax_k,
310 qctx.quantize_k,
311 qctx.stream,
312 )?;
313 T_QUANT.fetch_add(_t.elapsed().as_nanos() as u64, Ordering::Relaxed);
314 let _t = std::time::Instant::now();
315 gpu.free(w.ptr)?;
322 T_FREE.fetch_add(_t.elapsed().as_nanos() as u64, Ordering::Relaxed);
323 let c = N.fetch_add(1, Ordering::Relaxed) + 1;
324 if c.is_multiple_of(512) {
325 let ms = |a: &AtomicU64| a.load(Ordering::Relaxed) as f64 / 1.0e6;
326 tracing::info!(
327 "quantized_any(Bf16Raw) PROFILE after {c} calls (ms total): detect={:.1} \
328 store_get={:.1} quantize={:.1} free={:.1} | sum={:.1} per_call={:.3}ms",
329 ms(&T_DETECT),
330 ms(&T_GET),
331 ms(&T_QUANT),
332 ms(&T_FREE),
333 ms(&T_DETECT) + ms(&T_GET) + ms(&T_QUANT) + ms(&T_FREE),
334 (ms(&T_DETECT) + ms(&T_GET) + ms(&T_QUANT) + ms(&T_FREE)) / c as f64,
335 );
336 }
337 Ok(q)
338 }
339 }
340}
341
342pub(crate) fn quantized_from_fp8(
346 store: &WeightStore,
347 prefix: &str,
348 n: usize,
349 k: usize,
350 gpu: &dyn GpuBackend,
351 absmax_k: spark_runtime::gpu::KernelHandle,
352 quantize_k: spark_runtime::gpu::KernelHandle,
353 stream: u64,
354) -> Result<QuantizedWeight> {
355 let bf16 = dequant_fp8_blockscaled_to_bf16(store, prefix, gpu)?;
356 let result = quantize_to_nvfp4(&bf16, n, k, gpu, absmax_k, quantize_k, stream)?;
357 gpu.free(bf16.weight)?;
359 Ok(result)
360}
361
362#[allow(dead_code)]
368pub(crate) fn dense_from_fp8(
369 store: &WeightStore,
370 prefix: &str,
371 gpu: &dyn GpuBackend,
372) -> Result<DenseWeight> {
373 dequant_fp8_blockscaled_to_bf16(store, prefix, gpu)
374}
375
376#[allow(dead_code)]
378pub(crate) fn load_attention_qwen35(
379 store: &WeightStore,
380 layer_prefix: &str,
381 gpu: &dyn GpuBackend,
382) -> Result<AttentionWeights> {
383 let p = format!("{layer_prefix}.self_attn");
384 let (k_scale, v_scale) = load_kv_scales(store, &p, gpu);
385 Ok(AttentionWeights {
386 q_proj: dense(store, &format!("{p}.q_proj.weight_packed"))?,
389 k_proj: dense(store, &format!("{p}.k_proj.weight_packed"))?,
390 v_proj: dense(store, &format!("{p}.v_proj.weight_packed"))?,
391 o_proj: quantized_v2(store, &format!("{p}.o_proj"), gpu)?,
392 q_norm: dense(store, &format!("{p}.q_norm.weight"))?,
393 k_norm: dense(store, &format!("{p}.k_norm.weight"))?,
394 q_norm_full: None,
395 k_norm_full: None,
396 k_scale,
397 v_scale,
398 })
399}
400
401#[allow(dead_code)]
403pub(crate) fn load_quantized_proj_qwen35(
404 store: &WeightStore,
405 prefix: &str,
406 gpu: &dyn GpuBackend,
407) -> Result<QuantizedWeight> {
408 quantized_v2(store, prefix, gpu)
409}
410
411#[cfg(test)]
412mod ep_detection_tests {
413 use super::*;
414 use atlas_core::config::ModelConfig;
415 use spark_runtime::weights::WeightStore;
416
417 fn store_with(names: &[String]) -> WeightStore {
420 use std::collections::HashMap;
421 let map: HashMap<String, spark_runtime::weights::WeightTensor> = names
422 .iter()
423 .map(|n| {
424 (
425 n.clone(),
426 spark_runtime::weights::WeightTensor {
427 ptr: spark_runtime::gpu::DevicePtr::NULL,
428 shape: vec![1],
429 dtype: spark_runtime::weights::WeightDtype::FP8E4M3,
430 },
431 )
432 })
433 .collect();
434 WeightStore::from_map(map)
435 }
436
437 #[test]
438 fn alternate_layer0_fp8_dtype_is_detected_on_every_ep_rank() {
439 let mut cfg = ModelConfig::qwen3_next_80b_nvfp4();
440 cfg.quantization_config = None;
441 let store =
442 store_with(&["model.language_model.layers.0.self_attn.q_proj.weight".to_string()]);
443
444 cfg.ep_world_size = 2;
445 for ep_rank in 0..2 {
446 cfg.ep_rank = ep_rank;
447 assert_eq!(
448 detect_nvfp4_variant(&store, &cfg),
449 Nvfp4Variant::Fp8Dequanted,
450 "EP rank {ep_rank} must inspect the same layer-zero checkpoint marker"
451 );
452 }
453 }
454
455 #[test]
456 fn scale_inv_suffix_fallback_detects_an_unexpected_prefix() {
457 let mut cfg = ModelConfig::qwen3_next_80b_nvfp4();
458 cfg.quantization_config = None;
459 let store =
460 store_with(&["third_party.transformer.blocks.17.attn.q.weight_scale_inv".to_string()]);
461 assert_eq!(
462 detect_nvfp4_variant(&store, &cfg),
463 Nvfp4Variant::Fp8Dequanted
464 );
465 }
466}