Skip to main content

cherenkov/qwen4_exp/
packed.rs

1//! Zero-copy views over the packed qwen4-exp layout (`pack.rs`).
2
3use super::{Manifest, Qwen4ExpConfig};
4use crate::quant::{GROUP_SIZE, QLinear};
5use anyhow::{Context, Result};
6use half::bf16;
7use memmap2::Mmap;
8use std::path::{Path, PathBuf};
9
10pub struct Packed {
11    pub cfg: Qwen4ExpConfig,
12    pub manifest: Manifest,
13    pub dir: PathBuf,
14    pub dense: Mmap,
15    pub experts: Mmap,
16    /// Open handle on experts.bin for advisory reads (prefetch).
17    pub experts_file: std::fs::File,
18    pub ngram: Mmap,
19}
20
21/// One routed expert's three projections, viewed inside its record.
22pub struct ExpertRef<'a> {
23    pub gate: QLinear<'a>,
24    pub up: QLinear<'a>,
25    pub down: QLinear<'a>,
26}
27
28/// Hash parameters read from a PLE layer's dense tensors.
29#[derive(Debug)]
30pub(super) struct NgramMetadata {
31    pub multipliers: Vec<i64>,
32    pub head_offsets: Vec<u64>,
33    pub head_sizes: Vec<u64>,
34}
35
36fn map(path: &Path) -> Result<Mmap> {
37    let file = std::fs::File::open(path).with_context(|| format!("opening {}", path.display()))?;
38
39    // Safety: packed files are read-only while the process runs.
40    unsafe { Mmap::map(&file) }.with_context(|| format!("mmap {}", path.display()))
41}
42
43fn as_u32(bytes: &[u8]) -> &[u32] {
44    assert_eq!(
45        bytes.as_ptr() as usize % 4,
46        0,
47        "u32 view must be 4-byte aligned"
48    );
49
50    unsafe { std::slice::from_raw_parts(bytes.as_ptr().cast::<u32>(), bytes.len() / 4) }
51}
52
53fn as_bf16(bytes: &[u8]) -> &[bf16] {
54    assert_eq!(
55        bytes.as_ptr() as usize % 2,
56        0,
57        "bf16 view must be 2-byte aligned"
58    );
59
60    unsafe { std::slice::from_raw_parts(bytes.as_ptr().cast::<bf16>(), bytes.len() / 2) }
61}
62
63impl Packed {
64    /// `model_dir` holds config.json and tokenizer.json; the packed files
65    /// live in `model_dir/packed`, or directly in a standalone packed directory.
66    pub fn open(model_dir: &Path) -> Result<Self> {
67        let dir = if model_dir.join("manifest.json").is_file() {
68            model_dir.to_owned()
69        } else {
70            model_dir.join("packed")
71        };
72        // Imported checkpoints may disable an absent MTP head in their packed
73        // config. Older stores can still use metadata beside the source weights.
74        let config_dir = if dir.join("config.json").is_file() {
75            &dir
76        } else {
77            model_dir
78        };
79        let cfg = Qwen4ExpConfig::load(config_dir)?;
80        let manifest = Manifest::load(&dir)?;
81
82        anyhow::ensure!(
83            manifest.experts.group == GROUP_SIZE,
84            "expert group size must be 64"
85        );
86
87        let experts_path = dir.join("experts.bin");
88
89        Ok(Packed {
90            cfg,
91            manifest,
92            dense: map(&dir.join("dense.bin"))?,
93            experts: map(&experts_path)?,
94            experts_file: std::fs::File::open(&experts_path)
95                .with_context(|| format!("opening {}", experts_path.display()))?,
96            ngram: map(&dir.join("ngram.bin"))?,
97            dir,
98        })
99    }
100
101    pub fn dense_bytes(&self, name: &str) -> Result<&[u8]> {
102        let e = self.manifest.dense(name)?;
103
104        Ok(&self.dense[e.offset as usize..(e.offset + e.nbytes) as usize])
105    }
106
107    pub fn shape(&self, name: &str) -> Result<&[usize]> {
108        Ok(&self.manifest.dense(name)?.shape)
109    }
110
111    pub fn bf16(&self, name: &str) -> Result<&[bf16]> {
112        let e = self.manifest.dense(name)?;
113
114        anyhow::ensure!(e.dtype == "BF16", "{name}: expected BF16, got {}", e.dtype);
115
116        Ok(as_bf16(self.dense_bytes(name)?))
117    }
118
119    pub fn i64s(&self, name: &str) -> Result<Vec<i64>> {
120        let e = self.manifest.dense(name)?;
121
122        anyhow::ensure!(e.dtype == "I64", "{name}: expected I64, got {}", e.dtype);
123
124        Ok(self
125            .dense_bytes(name)?
126            .as_chunks::<8>()
127            .0
128            .iter()
129            .map(|&c| i64::from_le_bytes(c))
130            .collect())
131    }
132
133    /// Read a PLE embedding's hash parameters from the dense store for either backend.
134    pub(super) fn ngram_metadata(&self, prefix: &str) -> Result<NgramMetadata> {
135        Ok(NgramMetadata {
136            multipliers: self.i64s(&format!("{prefix}.layer_multipliers"))?,
137            head_offsets: self
138                .i64s(&format!("{prefix}.ngram_heads_offsets"))?
139                .into_iter()
140                .map(|value| value as u64)
141                .collect(),
142            head_sizes: self
143                .i64s(&format!("{prefix}.ngram_heads_vocab_sizes"))?
144                .into_iter()
145                .map(|value| value as u64)
146                .collect(),
147        })
148    }
149
150    /// Affine 4-bit group-64 linear layer `{prefix}.{weight,scales,biases}`.
151    pub fn qlinear(&self, prefix: &str) -> Result<QLinear<'_>> {
152        let w = self.manifest.dense(&format!("{prefix}.weight"))?;
153
154        anyhow::ensure!(w.dtype == "U32", "{prefix}.weight: expected U32");
155
156        let out_dim = w.shape[0];
157        let in_dim = w.shape[1] * 8;
158        let s = self.manifest.dense(&format!("{prefix}.scales"))?;
159
160        anyhow::ensure!(
161            s.shape == vec![out_dim, in_dim / GROUP_SIZE],
162            "{prefix}.scales shape {:?} is not group 64",
163            s.shape
164        );
165
166        Ok(QLinear {
167            out_dim,
168            in_dim,
169            weight: as_u32(self.dense_bytes(&format!("{prefix}.weight"))?),
170            scales: self.bf16(&format!("{prefix}.scales"))?,
171            biases: self.bf16(&format!("{prefix}.biases"))?,
172        })
173    }
174
175    /// Record index for a main decoder layer's expert block.
176    pub fn expert_layer(&self, layer: usize) -> usize {
177        layer
178    }
179
180    /// Record index for MTP layer `i`'s expert block.
181    pub fn mtp_expert_layer(&self, i: usize) -> Result<usize> {
182        let name = format!("mtp.layers.{i}.mlp.switch_mlp");
183
184        self.manifest
185            .experts
186            .layer_prefixes
187            .iter()
188            .position(|p| *p == name)
189            .with_context(|| format!("no expert records for {name}"))
190    }
191
192    pub fn record_offset(&self, record_layer: usize, expert: usize) -> usize {
193        let l = &self.manifest.experts;
194
195        ((record_layer * l.experts + expert) as u64 * l.record_stride) as usize
196    }
197
198    pub fn expert(&self, record_layer: usize, expert: usize) -> ExpertRef<'_> {
199        let l = &self.manifest.experts;
200        let base = self.record_offset(record_layer, expert);
201        let rec = &self.experts[base..base + l.record_bytes as usize];
202        let sl = |off: u64, len: u64| &rec[off as usize..(off + len) as usize];
203        let w_up = (l.inter * l.hidden / 2) as u64;
204        let s_up = (l.inter * (l.hidden / l.group) * 2) as u64;
205        let s_down = (l.hidden * (l.inter / l.group) * 2) as u64;
206
207        ExpertRef {
208            gate: QLinear {
209                out_dim: l.inter,
210                in_dim: l.hidden,
211                weight: as_u32(sl(l.gate_w, w_up)),
212                scales: as_bf16(sl(l.gate_s, s_up)),
213                biases: as_bf16(sl(l.gate_b, s_up)),
214            },
215            up: QLinear {
216                out_dim: l.inter,
217                in_dim: l.hidden,
218                weight: as_u32(sl(l.up_w, w_up)),
219                scales: as_bf16(sl(l.up_s, s_up)),
220                biases: as_bf16(sl(l.up_b, s_up)),
221            },
222            down: QLinear {
223                out_dim: l.hidden,
224                in_dim: l.inter,
225                weight: as_u32(sl(l.down_w, w_up)),
226                scales: as_bf16(sl(l.down_s, s_down)),
227                biases: as_bf16(sl(l.down_b, s_down)),
228            },
229        }
230    }
231
232    /// Dequantize one hashed n-gram row using the manifest's width and group size.
233    pub fn ngram_row(&self, id: u64, dst: &mut [f32]) {
234        let n = &self.manifest.ngram;
235
236        debug_assert_eq!(dst.len(), n.dim);
237
238        let base = (id * n.row_bytes) as usize;
239        let rec = &self.ngram[base..base + n.row_bytes as usize];
240        let wb = n.weight_bytes as usize;
241        let sb = n.scale_bytes as usize;
242        let scales = &rec[wb..wb + sb];
243        let biases = &rec[wb + sb..wb + 2 * sb];
244        let g = n.group;
245
246        for (i, out) in dst.iter_mut().enumerate() {
247            let word = u32::from_le_bytes(rec[i / 8 * 4..i / 8 * 4 + 4].try_into().unwrap());
248            let q = (word >> (4 * (i % 8))) & 0xF;
249            let gi = i / g;
250            let s =
251                bf16::from_bits(u16::from_le_bytes([scales[2 * gi], scales[2 * gi + 1]])).to_f32();
252            let b =
253                bf16::from_bits(u16::from_le_bytes([biases[2 * gi], biases[2 * gi + 1]])).to_f32();
254            *out = s * q as f32 + b;
255        }
256    }
257}
258
259#[cfg(test)]
260#[path = "../../tests/unit/qwen4_exp/packed.rs"]
261mod tests;