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 struct ModelWeights {
15 pub embed_tokens: DenseWeight,
17 pub final_norm: DenseWeight,
19 pub lm_head: DenseWeight,
21 pub layers: Vec<LayerWeights>,
23}
24
25pub(crate) fn ptr(store: &WeightStore, name: &str) -> Result<DevicePtr> {
27 Ok(store.get(name)?.ptr)
28}
29
30pub(crate) fn scalar_f32(store: &WeightStore, name: &str, gpu: &dyn GpuBackend) -> Result<f32> {
32 let w = store.get(name)?;
33 ensure!(
34 w.dtype == WeightDtype::FP32,
35 "Expected FP32 for {name}, got {:?}",
36 w.dtype
37 );
38 ensure!(
39 w.num_elements() == 1,
40 "Expected scalar for {name}, got {} elements",
41 w.num_elements()
42 );
43 let mut buf = [0u8; 4];
44 gpu.copy_d2h(w.ptr, &mut buf)?;
45 Ok(f32::from_le_bytes(buf))
46}
47
48pub(crate) fn load_kv_scales(
58 store: &WeightStore,
59 attn_prefix: &str,
60 gpu: &dyn GpuBackend,
61) -> (f32, f32) {
62 let k_key = format!("{attn_prefix}.k_proj.k_scale");
63 let v_key = format!("{attn_prefix}.v_proj.v_scale");
64
65 let k_scale = if store.contains(&k_key) {
66 match scalar_f32(store, &k_key, gpu) {
67 Ok(v) => {
68 tracing::debug!("Loaded k_scale={v:.6} from {k_key}");
69 v
70 }
71 Err(e) => {
72 tracing::warn!("Failed to load {k_key}: {e:#}, using 1.0");
73 1.0
74 }
75 }
76 } else {
77 tracing::debug!("No {k_key} in checkpoint, using k_scale=1.0");
78 1.0
79 };
80
81 let v_scale = if store.contains(&v_key) {
82 match scalar_f32(store, &v_key, gpu) {
83 Ok(v) => {
84 tracing::debug!("Loaded v_scale={v:.6} from {v_key}");
85 v
86 }
87 Err(e) => {
88 tracing::warn!("Failed to load {v_key}: {e:#}, using 1.0");
89 1.0
90 }
91 }
92 } else {
93 tracing::debug!("No {v_key} in checkpoint, using v_scale=1.0");
94 1.0
95 };
96
97 (k_scale, v_scale)
98}
99
100pub(crate) fn quantized(
104 store: &WeightStore,
105 prefix: &str,
106 gpu: &dyn GpuBackend,
107) -> Result<QuantizedWeight> {
108 let input_scale_key = format!("{prefix}.input_scale");
109 Ok(QuantizedWeight {
110 weight: ptr(store, &format!("{prefix}.weight"))?,
111 weight_scale: ptr(store, &format!("{prefix}.weight_scale"))?,
112 weight_scale_2: scalar_f32(store, &format!("{prefix}.weight_scale_2"), gpu)?,
113 input_scale: if store.contains(&input_scale_key) {
114 ptr(store, &input_scale_key)?
115 } else {
116 DevicePtr::NULL
117 },
118 weight_scale_2_vec: DevicePtr::NULL,
119 })
120}
121
122pub(crate) fn quantized_mxfp4_e8m0(store: &WeightStore, prefix: &str) -> Result<QuantizedWeight> {
137 let w = store.get(&format!("{prefix}.weight"))?;
138 let n = w.shape[0];
139 let k_packed = w.shape[1];
140 let total_nibbles = n * k_packed * 2;
141 let scale_t = store.get(&format!("{prefix}.scale"))?;
142 let num_groups = scale_t.num_elements();
143 ensure!(
144 num_groups > 0 && total_nibbles.is_multiple_of(num_groups),
145 "{prefix}: MXFP4 weight nibbles {total_nibbles} not divisible by E8M0 scale groups {num_groups}"
146 );
147 let block = total_nibbles / num_groups;
148 ensure!(
149 block == 32,
150 "{prefix}: native MXFP4 expects GROUP_SIZE=32, inferred {block} (scale groups {num_groups}) \
151 — refusing to land a non-MX checkpoint on the transcode-free path"
152 );
153 Ok(QuantizedWeight {
154 weight: ptr(store, &format!("{prefix}.weight"))?,
155 weight_scale: ptr(store, &format!("{prefix}.scale"))?,
156 weight_scale_2: 1.0, input_scale: DevicePtr::NULL,
158 weight_scale_2_vec: DevicePtr::NULL,
161 })
162}
163
164pub(crate) fn dense(store: &WeightStore, name: &str) -> Result<DenseWeight> {
165 let w = store.get(name)?;
166 Ok(DenseWeight { weight: w.ptr })
167}
168
169pub(crate) fn dense_auto_fp8_or_bf16(
191 store: &WeightStore,
192 prefix: &str,
193 gpu: &dyn GpuBackend,
194) -> Result<DenseWeight> {
195 let w = store.get(&format!("{prefix}.weight"))?;
196 match w.dtype {
197 WeightDtype::BF16 => Ok(DenseWeight { weight: w.ptr }),
198 WeightDtype::FP8E4M3 => dequant_fp8_blockscaled_to_bf16(store, prefix, gpu),
199 other => anyhow::bail!(
200 "dense_auto_fp8_or_bf16: unsupported dtype {:?} for {prefix}.weight",
201 other
202 ),
203 }
204}
205
206pub(crate) fn dense_f32_safe(
209 store: &WeightStore,
210 name: &str,
211 gpu: &dyn GpuBackend,
212) -> Result<DenseWeight> {
213 let w = store.get(name)?;
214 if w.dtype == WeightDtype::FP32 {
215 let n = w.num_elements();
223 let ptr = gpu.alloc(n * 2)?;
224 let trunc = gpu.kernel("quantize_nvfp4", "f32_to_bf16_trunc")?;
225 let blocks = (n.div_ceil(256) as u32).max(1);
226 spark_runtime::kernel_args::KernelLaunch::new(gpu, trunc)
227 .grid([blocks, 1, 1])
228 .block([256, 1, 1])
229 .arg_ptr(w.ptr)
230 .arg_ptr(ptr)
231 .arg_u32(n as u32)
232 .launch(gpu.default_stream())?;
233 Ok(DenseWeight { weight: ptr })
234 } else {
235 Ok(DenseWeight { weight: w.ptr })
236 }
237}
238
239pub(crate) fn dense_keep_f32(
249 store: &WeightStore,
250 name: &str,
251 gpu: &dyn GpuBackend,
252) -> Result<DenseWeight> {
253 let w = store.get(name)?;
254 match w.dtype {
255 WeightDtype::FP32 => {
256 Ok(DenseWeight { weight: w.ptr })
258 }
259 WeightDtype::BF16 => {
260 tracing::info!(
262 "dense_keep_f32: promoting {name} from BF16 to FP32 ({:?})",
263 w.shape
264 );
265 let n = w.num_elements();
266 let mut bf16_buf = vec![0u8; n * 2];
267 gpu.copy_d2h(w.ptr, &mut bf16_buf)?;
268 let f32_buf: Vec<u8> = bf16_buf
269 .chunks_exact(2)
270 .flat_map(|c| {
271 let bits = u16::from_le_bytes([c[0], c[1]]);
272 let f32_bits = (bits as u32) << 16;
273 f32_bits.to_le_bytes()
274 })
275 .collect();
276 let ptr = gpu.alloc(f32_buf.len())?;
277 gpu.copy_h2d(&f32_buf, ptr)?;
278 Ok(DenseWeight { weight: ptr })
279 }
280 other => {
281 bail!("dense_keep_f32: unsupported dtype {:?} for {name}", other);
282 }
283 }
284}
285
286pub(crate) fn dense_bf16_as_f32(
291 store: &WeightStore,
292 name: &str,
293 gpu: &dyn GpuBackend,
294) -> Result<DenseWeight> {
295 let w = store.get(name)?;
296 ensure!(
297 w.dtype == WeightDtype::BF16,
298 "Expected BF16 for {name}, got {:?}",
299 w.dtype
300 );
301 let n = w.num_elements();
302 let mut bf16_buf = vec![0u8; n * 2];
303 gpu.copy_d2h(w.ptr, &mut bf16_buf)?;
304 let f32_buf: Vec<u8> = bf16_buf
305 .chunks_exact(2)
306 .flat_map(|c| {
307 let bits = u16::from_le_bytes([c[0], c[1]]);
308 let f32_bits = (bits as u32) << 16;
309 f32_bits.to_le_bytes()
310 })
311 .collect();
312 let ptr = gpu.alloc(f32_buf.len())?;
313 gpu.copy_h2d(&f32_buf, ptr)?;
314 Ok(DenseWeight { weight: ptr })
315}
316
317pub(crate) fn dense_f32_as_bf16(
321 store: &WeightStore,
322 name: &str,
323 gpu: &dyn GpuBackend,
324) -> Result<DenseWeight> {
325 let w = store.get(name)?;
326 ensure!(
327 w.dtype == WeightDtype::FP32,
328 "Expected FP32 for {name}, got {:?}",
329 w.dtype
330 );
331 let n = w.num_elements();
332 let mut f32_buf = vec![0u8; n * 4];
333 gpu.copy_d2h(w.ptr, &mut f32_buf)?;
334 let bf16_buf: Vec<u8> = f32_buf
335 .chunks_exact(4)
336 .flat_map(|c| {
337 let val = f32::from_le_bytes([c[0], c[1], c[2], c[3]]);
338 f32_to_bf16(val).to_le_bytes()
339 })
340 .collect();
341 let ptr = gpu.alloc(bf16_buf.len())?;
342 gpu.copy_h2d(&bf16_buf, ptr)?;
343 Ok(DenseWeight { weight: ptr })
344}