1use 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
25pub 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 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
71pub(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
84pub(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
91pub(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 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 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 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 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 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; 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 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 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;