use super::*; impl Executor { pub(in crate::engine) fn save_checkpoint( &mut self, path: &Path, tag: [u8; 32], progress: &mut impl FnMut(u64), ) -> Result<(), String> { if let Some(parent) = path.parent() { fs::create_dir_all(parent).map_err(|error| error.to_string())?; } let temporary = path.with_extension("tmp"); let mut file = File::create(&temporary).map_err(|error| error.to_string())?; let mut reported = 0_u64; let mut pending = 0_u64; self.write_checkpoint(&mut file, tag, &mut |bytes| { reported += bytes; pending += bytes; if pending >= CHECKPOINT_IO_CHUNK as u64 { progress(pending); pending = 0; } })?; progress(pending); let total = file.metadata().map_err(|error| error.to_string())?.len(); progress(total.saturating_sub(reported)); file.sync_all().map_err(|error| error.to_string())?; fs::rename(&temporary, path).map_err(|error| error.to_string())?; self.checkpoint_tag = tag; Ok(()) } pub(in crate::engine) fn load_checkpoint( &mut self, path: &Path, progress: &mut impl FnMut(u64), ) -> Result { let mut file = match File::open(path) { Ok(file) => file, Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false), Err(error) => return Err(error.to_string()), }; let mut reported = 0_u64; let mut pending = 0_u64; self.read_checkpoint(&mut file, &mut |bytes| { reported += bytes; pending += bytes; if pending >= CHECKPOINT_IO_CHUNK as u64 { progress(pending); pending = 0; } })?; progress(pending); let total = file.metadata().map_err(|error| error.to_string())?.len(); progress(total.saturating_sub(reported)); Ok(true) } fn write_checkpoint( &self, file: &mut File, tag: [u8; 32], progress: &mut impl FnMut(u64), ) -> Result<(), String> { let shape = self.model.shape; file.write_all(CHECKPOINT_MAGIC) .map_err(|error| error.to_string())?; for value in [ CHECKPOINT_VERSION, self.session.context, self.session.raw_cap, shape.layers, shape.head_dim as u32, shape.indexer_head_dim as u32, shape.vocab as u32, u32::from(self.quality), ] { write_u32(file, value)?; } write_u64(file, self.model.main.len())?; write_u64(file, self.model_modified.0)?; write_u32(file, self.model_modified.1)?; for weight in [self.weights.token_embedding, self.weights.output] { write_u64(file, weight.offset)?; write_u64(file, weight.bytes)?; write_u32(file, weight.kind)?; } file.write_all(&tag).map_err(|error| error.to_string())?; let token_count = u32::try_from(self.tokens.len()) .map_err(|_| "KV checkpoint has too many tokens".to_owned())?; let raw_live = token_count.min(self.session.raw_cap); write_u32(file, token_count)?; write_u32(file, raw_live)?; for &token in &self.tokens { write_u32(file, token as u32)?; } for &logit in &self.logits { write_u32(file, logit.to_bits())?; } for layer in &self.session.layers { write_u32( file, layer.compression.as_ref().map_or(0, |state| state.rows), )?; } for layer in &self.session.layers { write_u32(file, layer.indexer.as_ref().map_or(0, |state| state.rows))?; } let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; let raw_first = token_count - raw_live; for layer in &self.session.layers { for position in raw_first..token_count { let physical = position % self.session.raw_cap; write_buffer( file, &layer.raw_cache, u64::from(physical) * shape.head_dim * 4, shape.head_dim * 4, &mut chunk, progress, )?; } if let Some(state) = &layer.compression { write_buffer( file, &state.cache, 0, u64::from(state.rows) * shape.head_dim * 2, &mut chunk, progress, )?; let bytes = compressor_state_bytes(state.ratio, shape.head_dim); write_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; write_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; } if let Some(state) = &layer.indexer { write_buffer( file, &state.cache, 0, u64::from(state.rows) * shape.indexer_head_dim * 4, &mut chunk, progress, )?; let bytes = compressor_state_bytes(state.ratio, shape.indexer_head_dim); write_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; write_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; } } Ok(()) } fn read_checkpoint( &mut self, file: &mut File, progress: &mut impl FnMut(u64), ) -> Result<(), String> { let mut magic = [0; 8]; file.read_exact(&mut magic) .map_err(|error| error.to_string())?; if &magic != CHECKPOINT_MAGIC { return Err("KV checkpoint has an invalid signature".into()); } let shape = self.model.shape; let header = [ CHECKPOINT_VERSION, self.session.context, self.session.raw_cap, shape.layers, shape.head_dim as u32, shape.indexer_head_dim as u32, shape.vocab as u32, u32::from(self.quality), ]; for expected in header { if read_u32(file)? != expected { return Err("KV checkpoint does not match the current executor".into()); } } if read_u64(file)? != self.model.main.len() { return Err("KV checkpoint was written for a different model".into()); } if read_u64(file)? != self.model_modified.0 || read_u32(file)? != self.model_modified.1 { return Err("KV checkpoint model file has changed".into()); } for weight in [self.weights.token_embedding, self.weights.output] { if read_u64(file)? != weight.offset || read_u64(file)? != weight.bytes || read_u32(file)? != weight.kind { return Err("KV checkpoint was written for a different model layout".into()); } } let mut checkpoint_tag = [0; 32]; file.read_exact(&mut checkpoint_tag) .map_err(|error| error.to_string())?; let token_count = read_u32(file)?; let raw_live = read_u32(file)?; if token_count > self.session.context || raw_live != token_count.min(self.session.raw_cap) { return Err("KV checkpoint token count is invalid".into()); } let mut tokens = Vec::with_capacity(token_count as usize); for _ in 0..token_count { let token = read_u32(file)?; if u64::from(token) >= shape.vocab { return Err("KV checkpoint contains an invalid token".into()); } tokens.push(token as i32); } let mut logits = Vec::with_capacity(shape.vocab as usize); for _ in 0..shape.vocab { logits.push(f32::from_bits(read_u32(file)?)); } let mut compressed_rows = Vec::with_capacity(shape.layers as usize); let mut indexer_rows = Vec::with_capacity(shape.layers as usize); for layer in 0..shape.layers { let rows = read_u32(file)?; let ratio = compression_ratio(layer); if rows != token_count.checked_div(ratio).unwrap_or(0) { return Err("KV checkpoint compressed row count is invalid".into()); } compressed_rows.push(rows); } for layer in 0..shape.layers { let rows = read_u32(file)?; let expected = if compression_ratio(layer) == 4 { token_count / 4 } else { 0 }; if rows != expected { return Err("KV checkpoint indexer row count is invalid".into()); } indexer_rows.push(rows); } self.reset()?; let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; let raw_first = token_count - raw_live; for (index, layer) in self.session.layers.iter_mut().enumerate() { for position in raw_first..token_count { let physical = position % self.session.raw_cap; read_buffer( file, &layer.raw_cache, u64::from(physical) * shape.head_dim * 4, shape.head_dim * 4, &mut chunk, progress, )?; } if let Some(state) = &mut layer.compression { state.rows = compressed_rows[index]; read_buffer( file, &state.cache, 0, u64::from(state.rows) * shape.head_dim * 2, &mut chunk, progress, )?; let bytes = compressor_state_bytes(state.ratio, shape.head_dim); read_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; read_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; } if let Some(state) = &mut layer.indexer { state.rows = indexer_rows[index]; read_buffer( file, &state.cache, 0, u64::from(state.rows) * shape.indexer_head_dim * 4, &mut chunk, progress, )?; let bytes = compressor_state_bytes(state.ratio, shape.indexer_head_dim); read_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; read_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; } } let mut trailing = [0]; if file .read(&mut trailing) .map_err(|error| error.to_string())? != 0 { self.reset()?; return Err("KV checkpoint has trailing data".into()); } self.session.position = token_count; self.tokens = tokens; self.logits = logits; self.checkpoint_tag = checkpoint_tag; Ok(()) } } fn compressor_state_bytes(ratio: u32, head_dim: u64) -> u64 { let coefficient = if ratio == 4 { 2 } else { 1 }; coefficient * head_dim * coefficient * u64::from(ratio) * 4 } fn write_buffer( file: &mut File, buffer: &Buffer, mut offset: u64, mut bytes: u64, chunk: &mut [u8], progress: &mut impl FnMut(u64), ) -> Result<(), String> { while bytes != 0 { let length = bytes.min(chunk.len() as u64) as usize; buffer.read(offset, &mut chunk[..length])?; file.write_all(&chunk[..length]) .map_err(|error| error.to_string())?; progress(length as u64); offset += length as u64; bytes -= length as u64; } Ok(()) } fn read_buffer( file: &mut File, buffer: &Buffer, mut offset: u64, mut bytes: u64, chunk: &mut [u8], progress: &mut impl FnMut(u64), ) -> Result<(), String> { while bytes != 0 { let length = bytes.min(chunk.len() as u64) as usize; file.read_exact(&mut chunk[..length]) .map_err(|error| error.to_string())?; buffer.write(offset, &chunk[..length])?; progress(length as u64); offset += length as u64; bytes -= length as u64; } Ok(()) } fn write_u32(file: &mut File, value: u32) -> Result<(), String> { file.write_all(&value.to_le_bytes()) .map_err(|error| error.to_string()) } fn write_u64(file: &mut File, value: u64) -> Result<(), String> { file.write_all(&value.to_le_bytes()) .map_err(|error| error.to_string()) } fn read_u32(file: &mut File) -> Result { let mut bytes = [0; 4]; file.read_exact(&mut bytes) .map_err(|error| error.to_string())?; Ok(u32::from_le_bytes(bytes)) } fn read_u64(file: &mut File) -> Result { let mut bytes = [0; 8]; file.read_exact(&mut bytes) .map_err(|error| error.to_string())?; Ok(u64::from_le_bytes(bytes)) }