Skip to main content

cherenkov/qwen4_exp/
pack.rs

1//! Pack qwen4-exp checkpoints into aligned files. MLX Q4 tensors are copied
2//! bit for bit; native BF16 tensors are converted as the files are written.
3//!
4//!   dense.bin    every tensor that is not an expert matrix, an n-gram
5//!                shard, or vision (64-byte aligned, name order)
6//!   experts.bin  `[layer][expert]` records (see `ExpertLayout`)
7//!   ngram.bin    `[row]` records of weight | scales | biases
8//!   manifest.json
9
10use super::{DenseEntry, ExpertLayout, Manifest, NgramLayout, PAGE};
11use crate::tensors::Dtype;
12use source::{Source, Tensor};
13
14mod affine;
15mod source;
16use crate::units::BYTES_PER_GB;
17use anyhow::{Context, Result, ensure};
18use std::fs::File;
19use std::io::{BufWriter, Write};
20use std::path::Path;
21use std::time::Instant;
22
23const DENSE_ALIGN: u64 = 64;
24
25/// Prepare selected precisions without loading the inference engine. Low-bit
26/// stores always derive from the aligned Q4 base, which is retained for reuse.
27pub fn prepare(model_dir: &Path, output: Option<&Path>, experts: &[u32]) -> Result<()> {
28    ensure!(
29        !experts.is_empty() && experts.iter().all(|bits| (2..=4).contains(bits)),
30        "select at least one expert precision: 4, 3 or 2"
31    );
32
33    let packed_input = model_dir.join("manifest.json").is_file();
34    let default_output = if packed_input {
35        model_dir.to_owned()
36    } else {
37        model_dir.join("packed")
38    };
39    let out = output.unwrap_or(&default_output);
40
41    if !out.join("manifest.json").exists() {
42        ensure!(
43            !packed_input,
44            "a packed input must use its existing store directory"
45        );
46        pack(model_dir, out)?;
47    }
48
49    let manifest = Manifest::load(out)?;
50    let source = model_dir.canonicalize()?;
51
52    // An explicit output must not silently select a different checkpoint.
53    ensure!(
54        out == default_output
55            || out.canonicalize()? == source
56            || Path::new(&manifest.source_dir).canonicalize().ok() == Some(source),
57        "the packed store at {} belongs to a different source model",
58        out.display()
59    );
60    copy_chat_metadata(model_dir, out)?;
61    eprintln!("4-bit base ready at {}", out.display());
62
63    let low_bits: Vec<_> = experts.iter().copied().filter(|&bits| bits != 4).collect();
64
65    super::lowbit::ensure_many(out, &manifest.experts, &low_bits)?;
66    eprintln!("requested expert stores ready");
67
68    Ok(())
69}
70
71/// Auxiliary files copied without replacing metadata already present in a store.
72pub(crate) const CHAT_METADATA_FILES: &[&str] = &[
73    "chat_template.jinja",
74    "tokenizer_config.json",
75    "generation_config.json",
76    "LICENSE",
77    "LICENSE.txt",
78    "LICENSE.md",
79    "NOTICE",
80    "NOTICE.txt",
81    "NOTICE.md",
82];
83
84/// Whether an available source has metadata that the prepared store lacks.
85pub(crate) fn missing_chat_metadata(source: &Path, target: &Path) -> bool {
86    CHAT_METADATA_FILES
87        .iter()
88        .any(|name| source.join(name).is_file() && !target.join(name).exists())
89}
90
91/// Older packed directories can acquire template metadata without repacking weights.
92pub(crate) fn copy_chat_metadata(model_dir: &Path, out_dir: &Path) -> Result<()> {
93    for name in CHAT_METADATA_FILES {
94        let source = model_dir.join(name);
95        let target = out_dir.join(name);
96
97        if !source.is_file() || target.exists() {
98            continue;
99        }
100
101        std::fs::copy(source, target)
102            .with_context(|| format!("copying {name} into the packed store"))?;
103    }
104
105    Ok(())
106}
107
108fn dtype_name(d: Dtype) -> &'static str {
109    match d {
110        Dtype::U32 => "U32",
111        Dtype::F32 => "F32",
112        Dtype::F16 => "F16",
113        Dtype::BF16 => "BF16",
114        Dtype::I64 => "I64",
115    }
116}
117
118enum Class {
119    Dense,
120    Expert,
121    Ngram,
122    Skip,
123}
124
125fn classify(name: &str) -> Class {
126    if name.starts_with("vision_tower") || name.contains(".visual.") {
127        Class::Skip
128    } else if name.contains(".mlp.switch_mlp.") {
129        Class::Expert
130    } else if name.contains(".ngram_embedding.shard") {
131        // MLX conversions name these "shards.N" or "shard_N".
132        Class::Ngram
133    } else {
134        Class::Dense
135    }
136}
137
138struct Progress {
139    start: Instant,
140    bytes: u64,
141    last: Instant,
142}
143
144impl Progress {
145    fn new() -> Self {
146        Self {
147            start: Instant::now(),
148            bytes: 0,
149            last: Instant::now(),
150        }
151    }
152
153    fn add(&mut self, n: u64, what: &str) {
154        self.bytes += n;
155
156        if self.last.elapsed().as_secs_f64() > 10.0 {
157            self.last = Instant::now();
158            let s = self.start.elapsed().as_secs_f64();
159
160            eprintln!(
161                "  {what}: {:.1} GB written, {:.2} GB/s, {:.0} s",
162                self.bytes as f64 / BYTES_PER_GB as f64,
163                self.bytes as f64 / BYTES_PER_GB as f64 / s,
164                s
165            );
166        }
167    }
168}
169
170fn pad_to(w: &mut BufWriter<File>, pos: &mut u64, align: u64) -> Result<()> {
171    let rem = *pos % align;
172
173    if rem != 0 {
174        let pad = (align - rem) as usize;
175
176        w.write_all(&vec![0u8; pad])?;
177
178        *pos += pad as u64;
179    }
180
181    Ok(())
182}
183
184pub fn pack(model_dir: &Path, out_dir: &Path) -> Result<()> {
185    ensure!(
186        !out_dir.join("manifest.json").exists(),
187        "a packed store already exists at {}; use it for inference",
188        out_dir.display()
189    );
190
191    let weights = Source::load(model_dir)?;
192
193    crate::storage::create_private_dir(out_dir)?;
194    ensure!(
195        weights.config.is_none() || model_dir.canonicalize()? != out_dir.canonicalize()?,
196        "BF16 import needs a separate output directory to preserve the source configuration"
197    );
198
199    // Estimate converted bytes, not BF16 source bytes. The space guard adds
200    // 2 GB for alignment and filesystem overhead.
201    crate::storage::require_space(out_dir, weights.estimated_bytes())?;
202
203    let mut names: Vec<&String> = weights.tensors.keys().collect();
204
205    names.sort();
206
207    let mut progress = Progress::new();
208    let dense = pack_dense(&weights, out_dir, &names, &mut progress)?;
209    let layout = pack_experts(&weights, out_dir, &mut progress)?;
210    let layout_n = pack_ngram(&weights, out_dir, &names, &mut progress)?;
211
212    // A custom --output directory is itself runnable, without finding its
213    // original source directory. Only small runtime metadata is copied.
214    if model_dir.canonicalize()? != out_dir.canonicalize()? {
215        for name in ["config.json", "tokenizer.json"] {
216            std::fs::copy(model_dir.join(name), out_dir.join(name))
217                .with_context(|| format!("copying {name} into the packed store"))?;
218        }
219
220        copy_chat_metadata(model_dir, out_dir)?;
221    }
222
223    if let Some(config) = &weights.config {
224        std::fs::write(
225            out_dir.join("config.json"),
226            serde_json::to_vec_pretty(config)?,
227        )?;
228    }
229
230    let manifest = Manifest {
231        version: 1,
232        source_dir: model_dir.display().to_string(),
233        dense,
234        experts: layout,
235        ngram: layout_n,
236    };
237
238    std::fs::write(
239        out_dir.join("manifest.json"),
240        serde_json::to_vec_pretty(&manifest)?,
241    )?;
242    eprintln!(
243        "packed into {} in {:.0} s ({:.1} GB)",
244        out_dir.display(),
245        progress.start.elapsed().as_secs_f64(),
246        progress.bytes as f64 / BYTES_PER_GB as f64
247    );
248
249    Ok(())
250}
251
252fn pack_dense(
253    weights: &Source,
254    out_dir: &Path,
255    names: &[&String],
256    progress: &mut Progress,
257) -> Result<Vec<DenseEntry>> {
258    // ---- dense ----
259    let mut dense = Vec::new();
260
261    {
262        let mut w = BufWriter::with_capacity(8 << 20, File::create(out_dir.join("dense.bin"))?);
263        let mut pos = 0u64;
264
265        for name in names {
266            if matches!(classify(name), Class::Dense) {
267                weights.prefetch(weights.tensor(name)?);
268            }
269        }
270
271        for name in names {
272            if !matches!(classify(name), Class::Dense) {
273                continue;
274            }
275
276            let t = weights.tensor(name)?;
277
278            pad_to(&mut w, &mut pos, DENSE_ALIGN)?;
279
280            weights.write(t, 0..t.nbytes, &mut w)?;
281            dense.push(DenseEntry {
282                name: (*name).clone(),
283                dtype: dtype_name(t.dtype).into(),
284                shape: t.shape.clone(),
285                offset: pos,
286                nbytes: t.nbytes as u64,
287            });
288
289            pos += t.nbytes as u64;
290
291            progress.add(t.nbytes as u64, "dense");
292        }
293
294        w.flush()?;
295        eprintln!(
296            "dense.bin: {} tensors, {:.2} GB",
297            dense.len(),
298            pos as f64 / BYTES_PER_GB as f64
299        );
300    }
301
302    Ok(dense)
303}
304
305fn pack_experts(weights: &Source, out_dir: &Path, progress: &mut Progress) -> Result<ExpertLayout> {
306    // ---- experts ----
307    let mut prefixes: Vec<String> = Vec::new();
308    let mut layer = 0;
309
310    while weights.tensors.contains_key(&format!(
311        "language_model.model.layers.{layer}.mlp.switch_mlp.gate_proj.weight"
312    )) {
313        prefixes.push(format!(
314            "language_model.model.layers.{layer}.mlp.switch_mlp"
315        ));
316
317        layer += 1;
318    }
319
320    let mut mtp = 0;
321
322    while weights
323        .tensors
324        .contains_key(&format!("mtp.layers.{mtp}.mlp.switch_mlp.gate_proj.weight"))
325    {
326        prefixes.push(format!("mtp.layers.{mtp}.mlp.switch_mlp"));
327
328        mtp += 1;
329    }
330
331    anyhow::ensure!(!prefixes.is_empty(), "no expert tensors found");
332
333    let gw = weights.tensor(&format!("{}.gate_proj.weight", prefixes[0]))?;
334    let gs = weights.tensor(&format!("{}.gate_proj.scales", prefixes[0]))?;
335    let dw = weights.tensor(&format!("{}.down_proj.weight", prefixes[0]))?;
336    let ds = weights.tensor(&format!("{}.down_proj.scales", prefixes[0]))?;
337    let experts = gw.shape[0];
338    let inter = gw.shape[1];
339    let hidden = gw.shape[2] * 8;
340
341    anyhow::ensure!(
342        dw.shape == vec![experts, hidden, inter / 8],
343        "down_proj shape {:?}",
344        dw.shape
345    );
346
347    let group = hidden / gs.shape[2];
348
349    anyhow::ensure!(group == inter / ds.shape[2], "group size mismatch");
350
351    let w_up = (inter * hidden / 2) as u64; // packed 4-bit bytes for gate/up
352    let w_down = (hidden * inter / 2) as u64;
353    let s_up = (inter * (hidden / group) * 2) as u64;
354    let s_down = (hidden * (inter / group) * 2) as u64;
355    let layout = {
356        let mut off = 0u64;
357        let mut take = |n: u64| {
358            let o = off;
359            off += n;
360
361            o
362        };
363        let gate_w = take(w_up);
364        let up_w = take(w_up);
365        let down_w = take(w_down);
366        let gate_s = take(s_up);
367        let gate_b = take(s_up);
368        let up_s = take(s_up);
369        let up_b = take(s_up);
370        let down_s = take(s_down);
371        let down_b = take(s_down);
372        let record_bytes = off;
373
374        ExpertLayout {
375            layers: prefixes.len(),
376            experts,
377            inter,
378            hidden,
379            group,
380            record_bytes,
381            record_stride: record_bytes.div_ceil(PAGE) * PAGE,
382            layer_prefixes: prefixes.clone(),
383            gate_w,
384            up_w,
385            down_w,
386            gate_s,
387            gate_b,
388            up_s,
389            up_b,
390            down_s,
391            down_b,
392        }
393    };
394
395    eprintln!(
396        "experts: {} layers x {} experts, record {} bytes (stride {}), {:.2} GB",
397        layout.layers,
398        experts,
399        layout.record_bytes,
400        layout.record_stride,
401        (layout.layers * experts) as f64 * layout.record_stride as f64 / BYTES_PER_GB as f64
402    );
403
404    {
405        let mut w = BufWriter::with_capacity(8 << 20, File::create(out_dir.join("experts.bin"))?);
406        let pad = vec![0u8; (layout.record_stride - layout.record_bytes) as usize];
407        let mut pos = 0u64;
408
409        const SUFFIXES: [&str; 9] = [
410            "gate_proj.weight",
411            "up_proj.weight",
412            "down_proj.weight",
413            "gate_proj.scales",
414            "gate_proj.biases",
415            "up_proj.scales",
416            "up_proj.biases",
417            "down_proj.scales",
418            "down_proj.biases",
419        ];
420
421        for s in SUFFIXES {
422            weights.prefetch(weights.tensor(&format!("{}.{s}", prefixes[0]))?);
423        }
424
425        for (pi, prefix) in prefixes.iter().enumerate() {
426            if let Some(next) = prefixes.get(pi + 1) {
427                for s in SUFFIXES {
428                    weights.prefetch(weights.tensor(&format!("{next}.{s}"))?);
429                }
430            }
431
432            let get = |suffix: &str| weights.tensor(&format!("{prefix}.{suffix}"));
433            let parts: [(&Tensor, u64); 9] = [
434                (get("gate_proj.weight")?, w_up),
435                (get("up_proj.weight")?, w_up),
436                (get("down_proj.weight")?, w_down),
437                (get("gate_proj.scales")?, s_up),
438                (get("gate_proj.biases")?, s_up),
439                (get("up_proj.scales")?, s_up),
440                (get("up_proj.biases")?, s_up),
441                (get("down_proj.scales")?, s_down),
442                (get("down_proj.biases")?, s_down),
443            ];
444
445            for (buf, per) in &parts {
446                anyhow::ensure!(
447                    buf.nbytes as u64 == *per * experts as u64,
448                    "{prefix}: tensor size {} != {} x {experts}",
449                    buf.nbytes,
450                    per
451                );
452            }
453
454            for e in 0..experts {
455                for (buf, per) in &parts {
456                    let start = (e as u64 * per) as usize;
457
458                    weights.write(buf, start..start + *per as usize, &mut w)?;
459                }
460
461                w.write_all(&pad)?;
462
463                pos += layout.record_stride;
464            }
465
466            progress.add(experts as u64 * layout.record_stride, prefix);
467        }
468
469        w.flush()?;
470        eprintln!("experts.bin: {:.2} GB", pos as f64 / BYTES_PER_GB as f64);
471    }
472
473    Ok(layout)
474}
475
476fn pack_ngram(
477    weights: &Source,
478    out_dir: &Path,
479    names: &[&String],
480    progress: &mut Progress,
481) -> Result<NgramLayout> {
482    // ---- n-gram ----
483    let ngram_prefix = names
484        .iter()
485        .find_map(|n| {
486            n.find(".ngram_embedding.")
487                .map(|i| n[..i + ".ngram_embedding".len()].to_string())
488        })
489        .context("n-gram shards not found")?;
490    let ple_prefix = ngram_prefix
491        .trim_end_matches(".ngram_embedding")
492        .to_string();
493    // Shard naming differs between converters: "shards.N" or "shard_N".
494    let shard = |s: usize, suffix: &str| -> String {
495        let dotted = format!("{ngram_prefix}.shards.{s}.{suffix}");
496
497        if weights.tensors.contains_key(&dotted) {
498            dotted
499        } else {
500            format!("{ngram_prefix}.shard_{s}.{suffix}")
501        }
502    };
503    let read_i64 = |name: &str| -> Result<Vec<i64>> {
504        let t = weights.tensor(name)?;
505
506        anyhow::ensure!(t.dtype == Dtype::I64, "{name}: expected I64");
507
508        Ok(weights
509            .bytes(t)?
510            .as_chunks::<8>()
511            .0
512            .iter()
513            .map(|&c| i64::from_le_bytes(c))
514            .collect())
515    };
516    let head_offsets = read_i64(&format!("{ple_prefix}.ngram_heads_offsets"))?;
517    let head_sizes = read_i64(&format!("{ple_prefix}.ngram_heads_vocab_sizes"))?;
518    let multipliers = read_i64(&format!("{ple_prefix}.layer_multipliers"))?;
519    let mut shards = 0;
520
521    while weights.tensors.contains_key(&shard(shards, "weight")) {
522        shards += 1;
523    }
524
525    anyhow::ensure!(shards > 0, "no n-gram shard tensors under {ngram_prefix}");
526
527    let s0w = weights.tensor(&shard(0, "weight"))?;
528    let s0s = weights.tensor(&shard(0, "scales"))?;
529    let rows_per_shard = s0w.shape[0];
530    let dim = s0w.shape[1] * 8;
531    let ngroup = dim / s0s.shape[1];
532    let weight_bytes = (dim / 2) as u64;
533    let scale_bytes = (dim / ngroup * 2) as u64;
534    let row_bytes = weight_bytes + 2 * scale_bytes;
535    let layout_n = NgramLayout {
536        rows: (shards * rows_per_shard) as u64,
537        row_bytes,
538        dim,
539        group: ngroup,
540        weight_bytes,
541        scale_bytes,
542        head_offsets: head_offsets.iter().map(|&v| v as u64).collect(),
543        head_vocab_sizes: head_sizes.iter().map(|&v| v as u64).collect(),
544        layer_multipliers: multipliers,
545    };
546
547    eprintln!(
548        "ngram: {shards} shards x {rows_per_shard} rows, dim {dim}, group {ngroup}, row {row_bytes} bytes, {:.2} GB",
549        layout_n.rows as f64 * row_bytes as f64 / BYTES_PER_GB as f64
550    );
551
552    {
553        let mut w = BufWriter::with_capacity(8 << 20, File::create(out_dir.join("ngram.bin"))?);
554
555        for part in ["weight", "scales", "biases"] {
556            weights.prefetch(weights.tensor(&shard(0, part))?);
557        }
558
559        for s in 0..shards {
560            if s + 1 < shards {
561                for part in ["weight", "scales", "biases"] {
562                    weights.prefetch(weights.tensor(&shard(s + 1, part))?);
563                }
564            }
565
566            let wt = weights.tensor(&shard(s, "weight"))?;
567            let st = weights.tensor(&shard(s, "scales"))?;
568            let bt = weights.tensor(&shard(s, "biases"))?;
569
570            anyhow::ensure!(wt.nbytes as u64 == rows_per_shard as u64 * weight_bytes);
571            anyhow::ensure!(st.nbytes as u64 == rows_per_shard as u64 * scale_bytes);
572            anyhow::ensure!(bt.nbytes == st.nbytes);
573
574            let (wb, sb) = (weight_bytes as usize, scale_bytes as usize);
575
576            for r in 0..rows_per_shard {
577                weights.write(wt, r * wb..(r + 1) * wb, &mut w)?;
578                weights.write(st, r * sb..(r + 1) * sb, &mut w)?;
579                weights.write(bt, r * sb..(r + 1) * sb, &mut w)?;
580            }
581
582            progress.add(
583                rows_per_shard as u64 * row_bytes,
584                &format!("ngram shard {s}"),
585            );
586        }
587
588        w.flush()?;
589    }
590
591    Ok(layout_n)
592}
593
594#[cfg(test)]
595#[path = "../../tests/unit/qwen4_exp/pack.rs"]
596mod tests;