1use 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 pub experts_file: std::fs::File,
18 pub ngram: Mmap,
19}
20
21pub struct ExpertRef<'a> {
23 pub gate: QLinear<'a>,
24 pub up: QLinear<'a>,
25 pub down: QLinear<'a>,
26}
27
28#[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 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 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 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 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 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 pub fn expert_layer(&self, layer: usize) -> usize {
177 layer
178 }
179
180 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 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;