Files
DS4Server/src/engine/metal/qwen_mtplx/session_checkpoint.rs
T

425 lines
14 KiB
Rust

//! 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"
);
}