use std::collections::BTreeMap;
use std::fmt;
use serde_json::Value;
const MAX_HEADER_BYTES: u64 = 64 * 1024 * 1024;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub enum Dtype {
Bf16,
F32,
}
impl Dtype {
#[must_use]
pub const fn size(self) -> usize {
match self {
Self::Bf16 => 2,
Self::F32 => 4,
}
}
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Bf16 => "BF16",
Self::F32 => "F32",
}
}
fn parse(raw: &str) -> Option<Self> {
match raw {
"BF16" => Some(Self::Bf16),
"F32" => Some(Self::F32),
_ => None,
}
}
}
impl fmt::Display for Dtype {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TensorEntry {
pub name: String,
pub dtype: Dtype,
pub shape: Vec<usize>,
pub begin: usize,
pub end: usize,
}
impl TensorEntry {
#[must_use]
pub fn element_count(&self) -> usize {
self.shape.iter().product()
}
#[must_use]
pub const fn byte_len(&self) -> usize {
self.end - self.begin
}
#[must_use]
pub fn row_len(&self) -> usize {
self.shape.iter().skip(1).product()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WeightsError {
TooShortForHeader {
len: usize,
},
HeaderLengthOutOfRange {
declared: u64,
available: usize,
},
HeaderNotJson {
detail: String,
},
HeaderNotObject,
MalformedEntry {
name: String,
detail: String,
},
UnsupportedDtype {
name: String,
raw: String,
},
SpanOutOfBounds {
name: String,
begin: usize,
end: usize,
payload_len: usize,
},
ShapeSpanMismatch {
name: String,
shape: Vec<usize>,
expected_bytes: usize,
actual_bytes: usize,
},
ShapeOverflow {
name: String,
shape: Vec<usize>,
},
}
impl fmt::Display for WeightsError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::TooShortForHeader { len } => {
write!(f, "not a safetensors file: {len} bytes, need at least 8")
}
Self::HeaderLengthOutOfRange {
declared,
available,
} => write!(
f,
"header length {declared} is out of range (file has {available} bytes, cap is \
{MAX_HEADER_BYTES})"
),
Self::HeaderNotJson { detail } => write!(f, "header is not valid JSON: {detail}"),
Self::HeaderNotObject => f.write_str("header JSON is not an object"),
Self::MalformedEntry { name, detail } => {
write!(f, "tensor `{name}`: {detail}")
}
Self::UnsupportedDtype { name, raw } => write!(
f,
"tensor `{name}`: unsupported dtype `{raw}` (accepted: BF16, F32)"
),
Self::SpanOutOfBounds {
name,
begin,
end,
payload_len,
} => write!(
f,
"tensor `{name}`: byte span {begin}..{end} escapes the {payload_len}-byte payload"
),
Self::ShapeSpanMismatch {
name,
shape,
expected_bytes,
actual_bytes,
} => write!(
f,
"tensor `{name}`: shape {shape:?} implies {expected_bytes} bytes but the span \
covers {actual_bytes}"
),
Self::ShapeOverflow { name, shape } => {
write!(f, "tensor `{name}`: shape {shape:?} overflows usize")
}
}
}
}
impl std::error::Error for WeightsError {}
#[derive(Clone, Debug)]
pub struct SafetensorsIndex {
entries: BTreeMap<String, TensorEntry>,
payload_begin: usize,
}
impl SafetensorsIndex {
pub fn parse(bytes: &[u8]) -> Result<Self, WeightsError> {
let Some(len_prefix) = bytes.get(..8) else {
return Err(WeightsError::TooShortForHeader { len: bytes.len() });
};
let header_len = u64::from_le_bytes(
len_prefix
.try_into()
.expect("slice of 8 bytes converts to [u8; 8]"),
);
if header_len > MAX_HEADER_BYTES {
return Err(WeightsError::HeaderLengthOutOfRange {
declared: header_len,
available: bytes.len(),
});
}
let header_len_usize =
usize::try_from(header_len).map_err(|_| WeightsError::HeaderLengthOutOfRange {
declared: header_len,
available: bytes.len(),
})?;
let payload_begin =
8usize
.checked_add(header_len_usize)
.ok_or(WeightsError::HeaderLengthOutOfRange {
declared: header_len,
available: bytes.len(),
})?;
let Some(header_bytes) = bytes.get(8..payload_begin) else {
return Err(WeightsError::HeaderLengthOutOfRange {
declared: header_len,
available: bytes.len(),
});
};
let parsed: Value =
serde_json::from_slice(header_bytes).map_err(|error| WeightsError::HeaderNotJson {
detail: error.to_string(),
})?;
let Value::Object(directory) = parsed else {
return Err(WeightsError::HeaderNotObject);
};
let payload_len = bytes.len() - payload_begin;
let mut entries = BTreeMap::new();
for (name, value) in directory {
if name == "__metadata__" {
continue;
}
let entry = parse_entry(&name, &value, payload_begin, payload_len)?;
entries.insert(name, entry);
}
Ok(Self {
entries,
payload_begin,
})
}
#[must_use]
pub const fn payload_begin(&self) -> usize {
self.payload_begin
}
#[must_use]
pub fn len(&self) -> usize {
self.entries.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
#[must_use]
pub fn entry(&self, name: &str) -> Option<&TensorEntry> {
self.entries.get(name)
}
pub fn entries(&self) -> impl Iterator<Item = &TensorEntry> {
self.entries.values()
}
pub fn names(&self) -> impl Iterator<Item = &str> {
self.entries.keys().map(String::as_str)
}
#[must_use]
pub fn total_tensor_bytes(&self) -> usize {
self.entries.values().map(TensorEntry::byte_len).sum()
}
#[must_use]
pub fn view<'a>(&self, name: &str, bytes: &'a [u8]) -> Option<TensorView<'a>> {
let entry = self.entries.get(name)?;
let raw = bytes.get(entry.begin..entry.end)?;
Some(TensorView {
dtype: entry.dtype,
shape: entry.shape.clone(),
raw,
})
}
}
fn parse_entry(
name: &str,
value: &Value,
payload_begin: usize,
payload_len: usize,
) -> Result<TensorEntry, WeightsError> {
let object = value
.as_object()
.ok_or_else(|| WeightsError::MalformedEntry {
name: name.to_owned(),
detail: "entry is not a JSON object".to_owned(),
})?;
let raw_dtype = object.get("dtype").and_then(Value::as_str).ok_or_else(|| {
WeightsError::MalformedEntry {
name: name.to_owned(),
detail: "missing string field `dtype`".to_owned(),
}
})?;
let dtype = Dtype::parse(raw_dtype).ok_or_else(|| WeightsError::UnsupportedDtype {
name: name.to_owned(),
raw: raw_dtype.to_owned(),
})?;
let raw_shape = object
.get("shape")
.and_then(Value::as_array)
.ok_or_else(|| WeightsError::MalformedEntry {
name: name.to_owned(),
detail: "missing array field `shape`".to_owned(),
})?;
let mut shape = Vec::with_capacity(raw_shape.len());
for dim in raw_shape {
let dim = dim
.as_u64()
.and_then(|d| usize::try_from(d).ok())
.ok_or_else(|| WeightsError::MalformedEntry {
name: name.to_owned(),
detail: "shape contains a non-usize dimension".to_owned(),
})?;
shape.push(dim);
}
let offsets = object
.get("data_offsets")
.and_then(Value::as_array)
.ok_or_else(|| WeightsError::MalformedEntry {
name: name.to_owned(),
detail: "missing array field `data_offsets`".to_owned(),
})?;
if offsets.len() != 2 {
return Err(WeightsError::MalformedEntry {
name: name.to_owned(),
detail: format!("`data_offsets` has {} entries, expected 2", offsets.len()),
});
}
let mut bound = [0usize; 2];
for (slot, raw) in bound.iter_mut().zip(offsets) {
*slot = raw
.as_u64()
.and_then(|v| usize::try_from(v).ok())
.ok_or_else(|| WeightsError::MalformedEntry {
name: name.to_owned(),
detail: "`data_offsets` contains a non-usize value".to_owned(),
})?;
}
let [begin, end] = bound;
if begin > end || end > payload_len {
return Err(WeightsError::SpanOutOfBounds {
name: name.to_owned(),
begin,
end,
payload_len,
});
}
let mut elements = 1usize;
for dim in &shape {
elements = elements
.checked_mul(*dim)
.ok_or_else(|| WeightsError::ShapeOverflow {
name: name.to_owned(),
shape: shape.clone(),
})?;
}
let expected_bytes =
elements
.checked_mul(dtype.size())
.ok_or_else(|| WeightsError::ShapeOverflow {
name: name.to_owned(),
shape: shape.clone(),
})?;
let actual_bytes = end - begin;
if expected_bytes != actual_bytes {
return Err(WeightsError::ShapeSpanMismatch {
name: name.to_owned(),
shape,
expected_bytes,
actual_bytes,
});
}
Ok(TensorEntry {
name: name.to_owned(),
dtype,
shape,
begin: payload_begin + begin,
end: payload_begin + end,
})
}
#[derive(Clone, Copy, Debug)]
pub struct TensorViewRef<'a> {
dtype: Dtype,
raw: &'a [u8],
}
#[derive(Clone, Debug)]
pub struct TensorView<'a> {
dtype: Dtype,
shape: Vec<usize>,
raw: &'a [u8],
}
#[must_use]
pub const fn bf16_bits_to_f32(bits: u16) -> f32 {
f32::from_bits((bits as u32) << 16)
}
impl<'a> TensorView<'a> {
#[must_use]
pub const fn dtype(&self) -> Dtype {
self.dtype
}
#[must_use]
pub fn shape(&self) -> &[usize] {
&self.shape
}
#[must_use]
pub fn len(&self) -> usize {
self.raw.len() / self.dtype.size()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn row_len(&self) -> usize {
self.shape.iter().skip(1).product()
}
#[must_use]
pub fn get_f32(&self, index: usize) -> Option<f32> {
let size = self.dtype.size();
let start = index.checked_mul(size)?;
let chunk = self.raw.get(start..start.checked_add(size)?)?;
Some(match self.dtype {
Dtype::Bf16 => bf16_bits_to_f32(u16::from_le_bytes([chunk[0], chunk[1]])),
Dtype::F32 => f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]),
})
}
#[must_use]
pub fn copy_row_f32(&self, row: usize, out: &mut [f32]) -> bool {
let row_len = self.row_len();
if row_len == 0 || out.len() != row_len {
return false;
}
let Some(base) = row.checked_mul(row_len) else {
return false;
};
if base.checked_add(row_len).is_none_or(|end| end > self.len()) {
return false;
}
for (offset, slot) in out.iter_mut().enumerate() {
match self.get_f32(base + offset) {
Some(value) => *slot = value,
None => return false,
}
}
true
}
#[must_use]
pub const fn as_bytes(&self) -> &'a [u8] {
self.raw
}
#[must_use]
pub const fn as_ref(&self) -> TensorViewRef<'a> {
TensorViewRef {
dtype: self.dtype,
raw: self.raw,
}
}
}
impl TensorViewRef<'_> {
#[must_use]
pub const fn dtype(&self) -> Dtype {
self.dtype
}
#[must_use]
pub fn len(&self) -> usize {
self.raw.len() / self.dtype.size()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[cfg(test)]
mod tests {
use super::*;
fn build(parts: &[(&str, Dtype, &[usize], &[u8])]) -> Vec<u8> {
let mut directory = serde_json::Map::new();
let mut payload = Vec::new();
for (name, dtype, shape, bytes) in parts {
let begin = payload.len();
payload.extend_from_slice(bytes);
directory.insert(
(*name).to_owned(),
serde_json::json!({
"dtype": dtype.as_str(),
"shape": shape,
"data_offsets": [begin, payload.len()],
}),
);
}
assemble(&Value::Object(directory), &payload)
}
fn assemble(header: &Value, payload: &[u8]) -> Vec<u8> {
let header_bytes = serde_json::to_vec(header).expect("header serializes");
let mut out = Vec::new();
out.extend_from_slice(&(header_bytes.len() as u64).to_le_bytes());
out.extend_from_slice(&header_bytes);
out.extend_from_slice(payload);
out
}
fn bf16_payload(values: &[u16]) -> Vec<u8> {
values.iter().flat_map(|v| v.to_le_bytes()).collect()
}
#[test]
fn parses_a_two_tensor_directory() {
let buffer = build(&[
("a", Dtype::Bf16, &[2, 2], &bf16_payload(&[0, 1, 2, 3])),
("b", Dtype::F32, &[2], &1.0f32.to_le_bytes().repeat(2)),
]);
let index = SafetensorsIndex::parse(&buffer).expect("parses");
assert_eq!(index.len(), 2);
assert_eq!(index.names().collect::<Vec<_>>(), vec!["a", "b"]);
let a = index.entry("a").expect("entry a");
assert_eq!(a.dtype, Dtype::Bf16);
assert_eq!(a.shape, vec![2, 2]);
assert_eq!(a.element_count(), 4);
assert_eq!(a.byte_len(), 8);
assert_eq!(a.row_len(), 2);
assert_eq!(index.total_tensor_bytes(), 16);
}
#[test]
fn metadata_key_is_not_a_tensor() {
let mut directory = serde_json::Map::new();
directory.insert(
"__metadata__".to_owned(),
serde_json::json!({"format": "pt"}),
);
directory.insert(
"w".to_owned(),
serde_json::json!({"dtype": "F32", "shape": [1], "data_offsets": [0, 4]}),
);
let buffer = assemble(&Value::Object(directory), &1.0f32.to_le_bytes());
let index = SafetensorsIndex::parse(&buffer).expect("parses");
assert_eq!(index.len(), 1);
assert!(index.entry("__metadata__").is_none());
}
#[test]
fn widening_bf16_is_exact_for_representable_values() {
for bits in [0x0000u16, 0x3f80, 0xbf80, 0x7f80, 0xff80, 0x0001, 0x8000] {
let widened = bf16_bits_to_f32(bits);
assert_eq!(widened.to_bits() >> 16, u32::from(bits));
assert_eq!(widened.to_bits() & 0x0000_ffff, 0);
}
assert_eq!(bf16_bits_to_f32(0x3f80), 1.0);
assert_eq!(bf16_bits_to_f32(0xbf80), -1.0);
assert_eq!(bf16_bits_to_f32(0x0000), 0.0);
assert!(bf16_bits_to_f32(0x7f80).is_infinite());
assert!(bf16_bits_to_f32(0x7fc0).is_nan());
}
#[test]
fn bf16_widening_round_trips_every_bit_pattern() {
for bits in 0..=u16::MAX {
let widened = bf16_bits_to_f32(bits);
assert_eq!(
(widened.to_bits() >> 16) as u16,
bits,
"bit pattern {bits:#06x} did not survive widening"
);
let exponent = bits & 0x7f80;
let mantissa = bits & 0x007f;
if exponent == 0x7f80 && mantissa != 0 {
assert!(widened.is_nan(), "{bits:#06x} should widen to NaN");
} else {
assert!(!widened.is_nan(), "{bits:#06x} should not widen to NaN");
}
}
}
#[test]
fn reads_elements_and_rows_without_materializing() {
let payload = bf16_payload(&[0x3f80, 0xbf80, 0x4000, 0xc000]);
let buffer = build(&[("w", Dtype::Bf16, &[2, 2], &payload)]);
let index = SafetensorsIndex::parse(&buffer).expect("parses");
let view = index.view("w", &buffer).expect("view");
assert_eq!(view.len(), 4);
assert_eq!(view.row_len(), 2);
assert_eq!(view.get_f32(0), Some(1.0));
assert_eq!(view.get_f32(1), Some(-1.0));
assert_eq!(view.get_f32(3), Some(-2.0));
assert_eq!(view.get_f32(4), None);
let mut row = [0.0f32; 2];
assert!(view.copy_row_f32(1, &mut row));
assert_eq!(row, [2.0, -2.0]);
assert!(!view.copy_row_f32(2, &mut row));
let mut wrong = [0.0f32; 3];
assert!(!view.copy_row_f32(0, &mut wrong));
}
#[test]
fn refuses_a_truncated_file() {
assert_eq!(
SafetensorsIndex::parse(&[0u8; 4]).expect_err("must refuse"),
WeightsError::TooShortForHeader { len: 4 }
);
}
#[test]
fn refuses_an_absurd_header_length() {
let mut buffer = u64::MAX.to_le_bytes().to_vec();
buffer.extend_from_slice(b"{}");
let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
assert!(matches!(error, WeightsError::HeaderLengthOutOfRange { .. }));
}
#[test]
fn refuses_a_header_longer_than_the_file() {
let mut buffer = 4096u64.to_le_bytes().to_vec();
buffer.extend_from_slice(b"{}");
let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
assert!(matches!(error, WeightsError::HeaderLengthOutOfRange { .. }));
}
#[test]
fn refuses_malformed_json_and_non_objects() {
let mut buffer = 5u64.to_le_bytes().to_vec();
buffer.extend_from_slice(b"{ not");
assert!(matches!(
SafetensorsIndex::parse(&buffer).expect_err("must refuse"),
WeightsError::HeaderNotJson { .. }
));
let buffer = assemble(&serde_json::json!([1, 2]), &[]);
assert_eq!(
SafetensorsIndex::parse(&buffer).expect_err("must refuse"),
WeightsError::HeaderNotObject
);
}
#[test]
fn refuses_an_unsupported_dtype() {
let buffer = assemble(
&serde_json::json!({"w": {"dtype": "I64", "shape": [1], "data_offsets": [0, 8]}}),
&[0u8; 8],
);
assert_eq!(
SafetensorsIndex::parse(&buffer).expect_err("must refuse"),
WeightsError::UnsupportedDtype {
name: "w".to_owned(),
raw: "I64".to_owned(),
}
);
}
#[test]
fn refuses_a_span_past_the_payload() {
let buffer = assemble(
&serde_json::json!({"w": {"dtype": "F32", "shape": [16], "data_offsets": [0, 64]}}),
&[0u8; 8],
);
let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
assert!(matches!(error, WeightsError::SpanOutOfBounds { .. }));
}
#[test]
fn refuses_reversed_offsets() {
let buffer = assemble(
&serde_json::json!({"w": {"dtype": "F32", "shape": [1], "data_offsets": [8, 4]}}),
&[0u8; 8],
);
let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
assert!(matches!(error, WeightsError::SpanOutOfBounds { .. }));
}
#[test]
fn refuses_a_shape_that_disagrees_with_its_span() {
let buffer = assemble(
&serde_json::json!({"w": {"dtype": "F32", "shape": [4], "data_offsets": [0, 8]}}),
&[0u8; 8],
);
let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
assert!(
matches!(
error,
WeightsError::ShapeSpanMismatch {
expected_bytes: 16,
actual_bytes: 8,
..
}
),
"wrong error: {error}"
);
}
#[test]
fn refuses_a_shape_that_overflows() {
let huge = usize::MAX;
let buffer = assemble(
&serde_json::json!({
"w": {"dtype": "F32", "shape": [huge, huge], "data_offsets": [0, 8]}
}),
&[0u8; 8],
);
let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
assert!(matches!(error, WeightsError::ShapeOverflow { .. }));
}
#[test]
fn view_refuses_a_buffer_that_is_not_the_parsed_one() {
let buffer = build(&[("w", Dtype::F32, &[1], &1.0f32.to_le_bytes())]);
let index = SafetensorsIndex::parse(&buffer).expect("parses");
assert!(index.view("w", &buffer[..4]).is_none());
assert!(index.view("missing", &buffer).is_none());
}
}
#[derive(Debug)]
pub struct SafetensorsFile {
mapping: ftts_kernels::mmap::MappedFile,
index: SafetensorsIndex,
}
impl SafetensorsFile {
pub fn open(path: impl AsRef<std::path::Path>) -> Result<Self, OpenError> {
let mapping = ftts_kernels::mmap::MappedFile::open(path).map_err(OpenError::Io)?;
let index = SafetensorsIndex::parse(mapping.as_slice()).map_err(OpenError::Weights)?;
Ok(Self { mapping, index })
}
pub fn advise_random(&self) {
self.mapping.advise_random();
}
#[must_use]
pub const fn index(&self) -> &SafetensorsIndex {
&self.index
}
#[must_use]
pub const fn mapped_len(&self) -> usize {
self.mapping.len()
}
#[must_use]
pub fn view(&self, name: &str) -> Option<TensorView<'_>> {
self.index.view(name, self.mapping.as_slice())
}
}
#[derive(Debug)]
pub enum OpenError {
Io(std::io::Error),
Weights(WeightsError),
}
impl fmt::Display for OpenError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io(error) => write!(f, "cannot open checkpoint: {error}"),
Self::Weights(error) => write!(f, "{error}"),
}
}
}
impl std::error::Error for OpenError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Io(error) => Some(error),
Self::Weights(error) => Some(error),
}
}
}