Save inference parity implementation and evaluation harness
This commit is contained in:
@@ -0,0 +1,424 @@
|
||||
//! A single atomic checkpoint for the existing DS4 KV index, not a second bank.
|
||||
//! Tensor payloads keep the reference codec's dtypes, shapes and raw bytes.
|
||||
use super::request::LiveSession;
|
||||
use super::session_codec::decode_payload;
|
||||
use super::stream::{Stream, Streams};
|
||||
use crate::engine::metal::checkpoint::{read_u64, write_u64};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::BTreeMap;
|
||||
use std::fs::{self, File};
|
||||
use std::io::{Read, Seek, SeekFrom, Write};
|
||||
use std::path::Path;
|
||||
|
||||
const MAGIC: &[u8; 8] = b"QWMTKV01";
|
||||
const HEADER_BYTES: u64 = 16;
|
||||
const MAX_MANIFEST_BYTES: u64 = 64 * 1024 * 1024;
|
||||
const IO_CHUNK: usize = 8 * 1024 * 1024;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub(super) struct Identity {
|
||||
pub(super) model: [u8; 32],
|
||||
pub(super) context: u32,
|
||||
pub(super) vocabulary: u32,
|
||||
pub(super) mtp: bool,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
struct Manifest {
|
||||
identity: Identity,
|
||||
tag: [u8; 32],
|
||||
tokens: Vec<i32>,
|
||||
lazy_kv: bool,
|
||||
payload: serde_json::Value,
|
||||
// Offsets are absolute; blob names are codec-local, never filesystem paths.
|
||||
blobs: BTreeMap<String, (u64, u64)>,
|
||||
}
|
||||
|
||||
fn check_abort(abort: &impl Fn() -> bool) -> Result<(), String> {
|
||||
if abort() {
|
||||
Err("Qwen checkpoint interrupted".into())
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_tokens(tokens: &[i32], identity: Identity) -> Result<(), String> {
|
||||
if identity.context == 0
|
||||
|| identity.vocabulary == 0
|
||||
|| tokens.is_empty()
|
||||
|| tokens.len() > identity.context as usize
|
||||
|| tokens
|
||||
.iter()
|
||||
.any(|&token| token < 0 || token as u32 >= identity.vocabulary)
|
||||
{
|
||||
return Err("Qwen checkpoint token prefix is invalid".into());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
impl LiveSession {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(super) fn save_checkpoint(
|
||||
&mut self,
|
||||
path: &Path,
|
||||
identity: Identity,
|
||||
tag: [u8; 32],
|
||||
streams: &Streams,
|
||||
stream: Stream,
|
||||
mut progress: impl FnMut(u64),
|
||||
abort: impl Fn() -> bool,
|
||||
) -> Result<(), String> {
|
||||
check_abort(&abort)?;
|
||||
let saved = self
|
||||
.current
|
||||
.as_ref()
|
||||
.ok_or("Qwen has no committed checkpoint")?;
|
||||
validate_tokens(&saved.token_ids, identity)?;
|
||||
if saved.mtp.is_some() != identity.mtp {
|
||||
return Err("Qwen checkpoint MTP configuration differs".into());
|
||||
}
|
||||
if let Some(parent) = path.parent() {
|
||||
fs::create_dir_all(parent).map_err(|e| e.to_string())?;
|
||||
}
|
||||
// Same serialized-writer, staging-file and sync/rename contract as DS4.
|
||||
let temporary = path.with_extension("tmp");
|
||||
let result = (|| -> Result<(), String> {
|
||||
let mut file = File::create(&temporary).map_err(|e| e.to_string())?;
|
||||
file.write_all(MAGIC).map_err(|e| e.to_string())?;
|
||||
write_u64(&mut file, 0)?;
|
||||
progress(HEADER_BYTES);
|
||||
let mut blobs = BTreeMap::new();
|
||||
let payload = saved.encode_payload(
|
||||
256,
|
||||
streams,
|
||||
stream,
|
||||
|name, bytes| {
|
||||
let offset = file.stream_position().map_err(|e| e.to_string())?;
|
||||
for chunk in bytes.chunks(IO_CHUNK) {
|
||||
check_abort(&abort)?;
|
||||
file.write_all(chunk).map_err(|e| e.to_string())?;
|
||||
progress(chunk.len() as u64);
|
||||
}
|
||||
if blobs.insert(name, (offset, bytes.len() as u64)).is_some() {
|
||||
return Err("duplicate Qwen checkpoint blob".into());
|
||||
}
|
||||
Ok(())
|
||||
},
|
||||
&abort,
|
||||
)?;
|
||||
check_abort(&abort)?;
|
||||
let offset = file.stream_position().map_err(|e| e.to_string())?;
|
||||
let manifest = serde_json::to_vec(&Manifest {
|
||||
identity,
|
||||
tag,
|
||||
tokens: saved.token_ids.clone(),
|
||||
lazy_kv: saved.lazy_kv,
|
||||
payload,
|
||||
blobs,
|
||||
})
|
||||
.map_err(|e| e.to_string())?;
|
||||
if manifest.len() as u64 > MAX_MANIFEST_BYTES {
|
||||
return Err("Qwen checkpoint metadata exceeds limit".into());
|
||||
}
|
||||
file.write_all(&manifest).map_err(|e| e.to_string())?;
|
||||
progress(manifest.len() as u64);
|
||||
file.seek(SeekFrom::Start(8)).map_err(|e| e.to_string())?;
|
||||
write_u64(&mut file, offset)?;
|
||||
file.sync_all().map_err(|e| e.to_string())?;
|
||||
check_abort(&abort)?;
|
||||
fs::rename(&temporary, path).map_err(|e| e.to_string())
|
||||
})();
|
||||
if result.is_err() {
|
||||
let _ = fs::remove_file(&temporary);
|
||||
} else {
|
||||
self.tag = tag;
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(super) fn load_checkpoint(
|
||||
&mut self,
|
||||
path: &Path,
|
||||
identity: Identity,
|
||||
streams: &Streams,
|
||||
stream: Stream,
|
||||
mut progress: impl FnMut(u64),
|
||||
abort: impl Fn() -> bool,
|
||||
) -> Result<bool, String> {
|
||||
check_abort(&abort)?;
|
||||
let mut file = match File::open(path) {
|
||||
Ok(file) => file,
|
||||
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(false),
|
||||
Err(e) => return Err(e.to_string()),
|
||||
};
|
||||
let mut magic = [0; 8];
|
||||
file.read_exact(&mut magic).map_err(|e| e.to_string())?;
|
||||
if &magic != MAGIC {
|
||||
return Err("Qwen checkpoint format differs; rebuild the cache".into());
|
||||
}
|
||||
let offset = read_u64(&mut file)?;
|
||||
let length = file.metadata().map_err(|e| e.to_string())?.len();
|
||||
if offset < HEADER_BYTES || offset >= length || length - offset > MAX_MANIFEST_BYTES {
|
||||
return Err("Qwen checkpoint metadata range is invalid".into());
|
||||
}
|
||||
file.seek(SeekFrom::Start(offset))
|
||||
.map_err(|e| e.to_string())?;
|
||||
let mut raw = vec![0; (length - offset) as usize];
|
||||
file.read_exact(&mut raw).map_err(|e| e.to_string())?;
|
||||
let mut manifest: Manifest = serde_json::from_slice(&raw).map_err(|e| e.to_string())?;
|
||||
progress(HEADER_BYTES + raw.len() as u64);
|
||||
if manifest.identity != identity {
|
||||
return Err(
|
||||
"Qwen checkpoint does not match the current model/context/MTP policy".into(),
|
||||
);
|
||||
}
|
||||
validate_tokens(&manifest.tokens, identity)?;
|
||||
if !manifest.payload["mtp_history_snapshot"].is_null() != identity.mtp {
|
||||
return Err("Qwen checkpoint MTP payload differs".into());
|
||||
}
|
||||
// Validate all extents before allocating any GPU state. Require exact
|
||||
// coverage, so overlapping blobs and unindexed/trailing bytes fail closed.
|
||||
let mut ranges = manifest.blobs.values().copied().collect::<Vec<_>>();
|
||||
ranges.sort_unstable();
|
||||
let mut end = HEADER_BYTES;
|
||||
for (start, size) in ranges {
|
||||
if start != end {
|
||||
return Err("Qwen checkpoint blob range is invalid".into());
|
||||
}
|
||||
end = start
|
||||
.checked_add(size)
|
||||
.filter(|&n| n <= offset)
|
||||
.ok_or("Qwen checkpoint blob exceeds payload")?;
|
||||
}
|
||||
if end != offset {
|
||||
return Err("Qwen checkpoint payload coverage differs".into());
|
||||
}
|
||||
let saved = decode_payload(
|
||||
&manifest.payload,
|
||||
manifest.tokens,
|
||||
manifest.lazy_kv,
|
||||
true,
|
||||
streams,
|
||||
stream,
|
||||
|name| {
|
||||
check_abort(&abort)?;
|
||||
let (start, size) = manifest
|
||||
.blobs
|
||||
.remove(name)
|
||||
.ok_or("missing or repeated Qwen checkpoint blob")?;
|
||||
file.seek(SeekFrom::Start(start))
|
||||
.map_err(|e| e.to_string())?;
|
||||
let mut bytes = Vec::new();
|
||||
let size = usize::try_from(size).map_err(|_| "Qwen checkpoint blob too large")?;
|
||||
bytes.try_reserve_exact(size).map_err(|e| e.to_string())?;
|
||||
bytes.resize(size, 0);
|
||||
for chunk in bytes.chunks_mut(IO_CHUNK) {
|
||||
check_abort(&abort)?;
|
||||
file.read_exact(chunk).map_err(|e| e.to_string())?;
|
||||
progress(chunk.len() as u64);
|
||||
}
|
||||
Ok(bytes)
|
||||
},
|
||||
)?;
|
||||
if !manifest.blobs.is_empty() {
|
||||
return Err("unreferenced Qwen checkpoint blobs".into());
|
||||
}
|
||||
check_abort(&abort)?;
|
||||
// Publish only after complete validation and materialization. Failure
|
||||
// leaves the live conversation and its prefix untouched.
|
||||
self.current = Some(saved);
|
||||
self.prefix = None;
|
||||
self.tag = manifest.tag;
|
||||
Ok(true)
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "model-free Metal checkpoint; external supervisor and DS4_QWEN38_CHECKPOINT_DIR required"]
|
||||
fn qwen_checkpoint_roundtrip_and_failed_restore_preserve_live_state() {
|
||||
use super::super::gpu::Context;
|
||||
use super::session_cache::{Boundary, SessionSnapshot, State};
|
||||
use super::stream::Device;
|
||||
use super::{
|
||||
allocator::Allocator,
|
||||
array::{Array, Dtype},
|
||||
configure_sources, scalar_buffer,
|
||||
};
|
||||
use std::cell::Cell;
|
||||
configure_sources().unwrap();
|
||||
let _context = Context::open_qwen(0).unwrap();
|
||||
let streams = Streams::new(Some(Allocator::new().unwrap()));
|
||||
let stream = streams.default_stream(Device::Gpu).unwrap();
|
||||
let leaf = Array::new(
|
||||
&[1, 1, 4],
|
||||
Dtype::BF16,
|
||||
scalar_buffer(&[0, 0, 128, 63, 0, 64, 64, 64]).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let state = Some(State::Recurrent(vec![Some(leaf.clone()), None]));
|
||||
let mut live = LiveSession {
|
||||
current: Some(SessionSnapshot {
|
||||
token_ids: vec![1, 2, 3, 4],
|
||||
logits: leaf.clone(),
|
||||
hidden: Some(leaf.clone()),
|
||||
trunk: vec![state.clone()],
|
||||
mtp: Some(vec![None]),
|
||||
lazy_kv: true,
|
||||
boundaries: vec![Boundary {
|
||||
tokens: 2,
|
||||
state: vec![state],
|
||||
hidden: Some(leaf),
|
||||
}],
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
let identity = Identity {
|
||||
model: [7; 32],
|
||||
context: 64,
|
||||
vocabulary: 256,
|
||||
mtp: true,
|
||||
};
|
||||
let directory =
|
||||
std::path::PathBuf::from(std::env::var_os("DS4_QWEN38_CHECKPOINT_DIR").unwrap());
|
||||
fs::create_dir_all(&directory).unwrap();
|
||||
let path = directory.join("roundtrip.bin");
|
||||
let fingerprint = |live: &LiveSession| {
|
||||
let mut blobs = BTreeMap::new();
|
||||
let payload = live
|
||||
.current
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.encode_payload(
|
||||
256,
|
||||
&streams,
|
||||
stream,
|
||||
|name, bytes| {
|
||||
blobs.insert(name, bytes);
|
||||
Ok(())
|
||||
},
|
||||
|| false,
|
||||
)
|
||||
.unwrap();
|
||||
(payload, blobs)
|
||||
};
|
||||
let original = fingerprint(&live);
|
||||
let mut written = 0;
|
||||
live.save_checkpoint(
|
||||
&path,
|
||||
identity,
|
||||
[9; 32],
|
||||
&streams,
|
||||
stream,
|
||||
|n| written += n,
|
||||
|| false,
|
||||
)
|
||||
.unwrap();
|
||||
let valid = fs::read(&path).unwrap();
|
||||
assert_eq!(written, valid.len() as u64);
|
||||
assert_eq!(live.tag, [9; 32]);
|
||||
live = LiveSession::default();
|
||||
let mut read = 0;
|
||||
assert!(
|
||||
live.load_checkpoint(&path, identity, &streams, stream, |n| read += n, || false)
|
||||
.unwrap()
|
||||
);
|
||||
assert_eq!(read, written);
|
||||
assert_eq!(fingerprint(&live), original);
|
||||
assert_eq!(live.tag, [9; 32]);
|
||||
assert_eq!(live.current.as_ref().unwrap().token_ids, [1, 2, 3, 4]);
|
||||
assert!(live.current.as_ref().unwrap().lazy_kv);
|
||||
assert!(
|
||||
!live
|
||||
.load_checkpoint(
|
||||
&directory.join("missing.bin"),
|
||||
identity,
|
||||
&streams,
|
||||
stream,
|
||||
|_| {},
|
||||
|| false
|
||||
)
|
||||
.unwrap()
|
||||
);
|
||||
|
||||
// Cancellation after IO has started must not replace a valid file/tag.
|
||||
let interrupted = Cell::new(false);
|
||||
assert!(
|
||||
live.save_checkpoint(
|
||||
&path,
|
||||
identity,
|
||||
[8; 32],
|
||||
&streams,
|
||||
stream,
|
||||
|_| interrupted.set(true),
|
||||
|| interrupted.get()
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
assert_eq!(fs::read(&path).unwrap(), valid);
|
||||
assert!(!path.with_extension("tmp").exists());
|
||||
interrupted.set(false);
|
||||
assert!(
|
||||
live.load_checkpoint(
|
||||
&path,
|
||||
identity,
|
||||
&streams,
|
||||
stream,
|
||||
|_| interrupted.set(true),
|
||||
|| interrupted.get()
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
for changed in [
|
||||
Identity {
|
||||
context: 65,
|
||||
..identity
|
||||
},
|
||||
Identity {
|
||||
mtp: false,
|
||||
..identity
|
||||
},
|
||||
Identity {
|
||||
model: [8; 32],
|
||||
..identity
|
||||
},
|
||||
] {
|
||||
assert!(
|
||||
live.load_checkpoint(&path, changed, &streams, stream, |_| {}, || false)
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
let offset = u64::from_le_bytes(valid[8..16].try_into().unwrap()) as usize;
|
||||
let broken = directory.join("broken.bin");
|
||||
let mutations: [fn(&mut Manifest); 5] = [
|
||||
|m| m.tokens[0] = -1,
|
||||
|m| m.tokens.resize(65, 1),
|
||||
|m| m.blobs.values_mut().next().unwrap().1 = u64::MAX,
|
||||
|m| m.blobs.values_mut().next().unwrap().0 = 0,
|
||||
|m| m.payload["logits"]["nbytes"] = serde_json::json!(99),
|
||||
];
|
||||
for mutate in mutations {
|
||||
let mut manifest: Manifest = serde_json::from_slice(&valid[offset..]).unwrap();
|
||||
mutate(&mut manifest);
|
||||
let mut bytes = valid[..offset].to_vec();
|
||||
bytes.extend(serde_json::to_vec(&manifest).unwrap());
|
||||
fs::write(&broken, bytes).unwrap();
|
||||
assert!(
|
||||
live.load_checkpoint(&broken, identity, &streams, stream, |_| {}, || false)
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
for bytes in [&b"DS4QWN01"[..], &valid[..offset], &valid[..7]] {
|
||||
fs::write(&broken, bytes).unwrap();
|
||||
assert!(
|
||||
live.load_checkpoint(&broken, identity, &streams, stream, |_| {}, || false)
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
assert_eq!(live.tag, [9; 32]);
|
||||
assert_eq!(fingerprint(&live), original);
|
||||
println!(
|
||||
"Qwen checkpoint: exact payload roundtrip, interrupted save, incompatible/corrupt restore passed"
|
||||
);
|
||||
}
|
||||
Reference in New Issue
Block a user