425 lines
14 KiB
Rust
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"
|
|
);
|
|
}
|