1use super::*;
7use std::fs::{File, OpenOptions};
8use atlas_tier::pio;
11use std::path::{Path, PathBuf};
12
13pub struct ExpertFileWriter {
18 dir: PathBuf,
19 index: ExpertIndex,
20 spec: ExpertRecordSpec,
21 layout: ExpertLayout,
22 files: Vec<File>,
23}
24
25impl ExpertFileWriter {
26 pub fn create(dir: &Path, index: ExpertIndex) -> Result<Self> {
27 std::fs::create_dir_all(dir).with_context(|| format!("mkdir {}", dir.display()))?;
28 let spec = index.spec();
29 let layout = index.layout();
30 let mut files = Vec::with_capacity(index.num_moe_layers as usize);
31 for l in 0..index.num_moe_layers {
32 let p = dir.join(index.file_name(l));
33 let f = OpenOptions::new()
34 .read(true)
35 .write(true)
36 .create(true)
37 .truncate(true)
38 .open(&p)
39 .with_context(|| format!("create {}", p.display()))?;
40 f.set_len(layout.bytes_per_layer())
41 .with_context(|| format!("set_len {}", p.display()))?;
42 files.push(f);
43 }
44 Ok(Self {
45 dir: dir.to_path_buf(),
46 index,
47 spec,
48 layout,
49 files,
50 })
51 }
52
53 pub fn spec(&self) -> &ExpertRecordSpec {
54 &self.spec
55 }
56
57 pub fn write_record(
59 &self,
60 key: ExpertKey,
61 header: &ExpertRecordHeader,
62 projs: &[ProjData; 3],
63 ) -> Result<()> {
64 if key.layer >= self.index.num_moe_layers {
65 bail!("layer {} out of range", key.layer);
66 }
67 if key.expert >= self.index.num_experts {
68 bail!("expert {} out of range", key.expert);
69 }
70 let rec = pack_record(&self.spec, self.layout.record_stride, header, projs)?;
71 let off = self.layout.file_offset(key);
72 pio::write_all_at(&self.files[key.layer as usize], &rec, off)
73 .with_context(|| format!("write record {:?} at {off}", key))?;
74 Ok(())
75 }
76
77 pub fn finish(self) -> Result<()> {
79 for f in &self.files {
80 f.sync_all().context("fsync layer file")?;
81 }
82 let p = self.dir.join(ExpertIndex::MANIFEST_NAME);
83 let json = serde_json::to_string_pretty(&self.index)?;
84 std::fs::write(&p, json).with_context(|| format!("write {}", p.display()))?;
85 Ok(())
86 }
87}
88
89pub struct ExpertFileReader {
93 index: ExpertIndex,
94 spec: ExpertRecordSpec,
95 layout: ExpertLayout,
96 files: Vec<File>,
97}
98
99impl ExpertFileReader {
100 pub fn open(dir: &Path) -> Result<Self> {
101 let mp = dir.join(ExpertIndex::MANIFEST_NAME);
102 let json =
103 std::fs::read_to_string(&mp).with_context(|| format!("read {}", mp.display()))?;
104 let index: ExpertIndex =
105 serde_json::from_str(&json).with_context(|| format!("parse {}", mp.display()))?;
106 if index.version != ExpertRecordHeader::VERSION {
107 bail!(
108 "manifest version {} != supported {}",
109 index.version,
110 ExpertRecordHeader::VERSION
111 );
112 }
113 let spec = index.spec();
114 let layout = index.layout();
115 let mut files = Vec::with_capacity(index.num_moe_layers as usize);
116 for l in 0..index.num_moe_layers {
117 let p = dir.join(index.file_name(l));
118 files.push(File::open(&p).with_context(|| format!("open {}", p.display()))?);
119 }
120 Ok(Self {
121 index,
122 spec,
123 layout,
124 files,
125 })
126 }
127
128 pub fn index(&self) -> &ExpertIndex {
129 &self.index
130 }
131 pub fn spec(&self) -> &ExpertRecordSpec {
132 &self.spec
133 }
134
135 pub fn read_record_raw(&self, key: ExpertKey) -> Result<Vec<u8>> {
137 if key.layer as usize >= self.files.len() {
140 bail!(
141 "ExpertFileReader: layer {} out of range ({} layer files)",
142 key.layer,
143 self.files.len()
144 );
145 }
146 let mut buf = vec![0u8; self.layout.record_stride as usize];
147 let off = self.layout.file_offset(key);
148 pio::read_exact_at(&self.files[key.layer as usize], &mut buf, off)
149 .with_context(|| format!("read record {:?} at {off}", key))?;
150 Ok(buf)
151 }
152}
153
154#[cfg(test)]
155mod fs_tests {
156 use super::*;
157
158 fn tmpdir(tag: &str) -> PathBuf {
159 let p = std::env::temp_dir().join(format!(
160 "atlas-xpr-{}-{}-{}",
161 tag,
162 std::process::id(),
163 std::time::SystemTime::now()
165 .duration_since(std::time::UNIX_EPOCH)
166 .map(|d| d.as_nanos())
167 .unwrap_or(0)
168 ));
169 std::fs::create_dir_all(&p).unwrap();
170 p
171 }
172
173 fn synth_index() -> ExpertIndex {
175 ExpertIndex::new(64, 128, 16, 256, 4096, vec![0, 1], 3)
177 }
178
179 fn synth_projs(spec: &ExpertRecordSpec, seed: u8) -> [Vec<(Vec<u8>, Vec<u8>)>; 1] {
180 let mut out = Vec::new();
181 for p in Proj::ALL {
182 let pb = spec.proj_bytes(p);
183 let packed: Vec<u8> = (0..pb.packed_bytes)
184 .map(|i| (i as u8).wrapping_add(seed).wrapping_add(p as u8))
185 .collect();
186 let scale: Vec<u8> = (0..pb.scale_bytes)
187 .map(|i| (i as u8).wrapping_mul(3).wrapping_add(seed))
188 .collect();
189 out.push((packed, scale));
190 }
191 [out]
192 }
193
194 #[test]
195 fn write_then_read_round_trips_bit_identical() {
196 let dir = tmpdir("rt");
197 let index = synth_index();
198 let spec = index.spec();
199
200 let mut expected = std::collections::HashMap::new();
202 {
203 let w = ExpertFileWriter::create(&dir, index.clone()).unwrap();
204 for layer in 0..index.num_moe_layers {
205 for expert in 0..index.num_experts {
206 let seed = (layer as u8) << 4 | expert as u8;
207 let raw = synth_projs(&spec, seed);
208 let projs = [
209 ProjData {
210 packed: &raw[0][0].0,
211 scale: &raw[0][0].1,
212 },
213 ProjData {
214 packed: &raw[0][1].0,
215 scale: &raw[0][1].1,
216 },
217 ProjData {
218 packed: &raw[0][2].0,
219 scale: &raw[0][2].1,
220 },
221 ];
222 let header = ExpertRecordHeader {
223 layer,
224 expert,
225 inter: index.inter as u32,
226 hidden: index.hidden as u32,
227 group_size: index.group_size as u32,
228 scale2: [seed as f32, seed as f32 + 0.5, seed as f32 + 1.0],
229 input_scale: [Some(1.0), None, Some(2.0)],
230 };
231 w.write_record(ExpertKey::new(layer, expert), &header, &projs)
232 .unwrap();
233 expected.insert((layer, expert), (raw, header));
234 }
235 }
236 w.finish().unwrap();
237 }
238
239 let r = ExpertFileReader::open(&dir).unwrap();
241 assert_eq!(r.index(), &index);
242 for layer in 0..index.num_moe_layers {
243 for expert in 0..index.num_experts {
244 let key = ExpertKey::new(layer, expert);
245 let buf = r.read_record_raw(key).unwrap();
246 let (hdr, views) = unpack_record(r.spec(), &buf).unwrap();
247 let (raw, exp_hdr) = &expected[&(layer, expert)];
248 assert_eq!(&hdr, exp_hdr, "header {:?}", key);
249 for p in Proj::ALL {
250 assert_eq!(
251 views[p as usize].packed,
252 &raw[0][p as usize].0[..],
253 "packed {:?} {:?}",
254 key,
255 p
256 );
257 assert_eq!(
258 views[p as usize].scale,
259 &raw[0][p as usize].1[..],
260 "scale {:?} {:?}",
261 key,
262 p
263 );
264 }
265 }
266 }
267 std::fs::remove_dir_all(&dir).ok();
268 }
269
270 #[test]
271 fn wrong_projection_length_errors() {
272 let index = synth_index();
273 let spec = index.spec();
274 let header = ExpertRecordHeader {
275 layer: 0,
276 expert: 0,
277 inter: index.inter as u32,
278 hidden: index.hidden as u32,
279 group_size: index.group_size as u32,
280 scale2: [1.0; 3],
281 input_scale: [Some(1.0); 3],
282 };
283 let bad = vec![0u8; 8]; let ok_scale = vec![0u8; spec.proj_bytes(Proj::Gate).scale_bytes as usize];
285 let projs = [
286 ProjData {
287 packed: &bad,
288 scale: &ok_scale,
289 },
290 ProjData {
291 packed: &bad,
292 scale: &ok_scale,
293 },
294 ProjData {
295 packed: &bad,
296 scale: &ok_scale,
297 },
298 ];
299 let err = pack_record(&spec, index.record_stride, &header, &projs);
300 assert!(err.is_err(), "short packed buffer must error");
301 }
302
303 #[test]
304 fn read_record_raw_rejects_out_of_range_layer() {
305 let dir = tmpdir("oob");
306 let index = synth_index(); ExpertFileWriter::create(&dir, index)
308 .unwrap()
309 .finish()
310 .unwrap();
311 let r = ExpertFileReader::open(&dir).unwrap();
312 assert!(r.read_record_raw(ExpertKey::new(0, 0)).is_ok());
314 assert!(r.read_record_raw(ExpertKey::new(99, 0)).is_err());
315 std::fs::remove_dir_all(&dir).ok();
316 }
317
318 #[test]
319 fn manifest_geometry_round_trips_through_json() {
320 let index = synth_index();
321 let json = serde_json::to_string(&index).unwrap();
322 let back: ExpertIndex = serde_json::from_str(&json).unwrap();
323 assert_eq!(index, back);
324 assert_eq!(index.layout().record_stride, back.layout().record_stride);
326 assert_eq!(index.spec().raw_bytes(), back.spec().raw_bytes());
327 }
328}