use std::io::{Read, Write};
use crate::tensor::{Device, DType, Result, Tensor, TensorError};
use super::buffer::Buffer;
use super::parameter::Parameter;
pub(crate) const MAGIC: [u8; 4] = *b"FDLC";
pub(crate) const VERSION: u32 = 2;
const MAX_VERSION: u32 = 2;
pub(crate) const HASH_LEN: usize = 32;
#[derive(Debug, Clone)]
pub struct LoadReport {
pub loaded: Vec<String>,
pub skipped: Vec<String>,
pub missing: Vec<String>,
}
pub fn save_checkpoint<W: Write>(
w: &mut W,
params: &[(String, Parameter)],
buffers: &[(String, Buffer)],
structural_hash: Option<&str>,
) -> Result<()> {
let total = (params.len() + buffers.len()) as u32;
write_checkpoint_header(w, total, structural_hash)?;
for (name, p) in params {
write_entry_name(w, name)?;
write_tensor_data(w, &p.variable.data())?;
}
for (name, b) in buffers {
write_entry_name(w, name)?;
write_tensor_data(w, &b.get())?;
}
Ok(())
}
pub(crate) fn write_checkpoint_header<W: Write>(
w: &mut W,
total: u32,
structural_hash: Option<&str>,
) -> Result<()> {
w.write_all(&MAGIC).map_err(io_err)?;
w.write_all(&VERSION.to_le_bytes()).map_err(io_err)?;
let hash_bytes = match structural_hash {
Some(hex) => hex_to_bytes(hex)?,
None => [0u8; HASH_LEN],
};
w.write_all(&hash_bytes).map_err(io_err)?;
w.write_all(&total.to_le_bytes()).map_err(io_err)?;
Ok(())
}
fn write_entry_name<W: Write>(w: &mut W, name: &str) -> Result<()> {
let name_bytes = name.as_bytes();
w.write_all(&(name_bytes.len() as u32).to_le_bytes()).map_err(io_err)?;
w.write_all(name_bytes).map_err(io_err)?;
Ok(())
}
pub(crate) struct RawCheckpointEntry<'a> {
pub name: &'a str,
pub shape: &'a [i64],
pub dtype_tag: u8,
pub raw: &'a [u8],
}
pub(crate) fn save_checkpoint_from_raw<W: Write>(
w: &mut W,
entries: &[RawCheckpointEntry<'_>],
structural_hash: Option<&str>,
) -> Result<()> {
write_checkpoint_header(w, entries.len() as u32, structural_hash)?;
for e in entries {
write_entry_name(w, e.name)?;
w.write_all(&(e.shape.len() as u32).to_le_bytes()).map_err(io_err)?;
for &s in e.shape {
w.write_all(&s.to_le_bytes()).map_err(io_err)?;
}
w.write_all(&[e.dtype_tag]).map_err(io_err)?;
w.write_all(&(e.raw.len() as u64).to_le_bytes()).map_err(io_err)?;
w.write_all(e.raw).map_err(io_err)?;
}
Ok(())
}
pub(crate) fn save_checkpoint_from_raw_file(
path: &str,
entries: &[RawCheckpointEntry<'_>],
structural_hash: Option<&str>,
) -> Result<()> {
write_file_atomic(path, |mut w| save_checkpoint_from_raw(&mut w, entries, structural_hash))
}
pub fn load_checkpoint<R: Read>(
r: &mut R,
params: &[(String, Parameter)],
buffers: &[(String, Buffer)],
structural_hash: Option<&str>,
) -> Result<LoadReport> {
let mut magic = [0u8; 4];
r.read_exact(&mut magic).map_err(io_err)?;
if magic != MAGIC {
return Err(TensorError::new(
"invalid checkpoint: bad magic (expected .fdl checkpoint)"
));
}
let version = read_u32(r)?;
if version == 0 || version > MAX_VERSION {
return Err(TensorError::new(&format!(
"unsupported checkpoint version {} (this build supports 1..={})",
version, MAX_VERSION,
)));
}
let mut file_hash = [0u8; HASH_LEN];
r.read_exact(&mut file_hash).map_err(io_err)?;
let file_nonzero = file_hash.iter().any(|&b| b != 0);
if let Some(expected_hex) = structural_hash {
let expected = hex_to_bytes(expected_hex)?;
let expected_nonzero = expected.iter().any(|&b| b != 0);
if file_nonzero && expected_nonzero && file_hash != expected {
return Err(TensorError::new(&format!(
"checkpoint architecture mismatch: file={} model={}",
bytes_to_hex(&file_hash),
expected_hex,
)));
}
}
let count = read_u32(r)? as usize;
let mut ckpt: std::collections::HashMap<String, (Vec<i64>, DType, Vec<u8>)> =
std::collections::HashMap::with_capacity(count);
for _ in 0..count {
let name = read_name(r)?;
let shape = read_shape(r)?;
let mut tag = [0u8; 1];
r.read_exact(&mut tag).map_err(io_err)?;
let dtype = dtype_from_tag(tag[0])?;
let byte_count = read_u64(r)? as usize;
let raw = read_payload(r, byte_count)?;
ckpt.insert(name, (shape, dtype, raw));
}
let mut loaded = Vec::new();
let mut missing = Vec::new();
for (name, p) in params {
if let Some((shape, dtype, raw)) = ckpt.remove(name) {
let model_shape = p.variable.shape();
if shape != model_shape {
return Err(TensorError::new(&format!(
"parameter {:?}: shape mismatch: checkpoint={:?} model={:?}",
name, shape, model_shape
)));
}
let t = tensor_from_raw_bytes(&raw, &shape, dtype)?;
let model_dtype = p.variable.data().dtype();
let t = if t.dtype() != model_dtype { t.to_dtype(model_dtype)? } else { t };
let dev = p.variable.data().device();
if dev != Device::CPU {
p.variable.set_data(t.to_device(dev)?);
} else {
p.variable.set_data(t);
}
loaded.push(name.clone());
} else {
missing.push(name.clone());
}
}
for (name, b) in buffers {
if let Some((shape, dtype, raw)) = ckpt.remove(name) {
let model_shape = b.shape();
if shape != model_shape {
return Err(TensorError::new(&format!(
"buffer {:?}: shape mismatch: checkpoint={:?} model={:?}",
name, shape, model_shape
)));
}
let t = tensor_from_raw_bytes(&raw, &shape, dtype)?;
let model_dtype = b.get().dtype();
let t = if t.dtype() != model_dtype { t.to_dtype(model_dtype)? } else { t };
let dev = b.device();
if dev != Device::CPU {
b.set(t.to_device(dev)?);
} else {
b.set(t);
}
loaded.push(name.clone());
} else {
missing.push(name.clone());
}
}
let skipped: Vec<String> = ckpt.into_keys().collect();
Ok(LoadReport { loaded, skipped, missing })
}
pub(crate) fn write_file_atomic<T>(
path: &str,
write: impl FnOnce(&mut dyn Write) -> Result<T>,
) -> Result<T> {
let is_gz = path.ends_with(".gz");
let tmp = format!("{path}.tmp");
let write_result = (|| -> Result<T> {
let f = std::fs::File::create(&tmp).map_err(io_err)?;
if is_gz {
let mut w = flate2::write::GzEncoder::new(f, flate2::Compression::default());
let v = write(&mut w)?;
w.finish().map_err(io_err)?;
Ok(v)
} else {
let mut w = std::io::BufWriter::new(f);
let v = write(&mut w)?;
w.flush().map_err(io_err)?;
Ok(v)
}
})();
match write_result {
Ok(v) => {
std::fs::rename(&tmp, path).map_err(io_err)?;
Ok(v)
}
Err(e) => {
let _ = std::fs::remove_file(&tmp);
Err(e)
}
}
}
pub fn save_checkpoint_file(
path: &str,
params: &[(String, Parameter)],
buffers: &[(String, Buffer)],
structural_hash: Option<&str>,
) -> Result<()> {
write_file_atomic(path, |mut w| save_checkpoint(&mut w, params, buffers, structural_hash))
}
pub fn load_checkpoint_file(
path: &str,
params: &[(String, Parameter)],
buffers: &[(String, Buffer)],
structural_hash: Option<&str>,
) -> Result<LoadReport> {
let f = std::fs::File::open(path).map_err(io_err)?;
if path.ends_with(".gz") {
let mut r = flate2::read::GzDecoder::new(f);
load_checkpoint(&mut r, params, buffers, structural_hash)
} else {
let mut r = std::io::BufReader::new(f);
load_checkpoint(&mut r, params, buffers, structural_hash)
}
}
pub fn checkpoint_keys(path: &str) -> Result<Vec<String>> {
let f = std::fs::File::open(path).map_err(io_err)?;
let mut r: Box<dyn Read> = if path.ends_with(".gz") {
Box::new(flate2::read::GzDecoder::new(f))
} else {
Box::new(std::io::BufReader::new(f))
};
let mut magic = [0u8; 4];
r.read_exact(&mut magic).map_err(io_err)?;
if magic != MAGIC {
return Err(TensorError::new(
"invalid checkpoint: bad magic (expected .fdl checkpoint)",
));
}
let version = read_u32(&mut r)?;
if version == 0 || version > MAX_VERSION {
return Err(TensorError::new(&format!(
"unsupported checkpoint version {} (this build supports 1..={})",
version, MAX_VERSION,
)));
}
let mut _hash = [0u8; HASH_LEN];
r.read_exact(&mut _hash).map_err(io_err)?;
let count = read_u32(&mut r)? as usize;
let mut keys = Vec::with_capacity(count);
for _ in 0..count {
keys.push(read_name(&mut r)?);
let _ = read_shape(&mut r)?;
let mut tag = [0u8; 1];
r.read_exact(&mut tag).map_err(io_err)?;
let byte_count = read_u64(&mut r)? as usize;
std::io::copy(&mut r.by_ref().take(byte_count as u64), &mut std::io::sink())
.map_err(io_err)?;
}
Ok(keys)
}
pub fn checkpoint_version(path: &str) -> Result<u32> {
let f = std::fs::File::open(path).map_err(io_err)?;
let mut r: Box<dyn Read> = if path.ends_with(".gz") {
Box::new(flate2::read::GzDecoder::new(f))
} else {
Box::new(std::io::BufReader::new(f))
};
let mut magic = [0u8; 4];
r.read_exact(&mut magic).map_err(io_err)?;
if magic != MAGIC {
return Err(TensorError::new(
"invalid checkpoint: bad magic (expected .fdl checkpoint)"
));
}
read_u32(&mut r)
}
pub(crate) fn write_tensor_state<W: Write>(w: &mut W, t: Option<&Tensor>) -> Result<()> {
match t {
None => {
w.write_all(&[0u8]).map_err(io_err)?;
}
Some(t) => {
w.write_all(&[1u8]).map_err(io_err)?;
write_tensor_data(w, t)?;
}
}
Ok(())
}
pub(crate) fn read_tensor_state<R: Read>(r: &mut R, device: Device) -> Result<Option<Tensor>> {
let mut present = [0u8; 1];
r.read_exact(&mut present).map_err(io_err)?;
if present[0] == 0 {
return Ok(None);
}
let t = read_tensor_data(r)?;
if device != Device::CPU {
Ok(Some(t.to_device(device)?))
} else {
Ok(Some(t))
}
}
pub(crate) fn dtype_tag(dtype: DType) -> u8 {
match dtype {
DType::Float16 => 1,
DType::BFloat16 => 2,
DType::Float32 => 3,
DType::Float64 => 4,
DType::Int32 => 5,
DType::Int64 => 6,
}
}
fn dtype_from_tag(tag: u8) -> Result<DType> {
match tag {
1 => Ok(DType::Float16),
2 => Ok(DType::BFloat16),
3 => Ok(DType::Float32),
4 => Ok(DType::Float64),
5 => Ok(DType::Int32),
6 => Ok(DType::Int64),
_ => Err(TensorError::new(&format!("unknown dtype tag: {}", tag))),
}
}
pub(crate) fn write_tensor_data<W: Write>(w: &mut W, t: &Tensor) -> Result<()> {
let shape = t.shape();
w.write_all(&(shape.len() as u32).to_le_bytes()).map_err(io_err)?;
for &s in &shape {
w.write_all(&s.to_le_bytes()).map_err(io_err)?;
}
let dtype = t.dtype();
w.write_all(&[dtype_tag(dtype)]).map_err(io_err)?;
let numel = t.numel() as usize;
let elem_size = dtype.element_size();
let byte_count = numel * elem_size;
let raw = copy_raw_bytes(t, byte_count)?;
w.write_all(&(byte_count as u64).to_le_bytes()).map_err(io_err)?;
w.write_all(&raw).map_err(io_err)?;
Ok(())
}
pub(crate) fn read_tensor_data<R: Read>(r: &mut R) -> Result<Tensor> {
let shape = read_shape(r)?;
let mut tag = [0u8; 1];
r.read_exact(&mut tag).map_err(io_err)?;
let dtype = dtype_from_tag(tag[0])?;
let byte_count = read_u64(r)? as usize;
let raw = read_payload(r, byte_count)?;
tensor_from_raw_bytes(&raw, &shape, dtype)
}
fn copy_raw_bytes(t: &Tensor, byte_count: usize) -> Result<Vec<u8>> {
let mut buf = vec![0u8; byte_count];
let err = unsafe {
flodl_sys::flodl_copy_data(
t.raw(),
buf.as_mut_ptr() as *mut std::ffi::c_void,
byte_count as i64,
)
};
check_err_raw(err)?;
Ok(buf)
}
fn tensor_from_raw_bytes(raw: &[u8], shape: &[i64], dtype: DType) -> Result<Tensor> {
match dtype {
DType::Float32 => {
let data: Vec<f32> = raw.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
Tensor::from_f32(&data, shape, Device::CPU)
}
DType::Float64 => {
let data: Vec<f64> = raw.chunks_exact(8)
.map(|c| f64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]))
.collect();
Tensor::from_f64(&data, shape, Device::CPU)
}
DType::Int64 => {
let data: Vec<i64> = raw.chunks_exact(8)
.map(|c| i64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]))
.collect();
Tensor::from_i64(&data, shape, Device::CPU)
}
DType::Float16 | DType::BFloat16 | DType::Int32 => {
Tensor::from_blob(raw, shape, dtype, Device::CPU)
}
}
}
#[derive(Debug, Clone)]
pub struct MigrateReport {
pub unchanged: Vec<String>,
pub remapped: Vec<(String, String)>,
pub dropped: Vec<String>,
pub missing: Vec<String>,
}
impl MigrateReport {
pub fn is_complete(&self) -> bool {
self.dropped.is_empty() && self.missing.is_empty()
}
}
impl std::fmt::Display for MigrateReport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
if !self.unchanged.is_empty() {
writeln!(f, "unchanged ({}):", self.unchanged.len())?;
for name in &self.unchanged { writeln!(f, " {}", name)?; }
}
if !self.remapped.is_empty() {
writeln!(f, "remapped ({}):", self.remapped.len())?;
for (old, new) in &self.remapped { writeln!(f, " {} -> {}", old, new)?; }
}
if !self.dropped.is_empty() {
writeln!(f, "dropped ({}):", self.dropped.len())?;
for name in &self.dropped { writeln!(f, " {}", name)?; }
}
if !self.missing.is_empty() {
writeln!(f, "missing ({}):", self.missing.len())?;
for name in &self.missing { writeln!(f, " {}", name)?; }
}
Ok(())
}
}
struct RawEntry {
name: String,
shape: Vec<i64>,
dtype: DType,
raw: Vec<u8>,
}
fn read_raw_checkpoint<R: Read>(r: &mut R) -> Result<Vec<RawEntry>> {
let mut magic = [0u8; 4];
r.read_exact(&mut magic).map_err(io_err)?;
if magic != MAGIC {
return Err(TensorError::new(
"invalid checkpoint: bad magic (expected .fdl checkpoint)"
));
}
let version = read_u32(r)?;
if version == 0 || version > MAX_VERSION {
return Err(TensorError::new(&format!(
"unsupported checkpoint version {} (this build supports 1..={})",
version, MAX_VERSION,
)));
}
let mut _hash = [0u8; HASH_LEN];
r.read_exact(&mut _hash).map_err(io_err)?;
let count = read_u32(r)? as usize;
let mut entries = Vec::with_capacity(count);
for _ in 0..count {
let name = read_name(r)?;
let shape = read_shape(r)?;
let mut tag = [0u8; 1];
r.read_exact(&mut tag).map_err(io_err)?;
let dtype = dtype_from_tag(tag[0])?;
let byte_count = read_u64(r)? as usize;
let raw = read_payload(r, byte_count)?;
entries.push(RawEntry { name, shape, dtype, raw });
}
Ok(entries)
}
fn write_raw_entry<W: Write>(w: &mut W, name: &str, e: &RawEntry) -> Result<()> {
let name_bytes = name.as_bytes();
w.write_all(&(name_bytes.len() as u32).to_le_bytes()).map_err(io_err)?;
w.write_all(name_bytes).map_err(io_err)?;
w.write_all(&(e.shape.len() as u32).to_le_bytes()).map_err(io_err)?;
for &s in &e.shape {
w.write_all(&s.to_le_bytes()).map_err(io_err)?;
}
w.write_all(&[dtype_tag(e.dtype)]).map_err(io_err)?;
w.write_all(&(e.raw.len() as u64).to_le_bytes()).map_err(io_err)?;
w.write_all(&e.raw).map_err(io_err)?;
Ok(())
}
pub fn migrate_checkpoint<R: Read, W: Write>(
r: &mut R,
w: &mut W,
params: &[(String, Parameter)],
buffers: &[(String, Buffer)],
) -> Result<MigrateReport> {
let entries = read_raw_checkpoint(r)?;
let mut targets: Vec<(String, Vec<i64>, DType)> = Vec::with_capacity(
params.len() + buffers.len()
);
for (name, p) in params {
targets.push((name.clone(), p.variable.shape(), p.variable.data().dtype()));
}
for (name, b) in buffers {
targets.push((name.clone(), b.shape(), b.get().dtype()));
}
let mut unchanged = Vec::new();
let mut remapped = Vec::new();
let mut missing = Vec::new();
let mut used = vec![false; entries.len()];
let mut output: Vec<(String, usize)> = Vec::new();
let name_index: std::collections::HashMap<&str, usize> =
entries.iter().enumerate().map(|(i, e)| (e.name.as_str(), i)).collect();
let mut unmatched: Vec<usize> = Vec::new();
for (mi, (name, shape, _)) in targets.iter().enumerate() {
if let Some(&ci) = name_index.get(name.as_str()) {
if !used[ci] && entries[ci].shape == *shape {
unchanged.push(name.clone());
used[ci] = true;
output.push((name.clone(), ci));
continue;
}
}
unmatched.push(mi);
}
for &mi in &unmatched {
let (name, shape, dtype) = &targets[mi];
let found = entries.iter().enumerate()
.find(|(ci, e)| !used[*ci] && e.shape == *shape && e.dtype == *dtype)
.map(|(ci, _)| ci);
if let Some(ci) = found {
remapped.push((entries[ci].name.clone(), name.clone()));
used[ci] = true;
output.push((name.clone(), ci));
} else {
missing.push(name.clone());
}
}
let dropped: Vec<String> = entries.iter().enumerate()
.filter(|(i, _)| !used[*i])
.map(|(_, e)| e.name.clone())
.collect();
w.write_all(&MAGIC).map_err(io_err)?;
w.write_all(&VERSION.to_le_bytes()).map_err(io_err)?;
w.write_all(&[0u8; HASH_LEN]).map_err(io_err)?;
w.write_all(&(output.len() as u32).to_le_bytes()).map_err(io_err)?;
for (name, ci) in &output {
write_raw_entry(w, name, &entries[*ci])?;
}
Ok(MigrateReport { unchanged, remapped, dropped, missing })
}
pub fn migrate_checkpoint_file(
src: &str,
dst: &str,
params: &[(String, Parameter)],
buffers: &[(String, Buffer)],
) -> Result<MigrateReport> {
let sf = std::fs::File::open(src).map_err(io_err)?;
write_file_atomic(dst, |mut w| {
if src.ends_with(".gz") {
let mut r = flate2::read::GzDecoder::new(sf);
migrate_checkpoint(&mut r, &mut w, params, buffers)
} else {
let mut r = std::io::BufReader::new(sf);
migrate_checkpoint(&mut r, &mut w, params, buffers)
}
})
}
pub(crate) fn io_err(e: impl std::fmt::Display) -> TensorError {
TensorError::new(&format!("io: {}", e))
}
const MAX_NAME_LEN: usize = 64 * 1024;
const MAX_NDIM: usize = 64;
fn read_name<R: Read>(r: &mut R) -> Result<String> {
let name_len = read_u32(r)? as usize;
if name_len > MAX_NAME_LEN {
return Err(TensorError::new(&format!(
"corrupt checkpoint: entry name length {name_len} exceeds {MAX_NAME_LEN}"
)));
}
let mut name_bytes = vec![0u8; name_len];
r.read_exact(&mut name_bytes).map_err(io_err)?;
Ok(String::from_utf8_lossy(&name_bytes).into_owned())
}
fn read_shape<R: Read>(r: &mut R) -> Result<Vec<i64>> {
let ndim = read_u32(r)? as usize;
if ndim > MAX_NDIM {
return Err(TensorError::new(&format!(
"corrupt checkpoint: tensor rank {ndim} exceeds {MAX_NDIM}"
)));
}
let mut shape = vec![0i64; ndim];
for s in &mut shape {
*s = read_i64(r)?;
}
Ok(shape)
}
fn read_payload<R: Read>(r: &mut R, byte_count: usize) -> Result<Vec<u8>> {
const PREALLOC_CAP: usize = 16 << 20;
let mut raw = Vec::with_capacity(byte_count.min(PREALLOC_CAP));
let n = r
.by_ref()
.take(byte_count as u64)
.read_to_end(&mut raw)
.map_err(io_err)?;
if n != byte_count {
return Err(TensorError::new(&format!(
"corrupt checkpoint: payload truncated: header claims {byte_count} bytes, \
{n} present"
)));
}
Ok(raw)
}
fn check_err_raw(err: *mut std::ffi::c_char) -> Result<()> {
if err.is_null() {
Ok(())
} else {
let msg = unsafe { std::ffi::CStr::from_ptr(err) }
.to_string_lossy()
.into_owned();
unsafe { flodl_sys::flodl_free_string(err) };
Err(TensorError::new(&msg))
}
}
fn read_u32<R: Read>(r: &mut R) -> Result<u32> {
let mut buf = [0u8; 4];
r.read_exact(&mut buf).map_err(io_err)?;
Ok(u32::from_le_bytes(buf))
}
fn read_u64<R: Read>(r: &mut R) -> Result<u64> {
let mut buf = [0u8; 8];
r.read_exact(&mut buf).map_err(io_err)?;
Ok(u64::from_le_bytes(buf))
}
fn read_i64<R: Read>(r: &mut R) -> Result<i64> {
let mut buf = [0u8; 8];
r.read_exact(&mut buf).map_err(io_err)?;
Ok(i64::from_le_bytes(buf))
}
pub(crate) fn read_f64_le<R: Read>(r: &mut R) -> Result<f64> {
let mut buf = [0u8; 8];
r.read_exact(&mut buf).map_err(io_err)?;
Ok(f64::from_le_bytes(buf))
}
pub(crate) fn write_f64_le<W: Write>(w: &mut W, v: f64) -> Result<()> {
w.write_all(&v.to_le_bytes()).map_err(io_err)?;
Ok(())
}
pub(crate) fn write_u32_le<W: Write>(w: &mut W, v: u32) -> Result<()> {
w.write_all(&v.to_le_bytes()).map_err(io_err)?;
Ok(())
}
pub(crate) fn write_i64_le<W: Write>(w: &mut W, v: i64) -> Result<()> {
w.write_all(&v.to_le_bytes()).map_err(io_err)?;
Ok(())
}
pub(crate) fn read_u32_le<R: Read>(r: &mut R) -> Result<u32> {
read_u32(r)
}
pub(crate) fn read_i64_le<R: Read>(r: &mut R) -> Result<i64> {
read_i64(r)
}
fn hex_to_bytes(hex: &str) -> Result<[u8; HASH_LEN]> {
if hex.len() != HASH_LEN * 2 {
return Err(TensorError::new(&format!(
"expected {} hex chars, got {}",
HASH_LEN * 2,
hex.len()
)));
}
let mut out = [0u8; HASH_LEN];
for (i, chunk) in hex.as_bytes().chunks(2).enumerate() {
let hi = hex_nibble(chunk[0])?;
let lo = hex_nibble(chunk[1])?;
out[i] = (hi << 4) | lo;
}
Ok(out)
}
fn hex_nibble(b: u8) -> Result<u8> {
match b {
b'0'..=b'9' => Ok(b - b'0'),
b'a'..=b'f' => Ok(b - b'a' + 10),
b'A'..=b'F' => Ok(b - b'A' + 10),
_ => Err(TensorError::new(&format!("invalid hex byte: {}", b))),
}
}
fn bytes_to_hex(bytes: &[u8]) -> String {
let mut s = String::with_capacity(bytes.len() * 2);
for &b in bytes {
use std::fmt::Write;
let _ = write!(s, "{:02x}", b);
}
s
}
#[cfg(test)]
#[path = "checkpoint_tests.rs"]
mod tests;