use crate::error::InferenceError;
use crate::weights::safetensors_layout::{
SafetensorsLayoutEntry, safetensors_dtype, validate_safetensors_layout,
};
use memmap2::Mmap;
use serde_json::Value;
use std::collections::{HashMap, HashSet};
use std::fs::File;
use std::path::{Path, PathBuf};
use std::sync::OnceLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum DType {
F32,
F16,
BF16,
F8E4M3,
F8E5M2,
Other {
label: &'static str,
},
}
impl DType {
fn name(self) -> &'static str {
match self {
Self::F32 => "F32",
Self::F16 => "F16",
Self::BF16 => "BF16",
Self::F8E4M3 => "F8_E4M3",
Self::F8E5M2 => "F8_E5M2",
Self::Other { label, .. } => label,
}
}
}
fn dtype_from_str(s: &str) -> Option<DType> {
Some(match s {
"F32" => DType::F32,
"F16" => DType::F16,
"BF16" => DType::BF16,
"F8_E4M3" => DType::F8E4M3,
"F8_E5M2" => DType::F8E5M2,
_ => {
let dtype = safetensors_dtype(s)?;
DType::Other { label: dtype.name }
}
})
}
#[derive(Debug)]
struct TensorMeta {
dtype: DType,
shape: Vec<usize>,
start: usize,
end: usize,
converted_f32: OnceLock<Box<[f32]>>,
validated: OnceLock<Result<(), String>>,
}
#[derive(Debug, Clone, Copy)]
pub struct Tensor2D<'a> {
pub data: &'a [f32],
pub rows: usize,
pub cols: usize,
}
#[derive(Debug, Clone, Copy)]
pub struct Tensor1D<'a> {
pub data: &'a [f32],
pub len: usize,
}
#[derive(Debug, Clone)]
pub struct TransformerLayerWeights<'a> {
pub query_weight: Tensor2D<'a>,
pub query_bias: Tensor1D<'a>,
pub key_weight: Tensor2D<'a>,
pub key_bias: Tensor1D<'a>,
pub value_weight: Tensor2D<'a>,
pub value_bias: Tensor1D<'a>,
pub attn_output_weight: Tensor2D<'a>,
pub attn_output_bias: Tensor1D<'a>,
pub attn_layer_norm_weight: Tensor1D<'a>,
pub attn_layer_norm_bias: Tensor1D<'a>,
pub ffn_intermediate_weight: Tensor2D<'a>,
pub ffn_intermediate_bias: Tensor1D<'a>,
pub ffn_output_weight: Tensor2D<'a>,
pub ffn_output_bias: Tensor1D<'a>,
pub ffn_layer_norm_weight: Tensor1D<'a>,
pub ffn_layer_norm_bias: Tensor1D<'a>,
}
#[derive(Debug, Clone)]
pub struct BertWeights<'a> {
pub word_embeddings: Tensor2D<'a>,
pub position_embeddings: Tensor2D<'a>,
pub token_type_embeddings: Tensor2D<'a>,
pub embedding_layer_norm_weight: Tensor1D<'a>,
pub embedding_layer_norm_bias: Tensor1D<'a>,
pub layers: Vec<TransformerLayerWeights<'a>>,
pub pooler_weight: Tensor2D<'a>,
pub pooler_bias: Tensor1D<'a>,
}
#[derive(Debug, Clone)]
pub struct CrossEncoderWeights {
pub classifier_weight: Vec<f32>,
pub classifier_bias: f32,
}
impl CrossEncoderWeights {
pub fn logit(&self, pooled: &[f32]) -> f32 {
debug_assert_eq!(self.classifier_weight.len(), pooled.len());
self.classifier_weight
.iter()
.zip(pooled.iter())
.map(|(w, v)| w * v)
.sum::<f32>()
+ self.classifier_bias
}
}
enum SafetensorsBacking {
Mapped(Mmap),
Owned(Vec<u8>),
}
impl SafetensorsBacking {
fn as_slice(&self) -> &[u8] {
match self {
SafetensorsBacking::Mapped(m) => &m[..],
SafetensorsBacking::Owned(v) => v.as_slice(),
}
}
}
impl std::fmt::Debug for SafetensorsBacking {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SafetensorsBacking::Mapped(m) => f.debug_tuple("Mapped").field(m).finish(),
SafetensorsBacking::Owned(v) => f.debug_struct("Owned").field("len", &v.len()).finish(),
}
}
}
#[derive(Debug)]
pub struct SafetensorsFile {
data: SafetensorsBacking,
data_offset: usize,
tensors: HashMap<String, TensorMeta>,
source: String,
}
impl SafetensorsFile {
pub fn open(path: &Path) -> Result<Self, InferenceError> {
let file = File::open(path).map_err(|e| {
InferenceError::InvalidSafetensors(format!("failed to open {}: {e}", path.display()))
})?;
Self::from_open_file(file, path)
}
pub(crate) fn from_open_file(file: File, path: &Path) -> Result<Self, InferenceError> {
let display_path = path.display().to_string();
let mmap = crate::weights::mmap_trust::map_after_untrusted_open(&file, path)
.map_err(InferenceError::InvalidSafetensors)?;
Self::from_backing(SafetensorsBacking::Mapped(mmap), display_path)
}
pub fn from_bytes(bytes: Vec<u8>) -> Result<Self, InferenceError> {
Self::from_backing(
SafetensorsBacking::Owned(bytes),
"<in-memory bytes>".to_string(),
)
}
fn from_backing(data: SafetensorsBacking, source: String) -> Result<Self, InferenceError> {
let bytes = data.as_slice();
if bytes.len() < 8 {
return Err(InferenceError::InvalidSafetensors(
"file too small to contain safetensors header".into(),
));
}
let header_len = u64::from_le_bytes(
bytes[0..8]
.try_into()
.map_err(|_| InferenceError::InvalidSafetensors("invalid header length".into()))?,
) as usize;
if header_len > MAX_SAFETENSORS_HEADER_BYTES {
return Err(InferenceError::InvalidSafetensors(format!(
"header length {header_len} exceeds limit of {MAX_SAFETENSORS_HEADER_BYTES} bytes"
)));
}
let data_offset = 8usize
.checked_add(header_len)
.ok_or_else(|| InferenceError::InvalidSafetensors("header length overflow".into()))?;
if data_offset > bytes.len() {
return Err(InferenceError::InvalidSafetensors(format!(
"header extends past end of file: header_end={}, file_len={}",
data_offset,
bytes.len()
)));
}
let header = std::str::from_utf8(&bytes[8..data_offset]).map_err(|e| {
InferenceError::InvalidSafetensors(format!("header is not valid UTF-8: {e}"))
})?;
let tensors = parse_safetensors_header(header)?;
let data_len = bytes.len() - data_offset;
let layout: Vec<_> = tensors
.iter()
.map(|(name, meta)| SafetensorsLayoutEntry {
name,
dtype: meta.dtype.name(),
shape: &meta.shape,
start: meta.start,
end: meta.end,
})
.collect();
validate_safetensors_layout(&source, data_len, &layout)?;
Ok(Self {
data,
data_offset,
tensors,
source,
})
}
pub fn tensor_shape(&self, name: &str) -> Option<&[usize]> {
self.tensors.get(name).map(|m| m.shape.as_slice())
}
pub fn tensor_names(&self) -> Vec<&str> {
self.tensors.keys().map(String::as_str).collect()
}
pub fn has_tensor(&self, name: &str) -> bool {
self.tensors.contains_key(name)
}
pub fn tensor_dtype(&self, name: &str) -> Option<&'static str> {
self.tensors.get(name).map(|m| m.dtype.name())
}
pub fn get_f32_tensor(&self, name: &str) -> Result<(&[f32], &[usize]), InferenceError> {
let meta = self
.tensors
.get(name)
.ok_or_else(|| InferenceError::MissingTensor(name.to_string()))?;
let start = self
.data_offset
.checked_add(meta.start)
.ok_or_else(|| InferenceError::InvalidSafetensors("tensor start overflow".into()))?;
let end = self
.data_offset
.checked_add(meta.end)
.ok_or_else(|| InferenceError::InvalidSafetensors("tensor end overflow".into()))?;
let bytes = &self.data.as_slice()[start..end];
let source = self.source.as_str();
let shape = meta.shape.as_slice();
let dtype_name = meta.dtype.name();
let (slice, validation_is_fused): (&[f32], bool) = match meta.dtype {
DType::F32 => {
#[cfg(target_endian = "little")]
if bytes.as_ptr().align_offset(std::mem::align_of::<f32>()) == 0 {
(bytes_to_f32_slice(bytes), false)
} else {
(
meta.converted_f32
.get_or_init(|| copy_bytes_to_f32_owned(bytes).into_boxed_slice())
.as_ref(),
false,
)
}
#[cfg(not(target_endian = "little"))]
{
(
meta.converted_f32
.get_or_init(|| copy_bytes_to_f32_owned(bytes).into_boxed_slice())
.as_ref(),
false,
)
}
}
DType::F16 => {
#[cfg(feature = "f16")]
{
(
meta.converted_f32
.get_or_init(|| {
let (values, has_non_finite) = convert_f16_bytes_to_f32(bytes);
let _ = meta.validated.get_or_init(|| {
let tensor = if has_non_finite {
crate::weights::ingress::IngestedTensor::decoded_f32(
source, name, shape, dtype_name, &values,
)
} else {
crate::weights::ingress::IngestedTensor::decoded_f32_known_finite(
source, name, shape, dtype_name, &values,
)
};
crate::weights::ingress::validate_ingested_tensor(tensor)
.map_err(|e| e.to_string())
});
values.into_boxed_slice()
})
.as_ref(),
true,
)
}
#[cfg(not(feature = "f16"))]
{
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {name} is F16 but lattice-inference was built without the f16 feature"
)));
}
}
DType::BF16 => {
#[cfg(feature = "f16")]
{
(
meta.converted_f32
.get_or_init(|| {
let (values, has_non_finite) = convert_bf16_bytes_to_f32(bytes);
let _ = meta.validated.get_or_init(|| {
let tensor = if has_non_finite {
crate::weights::ingress::IngestedTensor::decoded_f32(
source, name, shape, dtype_name, &values,
)
} else {
crate::weights::ingress::IngestedTensor::decoded_f32_known_finite(
source, name, shape, dtype_name, &values,
)
};
crate::weights::ingress::validate_ingested_tensor(tensor)
.map_err(|e| e.to_string())
});
values.into_boxed_slice()
})
.as_ref(),
true,
)
}
#[cfg(not(feature = "f16"))]
{
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {name} is BF16 but lattice-inference was built without the f16 feature"
)));
}
}
DType::F8E4M3 => {
#[cfg(feature = "f16")]
{
(
meta.converted_f32
.get_or_init(|| {
let (values, has_non_finite) = convert_f8_e4m3_bytes_to_f32(bytes);
let _ = meta.validated.get_or_init(|| {
let tensor = if has_non_finite {
crate::weights::ingress::IngestedTensor::decoded_f32(
source, name, shape, dtype_name, &values,
)
} else {
crate::weights::ingress::IngestedTensor::decoded_f32_known_finite(
source, name, shape, dtype_name, &values,
)
};
crate::weights::ingress::validate_ingested_tensor(tensor)
.map_err(|e| e.to_string())
});
values.into_boxed_slice()
})
.as_ref(),
true,
)
}
#[cfg(not(feature = "f16"))]
{
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {name} is F8_E4M3 but lattice-inference was built without the f16 feature"
)));
}
}
DType::F8E5M2 => {
#[cfg(feature = "f16")]
{
(
meta.converted_f32
.get_or_init(|| {
let (values, has_non_finite) = convert_f8_e5m2_bytes_to_f32(bytes);
let _ = meta.validated.get_or_init(|| {
let tensor = if has_non_finite {
crate::weights::ingress::IngestedTensor::decoded_f32(
source, name, shape, dtype_name, &values,
)
} else {
crate::weights::ingress::IngestedTensor::decoded_f32_known_finite(
source, name, shape, dtype_name, &values,
)
};
crate::weights::ingress::validate_ingested_tensor(tensor)
.map_err(|e| e.to_string())
});
values.into_boxed_slice()
})
.as_ref(),
true,
)
}
#[cfg(not(feature = "f16"))]
{
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {name} is F8_E5M2 but lattice-inference was built without the f16 feature"
)));
}
}
DType::Other { label, .. } => {
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {name} has unsupported dtype {label} (source: {}); only F32, F16, \
BF16, F8_E4M3, and F8_E5M2 tensors can be materialized as f32",
self.source
)));
}
};
let validation = if validation_is_fused {
meta.validated.get().ok_or_else(|| {
InferenceError::InvalidSafetensors(format!(
"{source}: tensor {name} ({dtype_name}) widening completed without a \
validation result"
))
})?
} else {
meta.validated.get_or_init(|| {
crate::weights::ingress::validate_ingested_tensor(
crate::weights::ingress::IngestedTensor::decoded_f32(
source, name, shape, dtype_name, slice,
),
)
.map_err(|e| e.to_string())
})
};
match validation {
Ok(()) => {}
Err(msg) => return Err(InferenceError::InvalidSafetensors(msg.clone())),
}
Ok((slice, meta.shape.as_slice()))
}
pub fn load_bert_weights(
&self,
num_layers: usize,
hidden_size: usize,
) -> Result<BertWeights<'_>, InferenceError> {
let word_shape = self.require_shape("embeddings.word_embeddings.weight")?;
let position_shape = self.require_shape("embeddings.position_embeddings.weight")?;
let token_type_shape = self.require_shape("embeddings.token_type_embeddings.weight")?;
if word_shape.len() != 2 || word_shape[1] != hidden_size {
return Err(InferenceError::ShapeMismatch {
name: "embeddings.word_embeddings.weight".into(),
expected: vec![word_shape[0], hidden_size],
actual: word_shape.to_vec(),
});
}
if position_shape.len() != 2 || position_shape[1] != hidden_size {
return Err(InferenceError::ShapeMismatch {
name: "embeddings.position_embeddings.weight".into(),
expected: vec![position_shape[0], hidden_size],
actual: position_shape.to_vec(),
});
}
if token_type_shape.len() != 2 || token_type_shape[1] != hidden_size {
return Err(InferenceError::ShapeMismatch {
name: "embeddings.token_type_embeddings.weight".into(),
expected: vec![token_type_shape[0], hidden_size],
actual: token_type_shape.to_vec(),
});
}
let word_embeddings = self.tensor2d(
"embeddings.word_embeddings.weight",
word_shape[0],
word_shape[1],
)?;
let position_embeddings = self.tensor2d(
"embeddings.position_embeddings.weight",
position_shape[0],
position_shape[1],
)?;
let token_type_embeddings = self.tensor2d(
"embeddings.token_type_embeddings.weight",
token_type_shape[0],
token_type_shape[1],
)?;
let embedding_layer_norm_weight =
self.tensor1d("embeddings.LayerNorm.weight", hidden_size)?;
let embedding_layer_norm_bias = self.tensor1d("embeddings.LayerNorm.bias", hidden_size)?;
let intermediate_shape = self.require_shape("encoder.layer.0.intermediate.dense.weight")?;
if intermediate_shape.len() != 2 || intermediate_shape[1] != hidden_size {
return Err(InferenceError::ShapeMismatch {
name: "encoder.layer.0.intermediate.dense.weight".into(),
expected: vec![intermediate_shape[0], hidden_size],
actual: intermediate_shape.to_vec(),
});
}
let intermediate_size = intermediate_shape[0];
let mut layers = Vec::with_capacity(num_layers);
for i in 0..num_layers {
let prefix = format!("encoder.layer.{i}");
layers.push(TransformerLayerWeights {
query_weight: self.tensor2d(
&format!("{prefix}.attention.self.query.weight"),
hidden_size,
hidden_size,
)?,
query_bias: self
.tensor1d(&format!("{prefix}.attention.self.query.bias"), hidden_size)?,
key_weight: self.tensor2d(
&format!("{prefix}.attention.self.key.weight"),
hidden_size,
hidden_size,
)?,
key_bias: self
.tensor1d(&format!("{prefix}.attention.self.key.bias"), hidden_size)?,
value_weight: self.tensor2d(
&format!("{prefix}.attention.self.value.weight"),
hidden_size,
hidden_size,
)?,
value_bias: self
.tensor1d(&format!("{prefix}.attention.self.value.bias"), hidden_size)?,
attn_output_weight: self.tensor2d(
&format!("{prefix}.attention.output.dense.weight"),
hidden_size,
hidden_size,
)?,
attn_output_bias: self.tensor1d(
&format!("{prefix}.attention.output.dense.bias"),
hidden_size,
)?,
attn_layer_norm_weight: self.tensor1d(
&format!("{prefix}.attention.output.LayerNorm.weight"),
hidden_size,
)?,
attn_layer_norm_bias: self.tensor1d(
&format!("{prefix}.attention.output.LayerNorm.bias"),
hidden_size,
)?,
ffn_intermediate_weight: self.tensor2d(
&format!("{prefix}.intermediate.dense.weight"),
intermediate_size,
hidden_size,
)?,
ffn_intermediate_bias: self.tensor1d(
&format!("{prefix}.intermediate.dense.bias"),
intermediate_size,
)?,
ffn_output_weight: self.tensor2d(
&format!("{prefix}.output.dense.weight"),
hidden_size,
intermediate_size,
)?,
ffn_output_bias: self
.tensor1d(&format!("{prefix}.output.dense.bias"), hidden_size)?,
ffn_layer_norm_weight: self
.tensor1d(&format!("{prefix}.output.LayerNorm.weight"), hidden_size)?,
ffn_layer_norm_bias: self
.tensor1d(&format!("{prefix}.output.LayerNorm.bias"), hidden_size)?,
});
}
let pooler_weight =
self.tensor2d_optional("pooler.dense.weight", hidden_size, hidden_size)?;
let pooler_bias = self.tensor1d_optional("pooler.dense.bias", hidden_size)?;
Ok(BertWeights {
word_embeddings,
position_embeddings,
token_type_embeddings,
embedding_layer_norm_weight,
embedding_layer_norm_bias,
layers,
pooler_weight,
pooler_bias,
})
}
pub fn load_cross_encoder_weights(
&self,
hidden_size: usize,
) -> Result<CrossEncoderWeights, InferenceError> {
let (weight, weight_shape) = self.get_f32_tensor("classifier.weight")?;
let classifier_weight = match weight_shape {
[1, cols] if *cols == hidden_size => weight.to_vec(),
[cols] if *cols == hidden_size => weight.to_vec(),
_ => {
return Err(InferenceError::ShapeMismatch {
name: "classifier.weight".to_string(),
expected: vec![1, hidden_size],
actual: weight_shape.to_vec(),
});
}
};
let (bias, bias_shape) = self.get_f32_tensor("classifier.bias")?;
if bias_shape != [1usize] {
return Err(InferenceError::ShapeMismatch {
name: "classifier.bias".to_string(),
expected: vec![1],
actual: bias_shape.to_vec(),
});
}
let classifier_bias = bias[0];
Ok(CrossEncoderWeights {
classifier_weight,
classifier_bias,
})
}
fn require_shape(&self, name: &str) -> Result<&[usize], InferenceError> {
self.tensor_shape(name)
.ok_or_else(|| InferenceError::MissingTensor(name.to_string()))
}
fn tensor1d(&self, name: &str, len: usize) -> Result<Tensor1D<'_>, InferenceError> {
let (data, shape) = self.get_f32_tensor(name)?;
if shape != [len] {
return Err(InferenceError::ShapeMismatch {
name: name.to_string(),
expected: vec![len],
actual: shape.to_vec(),
});
}
Ok(Tensor1D { data, len })
}
fn tensor2d(
&self,
name: &str,
rows: usize,
cols: usize,
) -> Result<Tensor2D<'_>, InferenceError> {
let (data, shape) = self.get_f32_tensor(name)?;
if shape != [rows, cols] {
return Err(InferenceError::ShapeMismatch {
name: name.to_string(),
expected: vec![rows, cols],
actual: shape.to_vec(),
});
}
Ok(Tensor2D { data, rows, cols })
}
fn tensor1d_optional(&self, name: &str, len: usize) -> Result<Tensor1D<'_>, InferenceError> {
if self.has_tensor(name) {
self.tensor1d(name, len)
} else {
Ok(Tensor1D { data: &[], len: 0 })
}
}
fn tensor2d_optional(
&self,
name: &str,
rows: usize,
cols: usize,
) -> Result<Tensor2D<'_>, InferenceError> {
if self.has_tensor(name) {
self.tensor2d(name, rows, cols)
} else {
Ok(Tensor2D {
data: &[],
rows: 0,
cols: 0,
})
}
}
}
#[derive(Debug, Clone)]
pub struct QwenLayerWeights<'a> {
pub q_proj_weight: Tensor2D<'a>,
pub k_proj_weight: Tensor2D<'a>,
pub v_proj_weight: Tensor2D<'a>,
pub o_proj_weight: Tensor2D<'a>,
pub q_norm_weight: Tensor1D<'a>,
pub k_norm_weight: Tensor1D<'a>,
pub input_layernorm_weight: Tensor1D<'a>,
pub gate_proj_weight: Tensor2D<'a>,
pub up_proj_weight: Tensor2D<'a>,
pub down_proj_weight: Tensor2D<'a>,
pub post_attention_layernorm_weight: Tensor1D<'a>,
pub fused_qkv: Vec<f32>,
pub qkv_out_dim: usize,
pub fused_gate_up: Vec<f32>,
pub gate_up_out_dim: usize,
}
#[derive(Debug, Clone)]
pub struct QwenWeights<'a> {
pub embed_tokens: Tensor2D<'a>,
pub norm_weight: Tensor1D<'a>,
pub layers: Vec<QwenLayerWeights<'a>>,
}
impl SafetensorsFile {
pub fn load_qwen_weights(
&self,
num_layers: usize,
hidden_size: usize,
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
intermediate_size: usize,
) -> Result<QwenWeights<'_>, InferenceError> {
let q_dim = num_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
let embed_shape = self.require_shape("embed_tokens.weight")?;
if embed_shape.len() != 2 || embed_shape[1] != hidden_size {
return Err(InferenceError::ShapeMismatch {
name: "embed_tokens.weight".into(),
expected: vec![embed_shape[0], hidden_size],
actual: embed_shape.to_vec(),
});
}
let embed_tokens = self.tensor2d("embed_tokens.weight", embed_shape[0], hidden_size)?;
let norm_weight = self.tensor1d("norm.weight", hidden_size)?;
let qkv_out_dim = q_dim + 2 * kv_dim;
let gate_up_out_dim = 2 * intermediate_size;
let mut layers = Vec::with_capacity(num_layers);
for i in 0..num_layers {
let p = format!("layers.{i}");
let q_proj_weight =
self.tensor2d(&format!("{p}.self_attn.q_proj.weight"), q_dim, hidden_size)?;
let k_proj_weight =
self.tensor2d(&format!("{p}.self_attn.k_proj.weight"), kv_dim, hidden_size)?;
let v_proj_weight =
self.tensor2d(&format!("{p}.self_attn.v_proj.weight"), kv_dim, hidden_size)?;
let gate_proj_weight = self.tensor2d(
&format!("{p}.mlp.gate_proj.weight"),
intermediate_size,
hidden_size,
)?;
let up_proj_weight = self.tensor2d(
&format!("{p}.mlp.up_proj.weight"),
intermediate_size,
hidden_size,
)?;
let mut fused_qkv = Vec::with_capacity(qkv_out_dim * hidden_size);
fused_qkv.extend_from_slice(q_proj_weight.data); fused_qkv.extend_from_slice(k_proj_weight.data); fused_qkv.extend_from_slice(v_proj_weight.data);
let mut fused_gate_up = Vec::with_capacity(gate_up_out_dim * hidden_size);
fused_gate_up.extend_from_slice(gate_proj_weight.data); fused_gate_up.extend_from_slice(up_proj_weight.data);
layers.push(QwenLayerWeights {
q_proj_weight,
k_proj_weight,
v_proj_weight,
o_proj_weight: self.tensor2d(
&format!("{p}.self_attn.o_proj.weight"),
hidden_size,
q_dim,
)?,
q_norm_weight: self.tensor1d(&format!("{p}.self_attn.q_norm.weight"), head_dim)?,
k_norm_weight: self.tensor1d(&format!("{p}.self_attn.k_norm.weight"), head_dim)?,
input_layernorm_weight: self
.tensor1d(&format!("{p}.input_layernorm.weight"), hidden_size)?,
gate_proj_weight,
up_proj_weight,
down_proj_weight: self.tensor2d(
&format!("{p}.mlp.down_proj.weight"),
hidden_size,
intermediate_size,
)?,
post_attention_layernorm_weight: self
.tensor1d(&format!("{p}.post_attention_layernorm.weight"), hidden_size)?,
fused_qkv,
qkv_out_dim,
fused_gate_up,
gate_up_out_dim,
});
}
Ok(QwenWeights {
embed_tokens,
norm_weight,
layers,
})
}
}
fn bytes_to_f32_slice(bytes: &[u8]) -> &[f32] {
assert!(bytes.len().is_multiple_of(4));
assert!(bytes.as_ptr().align_offset(std::mem::align_of::<f32>()) == 0);
unsafe { std::slice::from_raw_parts(bytes.as_ptr() as *const f32, bytes.len() / 4) }
}
fn copy_bytes_to_f32_owned(bytes: &[u8]) -> Vec<f32> {
debug_assert_eq!(bytes.len() % 4, 0);
let mut out = Vec::with_capacity(bytes.len() / 4);
for chunk in bytes.chunks_exact(4) {
out.push(f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]));
}
out
}
#[cfg(feature = "f16")]
fn convert_f16_bytes_to_f32(bytes: &[u8]) -> (Vec<f32>, bool) {
debug_assert_eq!(bytes.len() % 2, 0);
let mut out = Vec::with_capacity(bytes.len() / 2);
let mut has_non_finite = false;
for chunk in bytes.chunks_exact(2) {
let value =
crate::weights::half_bits::f16_bits_to_f32(u16::from_le_bytes([chunk[0], chunk[1]]));
has_non_finite |= !value.is_finite();
out.push(value);
}
(out, has_non_finite)
}
#[cfg(feature = "f16")]
fn convert_bf16_bytes_to_f32(bytes: &[u8]) -> (Vec<f32>, bool) {
debug_assert_eq!(bytes.len() % 2, 0);
let mut out = Vec::with_capacity(bytes.len() / 2);
let mut has_non_finite = false;
for chunk in bytes.chunks_exact(2) {
let value =
crate::weights::half_bits::bf16_bits_to_f32(u16::from_le_bytes([chunk[0], chunk[1]]));
has_non_finite |= !value.is_finite();
out.push(value);
}
(out, has_non_finite)
}
#[cfg(feature = "f16")]
fn convert_f8_e4m3_bytes_to_f32(bytes: &[u8]) -> (Vec<f32>, bool) {
let mut out = Vec::with_capacity(bytes.len());
let mut has_non_finite = false;
for &byte in bytes {
let value = crate::weights::half_bits::f8_e4m3_bits_to_f32(byte);
has_non_finite |= !value.is_finite();
out.push(value);
}
(out, has_non_finite)
}
#[cfg(feature = "f16")]
fn convert_f8_e5m2_bytes_to_f32(bytes: &[u8]) -> (Vec<f32>, bool) {
let mut out = Vec::with_capacity(bytes.len());
let mut has_non_finite = false;
for &byte in bytes {
let value = crate::weights::half_bits::f8_e5m2_bits_to_f32(byte);
has_non_finite |= !value.is_finite();
out.push(value);
}
(out, has_non_finite)
}
#[derive(Debug)]
pub struct ShardedQwenBacking {
pub embed_tokens: Vec<f32>,
pub norm_weight: Vec<f32>,
pub q_proj: Vec<Vec<f32>>,
pub k_proj: Vec<Vec<f32>>,
pub v_proj: Vec<Vec<f32>>,
pub o_proj: Vec<Vec<f32>>,
pub q_norm: Vec<Vec<f32>>,
pub k_norm: Vec<Vec<f32>>,
pub input_ln: Vec<Vec<f32>>,
pub gate_proj: Vec<Vec<f32>>,
pub up_proj: Vec<Vec<f32>>,
pub down_proj: Vec<Vec<f32>>,
pub post_ln: Vec<Vec<f32>>,
}
pub(crate) const MAX_SAFETENSORS_INDEX_BYTES: u64 = 67_108_864;
pub(crate) const MAX_WEIGHT_MAP_ENTRIES: usize = 1_000_000;
const MAX_SAFETENSORS_HEADER_BYTES: usize = 4_194_304;
const MAX_SAFETENSORS_HEADER_DEPTH: usize = 32;
#[derive(Debug, Clone, serde::Deserialize)]
pub struct SafetensorsIndex {
#[serde(default)]
pub metadata: Value,
#[serde(deserialize_with = "deserialize_weight_map_no_duplicates")]
pub weight_map: HashMap<String, String>,
}
fn deserialize_weight_map_no_duplicates<'de, D>(
deserializer: D,
) -> Result<HashMap<String, String>, D::Error>
where
D: serde::Deserializer<'de>,
{
struct WeightMapVisitor;
impl<'de> serde::de::Visitor<'de> for WeightMapVisitor {
type Value = HashMap<String, String>;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a map of tensor name to shard filename")
}
fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
where
A: serde::de::MapAccess<'de>,
{
let reserve = map.size_hint().unwrap_or(0).min(MAX_WEIGHT_MAP_ENTRIES);
let mut out = HashMap::with_capacity(reserve);
while let Some((name, shard)) = map.next_entry::<String, String>()? {
if out.len() >= MAX_WEIGHT_MAP_ENTRIES {
return Err(serde::de::Error::custom(format!(
"weight_map exceeds MAX_WEIGHT_MAP_ENTRIES ({MAX_WEIGHT_MAP_ENTRIES})"
)));
}
if out.insert(name.clone(), shard).is_some() {
return Err(serde::de::Error::custom(format!(
"duplicate tensor name in weight_map: {name}"
)));
}
}
Ok(out)
}
}
deserializer.deserialize_map(WeightMapVisitor)
}
#[derive(Debug, Clone)]
pub struct Tensor {
pub data: Vec<f32>,
pub shape: Vec<usize>,
}
pub trait TensorSource {
fn has_tensor(&mut self, name: &str) -> Result<bool, InferenceError>;
fn tensor_shape(&mut self, name: &str) -> Result<Option<Vec<usize>>, InferenceError>;
fn get_f32_tensor_owned(
&mut self,
name: &str,
) -> Result<(Vec<f32>, Vec<usize>), InferenceError>;
fn tensor_dtype(&mut self, _name: &str) -> Result<Option<String>, InferenceError> {
Ok(None)
}
}
impl TensorSource for SafetensorsFile {
fn has_tensor(&mut self, name: &str) -> Result<bool, InferenceError> {
Ok(SafetensorsFile::has_tensor(self, name))
}
fn tensor_shape(&mut self, name: &str) -> Result<Option<Vec<usize>>, InferenceError> {
Ok(SafetensorsFile::tensor_shape(self, name).map(<[usize]>::to_vec))
}
fn tensor_dtype(&mut self, name: &str) -> Result<Option<String>, InferenceError> {
Ok(SafetensorsFile::tensor_dtype(self, name).map(str::to_string))
}
fn get_f32_tensor_owned(
&mut self,
name: &str,
) -> Result<(Vec<f32>, Vec<usize>), InferenceError> {
let (data, shape) = self.get_f32_tensor(name)?;
Ok((data.to_vec(), shape.to_vec()))
}
}
#[derive(Debug)]
pub struct ShardedSafetensors {
root: PathBuf,
index: SafetensorsIndex,
shards: HashMap<String, SafetensorsFile>,
}
pub fn contained_shard_path(model_dir: &Path, shard_file: &str) -> Result<PathBuf, InferenceError> {
let rel = Path::new(shard_file);
if rel.as_os_str().is_empty() {
return Err(InferenceError::InvalidSafetensors(
"shard entry is empty; index entries must name a file within the model directory"
.to_string(),
));
}
if rel.is_absolute() {
return Err(InferenceError::InvalidSafetensors(format!(
"shard entry {shard_file:?} is an absolute path; \
index entries must stay within the model directory"
)));
}
for component in rel.components() {
match component {
std::path::Component::Normal(_) | std::path::Component::CurDir => {}
_ => {
return Err(InferenceError::InvalidSafetensors(format!(
"shard entry {shard_file:?} escapes the model directory; \
index entries must stay within the model directory"
)));
}
}
}
Ok(model_dir.join(rel))
}
pub fn parse_index(model_dir: &Path) -> Result<SafetensorsIndex, InferenceError> {
let index_path = model_dir.join("model.safetensors.index.json");
let file_len = std::fs::metadata(&index_path)
.map_err(InferenceError::Io)?
.len();
if file_len > MAX_SAFETENSORS_INDEX_BYTES {
return Err(InferenceError::InvalidSafetensors(format!(
"{} is {file_len} bytes, exceeding MAX_SAFETENSORS_INDEX_BYTES \
({MAX_SAFETENSORS_INDEX_BYTES})",
index_path.display()
)));
}
let json = std::fs::read_to_string(&index_path).map_err(InferenceError::Io)?;
serde_json::from_str(&json).map_err(|e| {
InferenceError::InvalidSafetensors(format!("failed to parse {}: {e}", index_path.display()))
})
}
pub(crate) fn open_manifest_entry_once(
model_root: &Path,
entry_name: &str,
) -> Result<(File, PathBuf), InferenceError> {
let candidate = contained_shard_path(model_root, entry_name)?;
let file = File::open(&candidate).map_err(|e| {
InferenceError::InvalidSafetensors(format!("failed to open {}: {e}", candidate.display()))
})?;
let real_path = real_path_of_open_file(&file, &candidate)?;
crate::weights::mmap_trust::reject_if_open_mmap_file_untrusted(&file, &real_path)
.map_err(InferenceError::InvalidSafetensors)?;
Ok((file, real_path))
}
#[cfg(target_os = "macos")]
fn real_path_of_open_file(file: &File, _candidate: &Path) -> Result<PathBuf, InferenceError> {
use std::os::unix::io::AsRawFd;
let mut buf = [0u8; libc::PATH_MAX as usize];
let ret = unsafe { libc::fcntl(file.as_raw_fd(), libc::F_GETPATH, buf.as_mut_ptr()) };
if ret == -1 {
return Err(InferenceError::Io(std::io::Error::last_os_error()));
}
let len = buf.iter().position(|&b| b == 0).unwrap_or(buf.len());
let s = std::str::from_utf8(&buf[..len]).map_err(|_| {
InferenceError::Inference("F_GETPATH returned a non-UTF-8 path".to_string())
})?;
Ok(PathBuf::from(s))
}
#[cfg(not(target_os = "macos"))]
fn real_path_of_open_file(_file: &File, candidate: &Path) -> Result<PathBuf, InferenceError> {
candidate.canonicalize().map_err(InferenceError::Io)
}
pub fn resolve_shard(
index: &SafetensorsIndex,
model_dir: &Path,
tensor_name: &str,
) -> Result<PathBuf, InferenceError> {
let shard_file = index
.weight_map
.get(tensor_name)
.ok_or_else(|| InferenceError::MissingTensor(tensor_name.to_string()))?;
contained_shard_path(model_dir, shard_file)
}
pub fn load_sharded(model_dir: &Path) -> Result<HashMap<String, Tensor>, InferenceError> {
let index = parse_index(model_dir)?;
let mut by_shard: std::collections::BTreeMap<String, Vec<String>> =
std::collections::BTreeMap::new();
for (tensor_name, shard_file) in &index.weight_map {
by_shard
.entry(shard_file.clone())
.or_default()
.push(tensor_name.clone());
}
let mut tensors = HashMap::with_capacity(index.weight_map.len());
for (shard_file, tensor_names) in by_shard {
let (file, real_path) = open_manifest_entry_once(model_dir, &shard_file)?;
let shard = SafetensorsFile::from_open_file(file, &real_path)?;
for tensor_name in tensor_names {
let (data, shape) = shard.get_f32_tensor(&tensor_name)?;
tensors.insert(
tensor_name,
Tensor {
data: data.to_vec(),
shape: shape.to_vec(),
},
);
}
}
Ok(tensors)
}
impl ShardedSafetensors {
pub fn open_index(index_path: &Path) -> Result<Self, InferenceError> {
let root = index_path
.parent()
.unwrap_or_else(|| Path::new("."))
.to_path_buf();
let index = parse_index(&root)?;
Ok(Self {
root,
index,
shards: HashMap::new(),
})
}
pub fn index(&self) -> &SafetensorsIndex {
&self.index
}
pub fn resolve_weight(&self, tensor_name: &str) -> Result<(PathBuf, String), InferenceError> {
let shard_file = self
.index
.weight_map
.get(tensor_name)
.ok_or_else(|| InferenceError::MissingTensor(tensor_name.to_string()))?;
let shard_path = contained_shard_path(&self.root, shard_file)?;
Ok((shard_path, tensor_name.to_string()))
}
fn shard_file_for(&self, name: &str) -> Result<String, InferenceError> {
self.index
.weight_map
.get(name)
.cloned()
.ok_or_else(|| InferenceError::MissingTensor(name.to_string()))
}
fn open_shard(&mut self, shard_file: &str) -> Result<&SafetensorsFile, InferenceError> {
if !self.shards.contains_key(shard_file) {
let (file, real_path) = open_manifest_entry_once(&self.root, shard_file)?;
let shard = SafetensorsFile::from_open_file(file, &real_path)?;
self.shards.insert(shard_file.to_string(), shard);
}
self.shards.get(shard_file).ok_or_else(|| {
InferenceError::InvalidSafetensors(format!("failed to cache shard {shard_file}"))
})
}
}
impl TensorSource for ShardedSafetensors {
fn has_tensor(&mut self, name: &str) -> Result<bool, InferenceError> {
Ok(self.index.weight_map.contains_key(name))
}
fn tensor_shape(&mut self, name: &str) -> Result<Option<Vec<usize>>, InferenceError> {
let shard_file = self.shard_file_for(name)?;
let shard = self.open_shard(&shard_file)?;
Ok(shard.tensor_shape(name).map(<[usize]>::to_vec))
}
fn get_f32_tensor_owned(
&mut self,
name: &str,
) -> Result<(Vec<f32>, Vec<usize>), InferenceError> {
let shard_file = self.shard_file_for(name)?;
let shard = self.open_shard(&shard_file)?;
let (data, shape) = shard.get_f32_tensor(name)?;
Ok((data.to_vec(), shape.to_vec()))
}
}
impl ShardedSafetensors {
pub fn load_qwen_weights_owned(
&mut self,
num_layers: usize,
hidden_size: usize,
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
intermediate_size: usize,
) -> Result<(Box<ShardedQwenBacking>, QwenWeights<'static>), InferenceError> {
let q_dim = num_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
let qkv_out_dim = q_dim + 2 * kv_dim;
let gate_up_out_dim = 2 * intermediate_size;
let (embed_data, embed_shape) = self.get_f32_tensor_owned("embed_tokens.weight")?;
if embed_shape.len() != 2 || embed_shape[1] != hidden_size {
return Err(InferenceError::ShapeMismatch {
name: "embed_tokens.weight".into(),
expected: vec![embed_shape[0], hidden_size],
actual: embed_shape,
});
}
let embed_vocab = embed_shape[0];
let (norm_data, norm_shape) = self.get_f32_tensor_owned("norm.weight")?;
if norm_shape != [hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: "norm.weight".into(),
expected: vec![hidden_size],
actual: norm_shape,
});
}
let mut q_proj_vecs = Vec::with_capacity(num_layers);
let mut k_proj_vecs = Vec::with_capacity(num_layers);
let mut v_proj_vecs = Vec::with_capacity(num_layers);
let mut o_proj_vecs = Vec::with_capacity(num_layers);
let mut q_norm_vecs = Vec::with_capacity(num_layers);
let mut k_norm_vecs = Vec::with_capacity(num_layers);
let mut input_ln_vecs = Vec::with_capacity(num_layers);
let mut gate_proj_vecs = Vec::with_capacity(num_layers);
let mut up_proj_vecs = Vec::with_capacity(num_layers);
let mut down_proj_vecs = Vec::with_capacity(num_layers);
let mut post_ln_vecs = Vec::with_capacity(num_layers);
for i in 0..num_layers {
let p = format!("layers.{i}");
let (q, qs) = self.get_f32_tensor_owned(&format!("{p}.self_attn.q_proj.weight"))?;
if qs != [q_dim, hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.self_attn.q_proj.weight"),
expected: vec![q_dim, hidden_size],
actual: qs,
});
}
q_proj_vecs.push(q);
let (k, ks) = self.get_f32_tensor_owned(&format!("{p}.self_attn.k_proj.weight"))?;
if ks != [kv_dim, hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.self_attn.k_proj.weight"),
expected: vec![kv_dim, hidden_size],
actual: ks,
});
}
k_proj_vecs.push(k);
let (v, vs) = self.get_f32_tensor_owned(&format!("{p}.self_attn.v_proj.weight"))?;
if vs != [kv_dim, hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.self_attn.v_proj.weight"),
expected: vec![kv_dim, hidden_size],
actual: vs,
});
}
v_proj_vecs.push(v);
let (o, os) = self.get_f32_tensor_owned(&format!("{p}.self_attn.o_proj.weight"))?;
if os != [hidden_size, q_dim] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.self_attn.o_proj.weight"),
expected: vec![hidden_size, q_dim],
actual: os,
});
}
o_proj_vecs.push(o);
let (qn, qns) = self.get_f32_tensor_owned(&format!("{p}.self_attn.q_norm.weight"))?;
if qns != [head_dim] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.self_attn.q_norm.weight"),
expected: vec![head_dim],
actual: qns,
});
}
q_norm_vecs.push(qn);
let (kn, kns) = self.get_f32_tensor_owned(&format!("{p}.self_attn.k_norm.weight"))?;
if kns != [head_dim] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.self_attn.k_norm.weight"),
expected: vec![head_dim],
actual: kns,
});
}
k_norm_vecs.push(kn);
let (iln, ilns) = self.get_f32_tensor_owned(&format!("{p}.input_layernorm.weight"))?;
if ilns != [hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.input_layernorm.weight"),
expected: vec![hidden_size],
actual: ilns,
});
}
input_ln_vecs.push(iln);
let (g, gs) = self.get_f32_tensor_owned(&format!("{p}.mlp.gate_proj.weight"))?;
if gs != [intermediate_size, hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.mlp.gate_proj.weight"),
expected: vec![intermediate_size, hidden_size],
actual: gs,
});
}
gate_proj_vecs.push(g);
let (u, us) = self.get_f32_tensor_owned(&format!("{p}.mlp.up_proj.weight"))?;
if us != [intermediate_size, hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.mlp.up_proj.weight"),
expected: vec![intermediate_size, hidden_size],
actual: us,
});
}
up_proj_vecs.push(u);
let (d, ds) = self.get_f32_tensor_owned(&format!("{p}.mlp.down_proj.weight"))?;
if ds != [hidden_size, intermediate_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.mlp.down_proj.weight"),
expected: vec![hidden_size, intermediate_size],
actual: ds,
});
}
down_proj_vecs.push(d);
let (pln, plns) =
self.get_f32_tensor_owned(&format!("{p}.post_attention_layernorm.weight"))?;
if plns != [hidden_size] {
return Err(InferenceError::ShapeMismatch {
name: format!("{p}.post_attention_layernorm.weight"),
expected: vec![hidden_size],
actual: plns,
});
}
post_ln_vecs.push(pln);
}
let backing = Box::new(ShardedQwenBacking {
embed_tokens: embed_data,
norm_weight: norm_data,
q_proj: q_proj_vecs,
k_proj: k_proj_vecs,
v_proj: v_proj_vecs,
o_proj: o_proj_vecs,
q_norm: q_norm_vecs,
k_norm: k_norm_vecs,
input_ln: input_ln_vecs,
gate_proj: gate_proj_vecs,
up_proj: up_proj_vecs,
down_proj: down_proj_vecs,
post_ln: post_ln_vecs,
});
let weights: QwenWeights<'static> = {
#[allow(clippy::explicit_auto_deref)]
let b: &ShardedQwenBacking = &*backing;
let b: &'static ShardedQwenBacking = unsafe { &*(b as *const ShardedQwenBacking) };
let embed_tokens = Tensor2D {
data: &b.embed_tokens,
rows: embed_vocab,
cols: hidden_size,
};
let norm_weight = Tensor1D {
data: &b.norm_weight,
len: hidden_size,
};
let mut layers = Vec::with_capacity(num_layers);
for i in 0..num_layers {
let mut fused_qkv = Vec::with_capacity(qkv_out_dim * hidden_size);
fused_qkv.extend_from_slice(&b.q_proj[i]);
fused_qkv.extend_from_slice(&b.k_proj[i]);
fused_qkv.extend_from_slice(&b.v_proj[i]);
let mut fused_gate_up = Vec::with_capacity(gate_up_out_dim * hidden_size);
fused_gate_up.extend_from_slice(&b.gate_proj[i]);
fused_gate_up.extend_from_slice(&b.up_proj[i]);
layers.push(QwenLayerWeights {
q_proj_weight: Tensor2D {
data: &b.q_proj[i],
rows: q_dim,
cols: hidden_size,
},
k_proj_weight: Tensor2D {
data: &b.k_proj[i],
rows: kv_dim,
cols: hidden_size,
},
v_proj_weight: Tensor2D {
data: &b.v_proj[i],
rows: kv_dim,
cols: hidden_size,
},
o_proj_weight: Tensor2D {
data: &b.o_proj[i],
rows: hidden_size,
cols: q_dim,
},
q_norm_weight: Tensor1D {
data: &b.q_norm[i],
len: head_dim,
},
k_norm_weight: Tensor1D {
data: &b.k_norm[i],
len: head_dim,
},
input_layernorm_weight: Tensor1D {
data: &b.input_ln[i],
len: hidden_size,
},
gate_proj_weight: Tensor2D {
data: &b.gate_proj[i],
rows: intermediate_size,
cols: hidden_size,
},
up_proj_weight: Tensor2D {
data: &b.up_proj[i],
rows: intermediate_size,
cols: hidden_size,
},
down_proj_weight: Tensor2D {
data: &b.down_proj[i],
rows: hidden_size,
cols: intermediate_size,
},
post_attention_layernorm_weight: Tensor1D {
data: &b.post_ln[i],
len: hidden_size,
},
fused_qkv,
qkv_out_dim,
fused_gate_up,
gate_up_out_dim,
});
}
QwenWeights {
embed_tokens,
norm_weight,
layers,
}
};
Ok((backing, weights))
}
}
fn parse_safetensors_header(json: &str) -> Result<HashMap<String, TensorMeta>, InferenceError> {
let mut parser = JsonParser::new(json);
parser.expect(b'{')?;
let mut tensors = HashMap::new();
let mut seen_keys = HashSet::new();
loop {
parser.skip_ws();
match parser.peek() {
Some(b'}') => {
parser.bump();
break;
}
Some(_) => {}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of header while parsing top-level object".into(),
));
}
}
let key = parser.parse_string()?;
if !seen_keys.insert(key.clone()) {
let kind = if key == "__metadata__" {
"metadata member"
} else {
"tensor name"
};
return Err(InferenceError::InvalidSafetensors(format!(
"duplicate {kind} in safetensors header: {key}"
)));
}
parser.skip_ws();
parser.expect(b':')?;
parser.skip_ws();
if key == "__metadata__" {
parser.parse_metadata()?;
} else {
let meta = parser.parse_tensor_meta(&key)?;
tensors.insert(key, meta);
}
parser.skip_ws();
match parser.peek() {
Some(b',') => {
parser.bump();
parser.skip_ws();
if matches!(parser.peek(), Some(b'}')) {
return Err(InferenceError::InvalidSafetensors(
"trailing comma in safetensors header object".into(),
));
}
}
Some(b'}') => {
parser.bump();
break;
}
Some(other) => {
return Err(InferenceError::InvalidSafetensors(format!(
"expected ',' or '}}' in header, found byte {other}"
)));
}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of header after tensor entry".into(),
));
}
}
}
while matches!(parser.peek(), Some(b' ')) {
parser.bump();
}
if let Some(other) = parser.peek() {
return Err(InferenceError::InvalidSafetensors(format!(
"non-space byte {other} after top-level safetensors header object"
)));
}
Ok(tensors)
}
struct JsonParser<'a> {
bytes: &'a [u8],
pos: usize,
depth: usize,
}
impl<'a> JsonParser<'a> {
fn new(s: &'a str) -> Self {
Self {
bytes: s.as_bytes(),
pos: 0,
depth: 0,
}
}
fn peek(&self) -> Option<u8> {
self.bytes.get(self.pos).copied()
}
fn bump(&mut self) -> Option<u8> {
let b = self.peek()?;
self.pos += 1;
Some(b)
}
fn skip_ws(&mut self) {
while matches!(self.peek(), Some(b' ' | b'\n' | b'\r' | b'\t')) {
self.pos += 1;
}
}
fn expect(&mut self, expected: u8) -> Result<(), InferenceError> {
match self.bump() {
Some(actual) if actual == expected => Ok(()),
Some(actual) => Err(InferenceError::InvalidSafetensors(format!(
"expected byte {:?}, found byte {:?} at position {}",
expected as char,
actual as char,
self.pos.saturating_sub(1)
))),
None => Err(InferenceError::InvalidSafetensors(format!(
"expected byte {:?}, found end of input",
expected as char
))),
}
}
fn parse_string(&mut self) -> Result<String, InferenceError> {
let start = self.pos;
self.expect(b'"')?;
let mut escaped = false;
loop {
let byte = self.bump().ok_or_else(|| {
InferenceError::InvalidSafetensors("unterminated string in header".into())
})?;
if escaped {
escaped = false;
} else if byte == b'\\' {
escaped = true;
} else if byte == b'"' {
break;
}
}
let raw = std::str::from_utf8(&self.bytes[start..self.pos]).map_err(|err| {
InferenceError::InvalidSafetensors(format!(
"safetensors header string is not valid UTF-8: {err}"
))
})?;
serde_json::from_str(raw).map_err(|err| {
InferenceError::InvalidSafetensors(format!(
"invalid JSON string in safetensors header: {err}"
))
})
}
fn parse_usize(&mut self) -> Result<usize, InferenceError> {
let start = self.pos;
match self.peek() {
Some(b'0') => {
self.pos += 1;
if matches!(self.peek(), Some(b'0'..=b'9')) {
return Err(InferenceError::InvalidSafetensors(format!(
"leading zero in unsigned integer at byte {start}"
)));
}
}
Some(b'1'..=b'9') => {
self.pos += 1;
while matches!(self.peek(), Some(b'0'..=b'9')) {
self.pos += 1;
}
}
_ => {}
}
if start == self.pos {
return Err(InferenceError::InvalidSafetensors(format!(
"expected unsigned integer at byte {}",
self.pos
)));
}
let s = std::str::from_utf8(&self.bytes[start..self.pos]).map_err(|e| {
InferenceError::InvalidSafetensors(format!("invalid number token in header: {e}"))
})?;
s.parse::<usize>().map_err(|e| {
InferenceError::InvalidSafetensors(format!("invalid usize value {s}: {e}"))
})
}
fn parse_usize_array(&mut self) -> Result<Vec<usize>, InferenceError> {
self.expect(b'[')?;
self.skip_ws();
let mut values = Vec::new();
if matches!(self.peek(), Some(b']')) {
self.bump();
return Ok(values);
}
loop {
self.skip_ws();
values.push(self.parse_usize()?);
self.skip_ws();
match self.peek() {
Some(b',') => {
self.bump();
self.skip_ws();
}
Some(b']') => {
self.bump();
break;
}
Some(other) => {
return Err(InferenceError::InvalidSafetensors(format!(
"expected ',' or ']' in array, found byte {other}"
)));
}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of input in array".into(),
));
}
}
}
Ok(values)
}
fn parse_metadata(&mut self) -> Result<(), InferenceError> {
self.expect(b'{')?;
self.skip_ws();
let mut seen_keys = HashSet::new();
if matches!(self.peek(), Some(b'}')) {
self.bump();
return Ok(());
}
loop {
self.skip_ws();
let key = self.parse_string()?;
if !seen_keys.insert(key.clone()) {
return Err(InferenceError::InvalidSafetensors(format!(
"duplicate __metadata__ key in safetensors header: {key}"
)));
}
self.skip_ws();
self.expect(b':')?;
self.skip_ws();
self.parse_string().map_err(|_| {
InferenceError::InvalidSafetensors(format!(
"safetensors __metadata__ value for {key:?} must be a string"
))
})?;
self.skip_ws();
match self.peek() {
Some(b',') => {
self.bump();
}
Some(b'}') => {
self.bump();
break;
}
Some(other) => {
return Err(InferenceError::InvalidSafetensors(format!(
"expected ',' or '}}' in safetensors __metadata__, found byte {other}"
)));
}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of safetensors __metadata__ object".into(),
));
}
}
}
Ok(())
}
fn skip_value(&mut self) -> Result<(), InferenceError> {
self.skip_ws();
match self.peek() {
Some(b'"') => {
self.parse_string()?;
Ok(())
}
Some(b'{') => {
self.depth += 1;
if self.depth > MAX_SAFETENSORS_HEADER_DEPTH {
return Err(InferenceError::InvalidSafetensors(format!(
"header nesting exceeds depth limit of {MAX_SAFETENSORS_HEADER_DEPTH}"
)));
}
self.bump();
self.skip_ws();
if matches!(self.peek(), Some(b'}')) {
self.bump();
self.depth -= 1;
return Ok(());
}
let mut seen_keys = HashSet::new();
loop {
self.skip_ws();
let key = self.parse_string()?;
if !seen_keys.insert(key.clone()) {
return Err(InferenceError::InvalidSafetensors(format!(
"duplicate key in safetensors header object: {key}"
)));
}
self.skip_ws();
self.expect(b':')?;
self.skip_ws();
self.skip_value()?;
self.skip_ws();
match self.peek() {
Some(b',') => {
self.bump();
}
Some(b'}') => {
self.bump();
break;
}
Some(other) => {
return Err(InferenceError::InvalidSafetensors(format!(
"expected ',' or '}}' while skipping object, found {other}"
)));
}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of input while skipping object".into(),
));
}
}
}
self.depth -= 1;
Ok(())
}
Some(b'[') => {
self.depth += 1;
if self.depth > MAX_SAFETENSORS_HEADER_DEPTH {
return Err(InferenceError::InvalidSafetensors(format!(
"header nesting exceeds depth limit of {MAX_SAFETENSORS_HEADER_DEPTH}"
)));
}
self.bump();
self.skip_ws();
if matches!(self.peek(), Some(b']')) {
self.bump();
self.depth -= 1;
return Ok(());
}
loop {
self.skip_ws();
self.skip_value()?;
self.skip_ws();
match self.peek() {
Some(b',') => {
self.bump();
}
Some(b']') => {
self.bump();
break;
}
Some(other) => {
return Err(InferenceError::InvalidSafetensors(format!(
"expected ',' or ']' while skipping array, found {other}"
)));
}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of input while skipping array".into(),
));
}
}
}
self.depth -= 1;
Ok(())
}
Some(b't') => self.skip_literal(b"true"),
Some(b'f') => self.skip_literal(b"false"),
Some(b'n') => self.skip_literal(b"null"),
Some(b'-' | b'0'..=b'9') => self.skip_number(),
Some(other) => Err(InferenceError::InvalidSafetensors(format!(
"unsupported JSON value starting with byte {other}"
))),
None => Err(InferenceError::InvalidSafetensors(
"unexpected end of input while skipping value".into(),
)),
}
}
fn skip_literal(&mut self, literal: &[u8]) -> Result<(), InferenceError> {
if self.pos + literal.len() > self.bytes.len()
|| &self.bytes[self.pos..self.pos + literal.len()] != literal
{
return Err(InferenceError::InvalidSafetensors(format!(
"expected literal {:?}",
std::str::from_utf8(literal).unwrap_or("<literal>")
)));
}
self.pos += literal.len();
Ok(())
}
fn skip_number(&mut self) -> Result<(), InferenceError> {
let start = self.pos;
if matches!(self.peek(), Some(b'-')) {
self.pos += 1;
}
match self.peek() {
Some(b'0') => {
self.pos += 1;
if matches!(self.peek(), Some(b'0'..=b'9')) {
return Err(InferenceError::InvalidSafetensors(format!(
"leading zero in JSON number at byte {start}"
)));
}
}
Some(b'1'..=b'9') => {
self.pos += 1;
while matches!(self.peek(), Some(b'0'..=b'9')) {
self.pos += 1;
}
}
_ => {
return Err(InferenceError::InvalidSafetensors(format!(
"invalid JSON number at byte {start}: expected an integer digit"
)));
}
}
if matches!(self.peek(), Some(b'.')) {
self.pos += 1;
let fraction_start = self.pos;
while matches!(self.peek(), Some(b'0'..=b'9')) {
self.pos += 1;
}
if self.pos == fraction_start {
return Err(InferenceError::InvalidSafetensors(format!(
"invalid JSON number at byte {start}: fraction has no digits"
)));
}
}
if matches!(self.peek(), Some(b'e' | b'E')) {
self.pos += 1;
if matches!(self.peek(), Some(b'+' | b'-')) {
self.pos += 1;
}
let exponent_start = self.pos;
while matches!(self.peek(), Some(b'0'..=b'9')) {
self.pos += 1;
}
if self.pos == exponent_start {
return Err(InferenceError::InvalidSafetensors(format!(
"invalid JSON number at byte {start}: exponent has no digits"
)));
}
}
Ok(())
}
fn parse_tensor_meta(&mut self, tensor_name: &str) -> Result<TensorMeta, InferenceError> {
self.expect(b'{')?;
self.skip_ws();
let mut dtype: Option<DType> = None;
let mut dtype_str: Option<String> = None;
let mut shape = None;
let mut data_offsets = None;
let mut seen_keys = HashSet::new();
if matches!(self.peek(), Some(b'}')) {
self.bump();
} else {
loop {
self.skip_ws();
let key = self.parse_string()?;
if !seen_keys.insert(key.clone()) {
return Err(InferenceError::InvalidSafetensors(format!(
"duplicate member in tensor {tensor_name} header object: {key}"
)));
}
self.skip_ws();
self.expect(b':')?;
self.skip_ws();
match key.as_str() {
"dtype" => {
let s = self.parse_string()?;
dtype = dtype_from_str(&s);
dtype_str = Some(s);
}
"shape" => shape = Some(self.parse_usize_array()?),
"data_offsets" => {
let arr = self.parse_usize_array()?;
if arr.len() != 2 {
return Err(InferenceError::InvalidSafetensors(format!(
"data_offsets must have length 2, got {}",
arr.len()
)));
}
data_offsets = Some((arr[0], arr[1]));
}
_ => self.skip_value()?,
}
self.skip_ws();
match self.peek() {
Some(b',') => {
self.bump();
}
Some(b'}') => {
self.bump();
break;
}
Some(other) => {
return Err(InferenceError::InvalidSafetensors(format!(
"expected ',' or '}}' in tensor object, found {other}"
)));
}
None => {
return Err(InferenceError::InvalidSafetensors(
"unexpected end of input in tensor object".into(),
));
}
}
}
}
let dtype = match (dtype, dtype_str) {
(Some(dtype), _) => dtype,
(None, Some(s)) => {
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {tensor_name} has unrecognized dtype {s:?}; its byte width is \
unknowable so extent validation cannot proceed"
)));
}
(None, None) => {
return Err(InferenceError::InvalidSafetensors(format!(
"tensor {tensor_name} entry missing dtype"
)));
}
};
let shape = shape.ok_or_else(|| {
InferenceError::InvalidSafetensors(format!("tensor {tensor_name} entry missing shape"))
})?;
let (start, end) = data_offsets.ok_or_else(|| {
InferenceError::InvalidSafetensors(format!(
"tensor {tensor_name} entry missing data_offsets"
))
})?;
Ok(TensorMeta {
dtype,
shape,
start,
end,
converted_f32: OnceLock::new(),
validated: OnceLock::new(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use std::io::Write;
use std::time::{SystemTime, UNIX_EPOCH};
struct TempFileGuard(std::path::PathBuf);
impl TempFileGuard {
fn into_path(self) -> std::path::PathBuf {
let path = self.0.clone();
std::mem::forget(self);
path
}
}
impl std::ops::Deref for TempFileGuard {
type Target = std::path::Path;
fn deref(&self) -> &std::path::Path {
&self.0
}
}
impl AsRef<std::path::Path> for TempFileGuard {
fn as_ref(&self) -> &std::path::Path {
&self.0
}
}
impl Drop for TempFileGuard {
fn drop(&mut self) {
let _ = fs::remove_file(&self.0);
}
}
struct TempDirGuard(std::path::PathBuf);
impl TempDirGuard {
fn into_path(self) -> std::path::PathBuf {
let path = self.0.clone();
std::mem::forget(self);
path
}
}
impl std::ops::Deref for TempDirGuard {
type Target = std::path::Path;
fn deref(&self) -> &std::path::Path {
&self.0
}
}
impl AsRef<std::path::Path> for TempDirGuard {
fn as_ref(&self) -> &std::path::Path {
&self.0
}
}
impl Drop for TempDirGuard {
fn drop(&mut self) {
let _ = fs::remove_dir_all(&self.0);
}
}
fn temp_path(name: &str) -> TempFileGuard {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("invariant: system time is after UNIX_EPOCH")
.as_nanos();
TempFileGuard(std::env::temp_dir().join(format!(
"{}_{}_{}.safetensors",
name,
std::process::id(),
nanos
)))
}
#[test]
fn temp_dir_guard_removes_directory_on_normal_drop() {
let dir = temp_dir("lattice_guard_drop_test");
let path = dir.to_path_buf();
assert!(path.is_dir(), "temp_dir must create the directory eagerly");
drop(dir);
assert!(
!path.exists(),
"TempDirGuard must remove its directory when dropped normally"
);
}
#[test]
fn temp_dir_guard_removes_directory_on_unwind() {
let dir = temp_dir("lattice_guard_unwind_test");
let path = dir.to_path_buf();
assert!(path.is_dir(), "temp_dir must create the directory eagerly");
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _dir = dir;
panic!("intentional panic to exercise TempDirGuard unwind cleanup");
}));
assert!(result.is_err(), "the inner closure must have panicked");
assert!(
!path.exists(),
"TempDirGuard must remove its directory when its owning frame unwinds"
);
}
#[test]
fn temp_dir_guard_into_path_opts_out_of_cleanup() {
let dir = temp_dir("lattice_guard_leak_opt_out_test");
let path = dir.into_path();
assert!(
path.is_dir(),
"into_path must hand back the still-existing directory without removing it"
);
fs::remove_dir_all(&path).ok();
}
#[test]
fn temp_file_guard_removes_file_on_normal_drop() {
let guard = temp_path("lattice_guard_file_drop_test");
write_raw_safetensors(&guard, r#"{"__metadata__":{"format":"pt"}}"#, &[]);
let path = guard.to_path_buf();
assert!(path.is_file(), "test setup must have created the file");
drop(guard);
assert!(
!path.exists(),
"TempFileGuard must remove its file when dropped normally"
);
}
#[test]
fn temp_file_guard_into_path_opts_out_of_cleanup() {
let guard = temp_path("lattice_guard_file_leak_opt_out_test");
write_raw_safetensors(&guard, r#"{"__metadata__":{"format":"pt"}}"#, &[]);
let path = guard.into_path();
assert!(
path.is_file(),
"into_path must hand back the still-existing file without removing it"
);
fs::remove_file(&path).ok();
}
#[test]
fn test_parse_small_safetensors_file() {
let path = temp_path("lattice_weights_test");
let header = r#"{
"__metadata__": {"format": "pt"},
"vec": {"dtype": "F32", "shape": [2], "data_offsets": [0, 8]},
"mat": {"dtype": "F32", "shape": [2, 2], "data_offsets": [8, 24]}
}"#
.replace(['\n', ' '], "");
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for value in [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0] {
bytes.extend_from_slice(&value.to_le_bytes());
}
let mut file = File::create(&path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
drop(file);
let st = SafetensorsFile::open(&path).expect("test setup: parse safetensors file");
let (vec_data, vec_shape) = st
.get_f32_tensor("vec")
.expect("test setup: vec tensor exists");
let (mat_data, mat_shape) = st
.get_f32_tensor("mat")
.expect("test setup: mat tensor exists");
assert_eq!(vec_shape, &[2]);
assert_eq!(mat_shape, &[2, 2]);
assert_eq!(vec_data, &[1.0, 2.0]);
assert_eq!(mat_data, &[3.0, 4.0, 5.0, 6.0]);
}
#[test]
fn duplicate_tensor_name_in_safetensors_header_is_rejected() {
let header = r#"{"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]},"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let err = SafetensorsFile::from_bytes(bytes)
.expect_err("duplicate tensor names must not collapse into the header map");
assert!(
err.to_string().contains("duplicate tensor name"),
"error must identify the duplicate header member: {err}"
);
}
#[test]
fn duplicate_tensor_object_members_are_rejected() {
let duplicate_members = [
r#""dtype":"F32","dtype":"F32","shape":[1],"data_offsets":[0,4]"#,
r#""dtype":"F32","shape":[1],"shape":[1],"data_offsets":[0,4]"#,
r#""dtype":"F32","shape":[1],"data_offsets":[0,4],"data_offsets":[0,4]"#,
r#""dtype":"F32","shape":[1],"data_offsets":[0,4],"extra":1,"extra":2"#,
r#""dtype":"F32","shape":[1],"data_offsets":[0,4],"extra":{"nested":1,"nested":2}"#,
];
for members in duplicate_members {
let header = format!(r#"{{"tensor":{{{members}}}}}"#);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let err = SafetensorsFile::from_bytes(bytes)
.expect_err("duplicate keys anywhere in a tensor object must be rejected");
assert!(
err.to_string().contains("duplicate"),
"error must identify the duplicate object member: {err}"
);
}
}
#[test]
fn metadata_must_be_a_unique_string_map() {
let invalid_headers = [
r#"{"__metadata__":{"source":"a","source":"b"},"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#,
r#"{"__metadata__":{"source":1},"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#,
r#"{"__metadata__":[],"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#,
r#"{"__metadata__":{},"__metadata__":{},"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#,
];
for header in invalid_headers {
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect_err("safetensors metadata must be one unique string-to-string map");
}
}
#[test]
fn utf8_strings_decode_canonically_before_duplicate_checks() {
let header = r#"{"__metadata__":{"作者":"café 🚀"},"t\u00e9nsor\ud83d\ude80":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let parsed = SafetensorsFile::from_bytes(bytes)
.expect("valid UTF-8 and a valid escaped surrogate pair must parse");
assert!(
parsed.has_tensor("ténsor🚀"),
"raw UTF-8 and escaped surrogate pairs must decode to canonical scalar values"
);
let duplicate = r#"{"__metadata__":{"é":"raw","\u00e9":"escaped"},"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(duplicate.len() as u64).to_le_bytes());
bytes.extend_from_slice(duplicate.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect_err("raw and escaped forms of the same metadata key are duplicates");
}
#[test]
fn malformed_json_numbers_are_rejected() {
let invalid_members = [
r#""extra":-"#,
r#""extra":1+2"#,
r#""extra":01"#,
r#""extra":-01"#,
r#""extra":1."#,
r#""extra":1e"#,
r#""extra":1e+"#,
r#""shape":[01]"#,
r#""data_offsets":[00,4]"#,
];
for invalid in invalid_members {
let members = match invalid.split_once(':').map(|(key, _)| key) {
Some(r#""shape""#) => {
format!(r#""dtype":"F32",{invalid},"data_offsets":[0,4]"#)
}
Some(r#""data_offsets""#) => {
format!(r#""dtype":"F32","shape":[1],{invalid}"#)
}
_ => format!(r#""dtype":"F32","shape":[1],"data_offsets":[0,4],{invalid}"#),
};
let header = format!(r#"{{"tensor":{{{members}}}}}"#);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
assert!(
SafetensorsFile::from_bytes(bytes).is_err(),
"invalid JSON number must fail closed: {header}"
);
}
}
#[test]
fn valid_json_numbers_in_unknown_values_are_accepted() {
let header = r#"{"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4],"extra":[0,-1,1.5,1e2,-0.25E+2]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect("RFC 8259 numbers in an unknown extension value must remain valid");
}
#[test]
fn malformed_unicode_surrogates_are_rejected() {
for escaped in [r#"\ud83d"#, r#"\ude80"#, r#"\ud83dX"#] {
let header = format!(
r#"{{"__metadata__":{{"value":"{escaped}"}},"tensor":{{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}}}"#
);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect_err("a malformed JSON Unicode surrogate must fail closed");
}
}
#[test]
fn trailing_comma_in_header_object_is_rejected() {
let header = r#"{"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]},}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect_err("the safetensors header must be valid JSON without a trailing comma");
}
#[test]
fn non_space_bytes_after_top_level_header_are_rejected() {
let tensor = r#"{"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
let header = format!("{tensor}THIS_IS_NOT_JSON");
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect_err("the declared header may contain only space padding after its JSON object");
}
#[test]
fn space_padding_after_top_level_header_is_accepted() {
let tensor = r#"{"tensor":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
let header = format!("{tensor} ");
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
SafetensorsFile::from_bytes(bytes)
.expect("ASCII space is valid safetensors header padding");
}
#[cfg(unix)]
#[test]
fn open_rejects_a_group_or_other_writable_checkpoint_file() {
use std::os::unix::fs::PermissionsExt;
let path = temp_path("lattice_weights_writable_checkpoint");
let header = r#"{"vec":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let mut file = File::create(&path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
drop(file);
fs::set_permissions(&path, fs::Permissions::from_mode(0o666)).expect("chmod 0o666");
let err = SafetensorsFile::open(&path)
.expect_err("a group/other-writable checkpoint file must be refused");
assert!(
matches!(&err, InferenceError::InvalidSafetensors(msg) if msg.contains("refusing to load")),
"expected a trust-boundary refusal, got: {err:?}"
);
fs::set_permissions(&path, fs::Permissions::from_mode(0o600)).expect("chmod 0o600");
SafetensorsFile::open(&path).expect("an owner-only checkpoint file must still be accepted");
}
#[test]
fn test_rejects_shape_byte_length_mismatch() {
let path = temp_path("lattice_weights_bad_shape");
let header = r#"{
"bad": {"dtype": "F32", "shape": [3], "data_offsets": [0, 8]}
}"#
.replace(['\n', ' '], "");
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for value in [1.0f32, 2.0] {
bytes.extend_from_slice(&value.to_le_bytes());
}
let mut file = File::create(&path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
drop(file);
let err = SafetensorsFile::open(&path)
.expect_err("shape byte length mismatch should be rejected at open");
assert!(
err.to_string().contains("byte length mismatch"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_one_byte_extra_payload() {
let path = temp_path("lattice_weights_extra_byte");
let header = r#"{
"extra": {"dtype": "F32", "shape": [2], "data_offsets": [0, 9]}
}"#
.replace(['\n', ' '], "");
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for value in [1.0f32, 2.0] {
bytes.extend_from_slice(&value.to_le_bytes());
}
bytes.push(0);
let mut file = File::create(&path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
drop(file);
let err = SafetensorsFile::open(&path)
.expect_err("one-byte-extra payload should be rejected at open");
assert!(
err.to_string().contains("byte length mismatch"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_aliased_tensor_ranges() {
let path = temp_path("lattice_weights_aliased_ranges");
let header = r#"{
"key": {"dtype": "F32", "shape": [1], "data_offsets": [0, 4]},
"value": {"dtype": "F32", "shape": [1], "data_offsets": [0, 4]}
}"#
.replace(['\n', ' '], "");
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let mut file = File::create(&path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
drop(file);
let err = SafetensorsFile::open(&path)
.expect_err("aliased tensor byte ranges must be rejected at open");
assert!(
err.to_string().contains("non-contiguous"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_leading_hole_in_data_section() {
let path = temp_path("lattice_weights_leading_hole");
let header = r#"{"t":{"dtype":"F32","shape":[1],"data_offsets":[4,8]}}"#;
write_raw_safetensors(&path, header, &[0; 8]);
let err = SafetensorsFile::open(&path)
.expect_err("the first tensor must start at data-section offset zero");
assert!(
err.to_string().contains("non-contiguous"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_internal_hole_in_data_section() {
let path = temp_path("lattice_weights_internal_hole");
let header = r#"{
"a":{"dtype":"F32","shape":[1],"data_offsets":[0,4]},
"b":{"dtype":"F32","shape":[1],"data_offsets":[8,12]}
}"#
.replace(['\n', ' '], "");
write_raw_safetensors(&path, &header, &[0; 12]);
let err = SafetensorsFile::open(&path)
.expect_err("tensor ranges must cover the data section contiguously");
assert!(
err.to_string().contains("non-contiguous"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_trailing_payload_outside_tensor_ranges() {
let path = temp_path("lattice_weights_trailing_payload");
let header = r#"{"t":{"dtype":"F32","shape":[1],"data_offsets":[0,4]}}"#;
write_raw_safetensors(&path, header, &[0; 8]);
let err = SafetensorsFile::open(&path)
.expect_err("tensor ranges must exhaust the complete data section");
assert!(
err.to_string().contains("trailing or missing payload"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_reversed_tensor_range() {
let path = temp_path("lattice_weights_reversed_range");
let header = r#"{"t":{"dtype":"F32","shape":[1],"data_offsets":[4,0]}}"#;
write_raw_safetensors(&path, header, &[0; 4]);
let err =
SafetensorsFile::open(&path).expect_err("a tensor range must not end before it starts");
assert!(
err.to_string().contains("invalid data_offsets"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_header_exceeding_depth_limit() {
let depth = MAX_SAFETENSORS_HEADER_DEPTH + 8;
let mut header =
String::from(r#"{"tensor":{"dtype":"F32","shape":[0],"data_offsets":[0,0],"extra":"#);
header.push_str(&"[".repeat(depth));
header.push_str(&"]".repeat(depth));
header.push_str("}}");
assert!(header.len() < MAX_SAFETENSORS_HEADER_BYTES);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
let err = SafetensorsFile::from_bytes(bytes)
.expect_err("header nesting past the depth limit must be rejected");
assert!(
err.to_string().contains("depth limit"),
"unexpected error: {err}"
);
}
#[test]
fn test_rejects_header_exceeding_size_limit() {
let prefix = r#"{"__metadata__":""#;
let suffix = r#""}"#;
let pad_len = MAX_SAFETENSORS_HEADER_BYTES + 1024 - prefix.len() - suffix.len();
let header_len = prefix.len() + pad_len + suffix.len();
assert!(header_len > MAX_SAFETENSORS_HEADER_BYTES);
let mut bytes = Vec::with_capacity(8 + header_len);
bytes.extend_from_slice(&(header_len as u64).to_le_bytes());
bytes.extend_from_slice(prefix.as_bytes());
bytes.extend_from_slice(&vec![b'a'; pad_len]);
bytes.extend_from_slice(suffix.as_bytes());
let err = SafetensorsFile::from_bytes(bytes)
.expect_err("a header past the size limit must be rejected");
assert!(
err.to_string().contains("exceeds limit"),
"unexpected error: {err}"
);
}
#[test]
fn test_parses_header_within_depth_and_size_bounds() {
let header = r#"{
"__metadata__": {"format": "pt"},
"t": {"dtype": "F32", "shape": [2], "data_offsets": [0, 8]}
}"#
.replace(['\n', ' '], "");
assert!(header.len() < MAX_SAFETENSORS_HEADER_BYTES);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for value in [1.0f32, 2.0] {
bytes.extend_from_slice(&value.to_le_bytes());
}
let sf = SafetensorsFile::from_bytes(bytes)
.expect("a header within both bounds must still parse");
let (data, shape) = sf
.get_f32_tensor("t")
.expect("tensor within a bounded header must still load");
assert_eq!(shape, &[2]);
assert_eq!(data, &[1.0, 2.0]);
}
fn write_raw_safetensors(path: &std::path::Path, header: &str, raw: &[u8]) {
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(raw);
let mut file = File::create(path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
}
fn write_raw_tensor(
path: &std::path::Path,
name: &str,
dtype: &str,
shape: &[usize],
raw: &[u8],
) {
let shape_str = format!(
"[{}]",
shape
.iter()
.map(std::string::ToString::to_string)
.collect::<Vec<_>>()
.join(",")
);
let header = format!(
r#"{{"{name}":{{"dtype":"{dtype}","shape":{shape_str},"data_offsets":[0,{}]}}}}"#,
raw.len()
);
write_raw_safetensors(path, &header, raw);
}
#[test]
fn test_open_mapped_f32_rejects_nan() {
let path = temp_path("lattice_weights_f32_nan_mapped");
let mut raw = Vec::new();
raw.extend_from_slice(&1.0f32.to_le_bytes());
raw.extend_from_slice(&f32::NAN.to_le_bytes());
write_raw_tensor(&path, "t", "F32", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("NaN in mapped F32 data must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[test]
fn test_open_mapped_f32_rejects_positive_infinity() {
let path = temp_path("lattice_weights_f32_posinf_mapped");
let mut raw = Vec::new();
raw.extend_from_slice(&f32::INFINITY.to_le_bytes());
raw.extend_from_slice(&2.0f32.to_le_bytes());
write_raw_tensor(&path, "t", "F32", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("+inf in mapped F32 data must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[test]
fn test_open_mapped_f32_rejects_negative_infinity() {
let path = temp_path("lattice_weights_f32_neginf_mapped");
let mut raw = Vec::new();
raw.extend_from_slice(&1.0f32.to_le_bytes());
raw.extend_from_slice(&f32::NEG_INFINITY.to_le_bytes());
write_raw_tensor(&path, "t", "F32", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("-inf in mapped F32 data must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[test]
fn test_from_bytes_owned_f32_rejects_non_finite() {
let header = r#"{"t":{"dtype":"F32","shape":[2],"data_offsets":[0,8]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&f32::NAN.to_le_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let sf = SafetensorsFile::from_bytes(bytes).expect("from_bytes: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("NaN via from_bytes (owned backing) must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[test]
fn test_get_f32_tensor_caches_validation_after_first_call() {
let path = temp_path("lattice_weights_validation_cache");
write_single_f32_tensor(&path, "t", &[1.0, 2.0, 3.0]);
let sf = SafetensorsFile::open(&path).expect("open: valid file");
let meta = sf.tensors.get("t").expect("tensor tracked");
assert!(
meta.validated.get().is_none(),
"must not be marked validated before first access"
);
sf.get_f32_tensor("t").expect("first access validates");
assert!(
meta.validated.get().is_some(),
"must be marked validated after first access"
);
sf.get_f32_tensor("t")
.expect("second access reuses cached validation");
}
#[test]
fn test_unsupported_dtype_tensor_fails_as_format_error_not_missing_tensor() {
let path = temp_path("lattice_weights_unsupported_dtype");
write_raw_tensor(&path, "counts", "I64", &[2], &8i64.to_le_bytes().repeat(2));
let sf = SafetensorsFile::open(&path).expect("open: known dtype, valid extent");
assert!(
sf.has_tensor("counts"),
"unsupported-dtype tensor must remain structurally tracked"
);
assert_eq!(sf.tensor_shape("counts"), Some(&[2usize][..]));
let err = sf
.get_f32_tensor("counts")
.expect_err("requesting an I64 tensor as f32 must fail");
assert!(
matches!(err, InferenceError::InvalidSafetensors(_)),
"must fail as a dtype/format error, not MissingTensor: {err:?}"
);
assert!(
!matches!(err, InferenceError::MissingTensor(_)),
"must not be reported as a missing tensor: {err:?}"
);
}
#[test]
fn test_standard_sub_byte_and_fnuz_dtypes_are_tracked_structurally() {
let cases = [
("F4", 2usize, 1usize),
("F6_E2M3", 4, 3),
("F6_E3M2", 4, 3),
("F8_E4M3FNUZ", 1, 1),
("F8_E5M2FNUZ", 1, 1),
];
for (dtype, elements, bytes) in cases {
let path = temp_path(&format!("lattice_weights_{dtype}"));
write_raw_tensor(&path, "t", dtype, &[elements], &vec![0; bytes]);
let sf = SafetensorsFile::open(&path)
.unwrap_or_else(|err| panic!("{dtype} with valid extent must open: {err}"));
assert!(sf.has_tensor("t"), "{dtype} tensor must remain tracked");
assert_eq!(sf.tensor_dtype("t"), Some(dtype));
let err = sf
.get_f32_tensor("t")
.expect_err("unsupported dtype must not materialize as f32");
assert!(
err.to_string().contains("unsupported dtype"),
"unexpected {dtype} materialization error: {err}"
);
}
}
#[test]
fn test_sub_byte_dtype_requires_byte_aligned_shape() {
let path = temp_path("lattice_weights_f4_unaligned");
write_raw_tensor(&path, "t", "F4", &[1], &[]);
let err = SafetensorsFile::open(&path)
.expect_err("a one-element F4 tensor must not round down to zero bytes");
assert!(
err.to_string().contains("not byte-aligned"),
"unexpected error: {err}"
);
}
#[test]
fn test_sub_byte_dtype_rejects_wrong_byte_extent() {
let path = temp_path("lattice_weights_f4_bad_extent");
write_raw_tensor(&path, "t", "F4", &[4], &[0]);
let err =
SafetensorsFile::open(&path).expect_err("four F4 elements require two payload bytes");
assert!(
err.to_string().contains("byte length mismatch"),
"unexpected error: {err}"
);
}
#[test]
fn test_unrecognized_dtype_string_rejected_at_parse_time() {
let path = temp_path("lattice_weights_bogus_dtype");
let header = r#"{"t":{"dtype":"BANANA","shape":[1],"data_offsets":[0,4]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
bytes.extend_from_slice(&1.0f32.to_le_bytes());
let mut file = File::create(&path).expect("test setup: create safetensors file");
file.write_all(&bytes)
.expect("test setup: write safetensors bytes");
drop(file);
let err = SafetensorsFile::open(&path)
.expect_err("an unrecognized dtype string must fail at open/parse time");
assert!(
err.to_string().contains("unrecognized dtype"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_f16_tensor_rejects_nan_bit_pattern_through_safetensors_path() {
let path = temp_path("lattice_weights_f16_nan");
let nan_bits: u16 = 0x7E00;
let one_bits: u16 = 0x3C00; let mut raw = Vec::new();
raw.extend_from_slice(&one_bits.to_le_bytes());
raw.extend_from_slice(&nan_bits.to_le_bytes());
write_raw_tensor(&path, "t", "F16", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("F16 NaN bit pattern must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
assert!(
err.to_string().contains("element index 1"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_f16_tensor_rejects_infinity_bit_pattern_through_safetensors_path() {
let path = temp_path("lattice_weights_f16_inf");
let inf_bits: u16 = 0x7C00;
let one_bits: u16 = 0x3C00;
let mut raw = Vec::new();
raw.extend_from_slice(&inf_bits.to_le_bytes());
raw.extend_from_slice(&one_bits.to_le_bytes());
write_raw_tensor(&path, "t", "F16", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("F16 +inf bit pattern must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_bf16_tensor_rejects_nan_bit_pattern_through_safetensors_path() {
let path = temp_path("lattice_weights_bf16_nan");
let nan_bits: u16 = 0x7FC0;
let one_bits: u16 = 0x3F80; let mut raw = Vec::new();
raw.extend_from_slice(&one_bits.to_le_bytes());
raw.extend_from_slice(&nan_bits.to_le_bytes());
write_raw_tensor(&path, "t", "BF16", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("BF16 NaN bit pattern must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
assert!(
err.to_string().contains("element index 1"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_bf16_tensor_rejects_infinity_bit_pattern_through_safetensors_path() {
let path = temp_path("lattice_weights_bf16_inf");
let inf_bits: u16 = 0x7F80;
let one_bits: u16 = 0x3F80;
let mut raw = Vec::new();
raw.extend_from_slice(&inf_bits.to_le_bytes());
raw.extend_from_slice(&one_bits.to_le_bytes());
write_raw_tensor(&path, "t", "BF16", &[2], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("BF16 +inf bit pattern must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_f8_e4m3_tensor_decodes_through_safetensors_path() {
let path = temp_path("lattice_weights_f8_e4m3_decode");
let raw: [u8; 4] = [0x00, 0x38, 0xb8, 0x7e];
write_raw_tensor(&path, "t", "F8_E4M3", &[raw.len()], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let (values, shape) = sf
.get_f32_tensor("t")
.expect("finite F8_E4M3 values must decode");
assert_eq!(shape, &[raw.len()]);
assert_eq!(values, &[0.0f32, 1.0, -1.0, 448.0]);
}
#[cfg(feature = "f16")]
#[test]
fn test_f8_e4m3_tensor_rejects_nan_bit_pattern_through_safetensors_path() {
let path = temp_path("lattice_weights_f8_e4m3_nan");
let one_bits: u8 = 0x38;
let nan_bits: u8 = 0x7f;
write_raw_tensor(&path, "t", "F8_E4M3", &[2], &[one_bits, nan_bits]);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("F8_E4M3 NaN bit pattern must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
assert!(
err.to_string().contains("element index 1"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_f8_e5m2_tensor_decodes_through_safetensors_path() {
let path = temp_path("lattice_weights_f8_e5m2_decode");
let raw: [u8; 4] = [0x00, 0x3c, 0xbc, 0x7b];
write_raw_tensor(&path, "t", "F8_E5M2", &[raw.len()], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let (values, shape) = sf
.get_f32_tensor("t")
.expect("finite F8_E5M2 values must decode");
assert_eq!(shape, &[raw.len()]);
assert_eq!(values, &[0.0f32, 1.0, -1.0, 57344.0]);
}
#[cfg(feature = "f16")]
#[test]
fn test_f8_e5m2_tensor_rejects_infinity_bit_pattern_through_safetensors_path() {
let path = temp_path("lattice_weights_f8_e5m2_inf");
let inf_bits: u8 = 0x7c;
let one_bits: u8 = 0x3c;
write_raw_tensor(&path, "t", "F8_E5M2", &[2], &[inf_bits, one_bits]);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let err = sf
.get_f32_tensor("t")
.expect_err("F8_E5M2 +inf bit pattern must be rejected");
assert!(
err.to_string().contains("non-finite"),
"unexpected error: {err}"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_f16_widening_preserves_finite_edge_values() {
let path = temp_path("lattice_weights_f16_fused_finite");
let bits = [0x0000u16, 0x8000, 0x0001, 0x3c00];
let raw = bits
.iter()
.flat_map(|bits| bits.to_le_bytes())
.collect::<Vec<_>>();
write_raw_tensor(&path, "t", "F16", &[bits.len()], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let (values, shape) = sf
.get_f32_tensor("t")
.expect("finite F16 values must widen");
let expected = bits.map(crate::weights::half_bits::f16_bits_to_f32);
assert_eq!(shape, &[bits.len()]);
assert_eq!(
values
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
expected
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>()
);
assert!(
sf.tensors
.get("t")
.expect("tensor tracked")
.validated
.get()
.is_some(),
"widening must publish its fused validation result"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_f16_widening_skips_second_finite_scan() {
let path = temp_path("lattice_weights_f16_fused_scan");
let bits = [0x0000u16, 0x8000, 0x0001, 0x3c00];
let raw = bits
.iter()
.flat_map(|bits| bits.to_le_bytes())
.collect::<Vec<_>>();
write_raw_tensor(&path, "t", "F16", &[bits.len()], &raw);
crate::weights::ingress::reset_decoded_f32_finite_check_count();
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
sf.get_f32_tensor("t")
.expect("finite F16 values must widen");
assert_eq!(
crate::weights::ingress::decoded_f32_finite_check_count(),
0,
"fused F16 validation must not rescan widened values"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_bf16_widening_preserves_finite_edge_values() {
let path = temp_path("lattice_weights_bf16_fused_finite");
let bits = [0x0000u16, 0x8000, 0x0001, 0x3f80];
let raw = bits
.iter()
.flat_map(|bits| bits.to_le_bytes())
.collect::<Vec<_>>();
write_raw_tensor(&path, "t", "BF16", &[bits.len()], &raw);
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
let (values, shape) = sf
.get_f32_tensor("t")
.expect("finite BF16 values must widen");
let expected = bits.map(crate::weights::half_bits::bf16_bits_to_f32);
assert_eq!(shape, &[bits.len()]);
assert_eq!(
values
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
expected
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>()
);
assert!(
sf.tensors
.get("t")
.expect("tensor tracked")
.validated
.get()
.is_some(),
"widening must publish its fused validation result"
);
}
#[cfg(feature = "f16")]
#[test]
fn test_bf16_widening_skips_second_finite_scan() {
let path = temp_path("lattice_weights_bf16_fused_scan");
let bits = [0x0000u16, 0x8000, 0x0001, 0x3f80];
let raw = bits
.iter()
.flat_map(|bits| bits.to_le_bytes())
.collect::<Vec<_>>();
write_raw_tensor(&path, "t", "BF16", &[bits.len()], &raw);
crate::weights::ingress::reset_decoded_f32_finite_check_count();
let sf = SafetensorsFile::open(&path).expect("open: header/extent are valid");
sf.get_f32_tensor("t")
.expect("finite BF16 values must widen");
assert_eq!(
crate::weights::ingress::decoded_f32_finite_check_count(),
0,
"fused BF16 validation must not rescan widened values"
);
}
fn temp_dir(name: &str) -> TempDirGuard {
let nanos = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("invariant: system time is after UNIX_EPOCH")
.as_nanos();
let path = std::env::temp_dir().join(format!("{}_{}_{}", name, std::process::id(), nanos));
fs::create_dir_all(&path).expect("test setup: create temp dir");
TempDirGuard(path)
}
fn write_single_f32_tensor(path: &std::path::Path, name: &str, values: &[f32]) {
let byte_len = std::mem::size_of_val(values);
let header = format!(
r#"{{"{name}":{{"dtype":"F32","shape":[{}],"data_offsets":[0,{byte_len}]}}}}"#,
values.len()
);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for value in values {
bytes.extend_from_slice(&value.to_le_bytes());
}
let mut file = File::create(path).expect("test setup: create shard");
file.write_all(&bytes).expect("test setup: write shard");
}
#[test]
fn test_parse_safetensors_index_and_resolve_weight_map() {
let dir = temp_dir("lattice_sharded_weights_test");
let shard_a = dir.join("model-00001-of-00002.safetensors");
let shard_b = dir.join("model-00002-of-00002.safetensors");
write_single_f32_tensor(&shard_a, "tensor.a", &[1.0, 2.0]);
write_single_f32_tensor(&shard_b, "tensor.b", &[3.0, 4.0, 5.0]);
let index_path = dir.join("model.safetensors.index.json");
fs::write(
&index_path,
r#"{
"metadata": {"total_size": 20.0},
"weight_map": {
"tensor.a": "model-00001-of-00002.safetensors",
"tensor.b": "model-00002-of-00002.safetensors"
}
}"#,
)
.expect("test setup: write index");
let mut st = ShardedSafetensors::open_index(&index_path).expect("index parses");
let (path_a, name_a) = st.resolve_weight("tensor.a").expect("tensor a resolves");
assert_eq!(
path_a.file_name().unwrap().to_string_lossy(),
"model-00001-of-00002.safetensors"
);
assert_eq!(name_a, "tensor.a");
let (a, a_shape) = st.get_f32_tensor_owned("tensor.a").expect("tensor a loads");
let (b, b_shape) = st.get_f32_tensor_owned("tensor.b").expect("tensor b loads");
assert_eq!(a_shape, vec![2]);
assert_eq!(b_shape, vec![3]);
assert_eq!(a, vec![1.0, 2.0]);
assert_eq!(b, vec![3.0, 4.0, 5.0]);
assert!(st.has_tensor("tensor.a").unwrap());
assert!(!st.has_tensor("tensor.c").unwrap());
let index = parse_index(&dir).expect("parse_index succeeds");
assert_eq!(index.weight_map.len(), 2);
let shard_a_path = resolve_shard(&index, &dir, "tensor.a").expect("resolve_shard tensor.a");
assert_eq!(shard_a_path, shard_a);
assert!(resolve_shard(&index, &dir, "tensor.missing").is_err());
let loaded = load_sharded(&dir).expect("load_sharded succeeds");
assert_eq!(loaded.len(), 2);
let ta = loaded.get("tensor.a").expect("tensor.a in loaded map");
let tb = loaded.get("tensor.b").expect("tensor.b in loaded map");
assert_eq!(ta.data, vec![1.0_f32, 2.0]);
assert_eq!(ta.shape, vec![2]);
assert_eq!(tb.data, vec![3.0_f32, 4.0, 5.0]);
assert_eq!(tb.shape, vec![3]);
}
#[test]
fn contained_shard_path_accepts_entries_beneath_model_dir() {
let dir = temp_dir("lattice_containment_ok_test");
fs::write(dir.join("shard.safetensors"), b"x").expect("test setup");
fs::create_dir_all(dir.join("sub")).expect("test setup");
fs::write(dir.join("sub").join("nested.safetensors"), b"x").expect("test setup");
let plain = contained_shard_path(&dir, "shard.safetensors").expect("plain name resolves");
assert_eq!(
plain.file_name().unwrap().to_string_lossy(),
"shard.safetensors"
);
contained_shard_path(&dir, "sub/nested.safetensors")
.expect("subdirectory entry beneath the model dir resolves");
}
#[test]
fn contained_shard_path_rejects_absolute_entry() {
let dir = temp_dir("lattice_containment_abs_test");
let outside = temp_dir("lattice_containment_abs_outside");
let target = outside.join("outside.safetensors");
fs::write(&target, b"x").expect("test setup");
let err = contained_shard_path(&dir, &target.to_string_lossy())
.expect_err("absolute entry must be rejected");
assert!(
matches!(err, InferenceError::InvalidSafetensors(_)),
"unexpected error kind: {err}"
);
let msg = err.to_string();
assert!(
msg.contains("must stay within the model directory")
|| msg.contains("escapes model root"),
"unexpected error: {err}"
);
}
#[test]
fn contained_shard_path_rejects_parent_traversal() {
let outer = temp_dir("lattice_containment_dotdot_test");
let dir = outer.join("model");
fs::create_dir_all(&dir).expect("test setup");
fs::write(outer.join("secret.bin"), b"x").expect("test setup");
let err = contained_shard_path(&dir, "../secret.bin")
.expect_err("parent traversal must be rejected");
let msg = err.to_string();
assert!(
msg.contains("must stay within the model directory")
|| msg.contains("escapes model root"),
"unexpected error: {err}"
);
assert!(
contained_shard_path(&dir, "../nonexistent.bin").is_err(),
"traversal to a missing target must also fail"
);
}
#[cfg(unix)]
#[test]
fn sharded_loader_follows_hub_cache_snapshot_layout() {
let outer = temp_dir("lattice_containment_hub_layout_test");
let blobs = outer.join("blobs");
let dir = outer.join("snapshots").join("rev");
fs::create_dir_all(&blobs).expect("test setup");
fs::create_dir_all(&dir).expect("test setup");
write_single_f32_tensor(&blobs.join("blob-sha"), "tensor.a", &[1.0, 2.0]);
std::os::unix::fs::symlink(
Path::new("../../blobs/blob-sha"),
dir.join("model.safetensors"),
)
.expect("test setup");
fs::write(
dir.join("model.safetensors.index.json"),
r#"{"metadata": {}, "weight_map": {"tensor.a": "model.safetensors"}}"#,
)
.expect("test setup: write index");
let tensors =
load_sharded(&dir).expect("hub-cache snapshot layout must load through the symlink");
assert_eq!(tensors["tensor.a"].data, vec![1.0, 2.0]);
}
#[test]
fn sharded_loaders_reject_escaping_index_entry() {
let outer = temp_dir("lattice_containment_loader_test");
let dir = outer.join("model");
fs::create_dir_all(&dir).expect("test setup");
let evil = outer.join("evil.safetensors");
write_single_f32_tensor(&evil, "tensor.a", &[1.0, 2.0]);
fs::write(
dir.join("model.safetensors.index.json"),
r#"{"metadata": {}, "weight_map": {"tensor.a": "../evil.safetensors"}}"#,
)
.expect("test setup: write index");
let err = load_sharded(&dir).expect_err("load_sharded must reject the escaping entry");
let msg = err.to_string();
assert!(
msg.contains("must stay within the model directory")
|| msg.contains("escapes model root"),
"unexpected error: {err}"
);
let index_path = dir.join("model.safetensors.index.json");
let mut st = ShardedSafetensors::open_index(&index_path).expect("index itself parses");
let err = st
.get_f32_tensor_owned("tensor.a")
.expect_err("lazy loader must reject the escaping entry at access time");
let msg = err.to_string();
assert!(
msg.contains("must stay within the model directory")
|| msg.contains("escapes model root"),
"unexpected error: {err}"
);
let err = st
.resolve_weight("tensor.a")
.expect_err("resolve_weight must reject the escaping entry");
let msg = err.to_string();
assert!(
msg.contains("must stay within the model directory")
|| msg.contains("escapes model root"),
"unexpected error: {err}"
);
}
#[test]
fn test_load_sharded_multi_tensor_per_shard() {
let dir = temp_dir("lattice_multi_tensor_shard_test");
let shard = dir.join("model-00001-of-00001.safetensors");
let x_byte_count = 2_usize * std::mem::size_of::<f32>(); let y_byte_count = 3_usize * std::mem::size_of::<f32>(); let xy_end = x_byte_count + y_byte_count;
let header = format!(
r#"{{"tensor.x":{{"dtype":"F32","shape":[2],"data_offsets":[0,{x_byte_count}]}},"tensor.y":{{"dtype":"F32","shape":[3],"data_offsets":[{x_byte_count},{xy_end}]}}}}"#
);
let mut bytes: Vec<u8> = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for v in [1.0_f32, 2.0] {
bytes.extend_from_slice(&v.to_le_bytes());
}
for v in [3.0_f32, 4.0, 5.0] {
bytes.extend_from_slice(&v.to_le_bytes());
}
fs::write(&shard, &bytes).expect("test setup: write two-tensor shard");
let index_path = dir.join("model.safetensors.index.json");
fs::write(
&index_path,
r#"{"weight_map":{"tensor.x":"model-00001-of-00001.safetensors","tensor.y":"model-00001-of-00001.safetensors"}}"#,
)
.expect("test setup: write index");
let loaded = load_sharded(&dir).expect("load_sharded with multi-tensor shard succeeds");
assert_eq!(loaded.len(), 2);
let tx = loaded.get("tensor.x").expect("tensor.x in loaded map");
let ty = loaded.get("tensor.y").expect("tensor.y in loaded map");
assert_eq!(tx.data, vec![1.0_f32, 2.0]);
assert_eq!(tx.shape, vec![2]);
assert_eq!(ty.data, vec![3.0_f32, 4.0, 5.0]);
assert_eq!(ty.shape, vec![3]);
}
#[test]
fn test_load_sharded_missing_shard_returns_err() {
let dir = temp_dir("lattice_missing_shard_test");
let index_path = dir.join("model.safetensors.index.json");
fs::write(
&index_path,
r#"{"weight_map":{"tensor.a":"model-00001-of-00001.safetensors"}}"#,
)
.expect("test setup: write index");
let result = load_sharded(&dir);
assert!(
result.is_err(),
"load_sharded must return Err when indexed shard file is absent"
);
}
#[test]
fn test_cross_encoder_weights_logit_known_value() {
let weights = CrossEncoderWeights {
classifier_weight: vec![1.0f32, 2.0, -1.0],
classifier_bias: 0.5,
};
let pooled = [2.0f32, 1.0, 3.0];
let logit = weights.logit(&pooled);
assert!(
(logit - 1.5f32).abs() < 1e-6,
"expected logit=1.5, got {logit}"
);
}
fn write_classifier_safetensors(
path: &std::path::Path,
weight_shape_header: &str,
weight_values: &[f32],
bias_value: f32,
weight_end: usize,
) {
let bias_end = weight_end + 4; let header = format!(
r#"{{"classifier.weight":{{{weight_shape_header},"data_offsets":[0,{weight_end}]}},"classifier.bias":{{"dtype":"F32","shape":[1],"data_offsets":[{weight_end},{bias_end}]}}}}"#
);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header.as_bytes());
for v in weight_values {
bytes.extend_from_slice(&v.to_le_bytes());
}
bytes.extend_from_slice(&bias_value.to_le_bytes());
let mut file = File::create(path).expect("test setup: create safetensors");
file.write_all(&bytes)
.expect("test setup: write safetensors");
}
#[test]
fn test_load_cross_encoder_weights_2d_shape() {
let path = temp_path("lattice_ce_weights_2d");
write_classifier_safetensors(
&path,
r#""dtype":"F32","shape":[1,3]"#,
&[1.0f32, 2.0, 3.0],
0.5,
12,
);
let st = SafetensorsFile::open(&path).unwrap();
let weights = st.load_cross_encoder_weights(3).unwrap();
assert_eq!(weights.classifier_weight, vec![1.0f32, 2.0, 3.0]);
assert!(
(weights.classifier_bias - 0.5f32).abs() < 1e-6,
"bias mismatch: {}",
weights.classifier_bias
);
}
#[test]
fn test_load_cross_encoder_weights_1d_shape() {
let path = temp_path("lattice_ce_weights_1d");
write_classifier_safetensors(
&path,
r#""dtype":"F32","shape":[3]"#,
&[4.0f32, 5.0, 6.0],
1.0,
12,
);
let st = SafetensorsFile::open(&path).unwrap();
let weights = st.load_cross_encoder_weights(3).unwrap();
assert_eq!(weights.classifier_weight, vec![4.0f32, 5.0, 6.0]);
assert!(
(weights.classifier_bias - 1.0f32).abs() < 1e-6,
"bias mismatch: {}",
weights.classifier_bias
);
}
#[test]
fn test_load_cross_encoder_weights_shape_mismatch_error() {
let path = temp_path("lattice_ce_weights_mismatch");
write_classifier_safetensors(
&path,
r#""dtype":"F32","shape":[1,4]"#,
&[1.0f32, 2.0, 3.0, 4.0],
0.0,
16,
);
let st = SafetensorsFile::open(&path).unwrap();
let err = st
.load_cross_encoder_weights(3)
.expect_err("shape mismatch should be an error");
assert!(
matches!(err, InferenceError::ShapeMismatch { .. }),
"expected ShapeMismatch, got {err:?}"
);
}
#[test]
fn parse_index_rejects_oversized_index_file() {
let root = temp_dir("lattice_index_oversized");
let oversized = format!(
r#"{{"metadata":{{}},"weight_map":{{}}, "pad":"{}"}}"#,
"x".repeat(MAX_SAFETENSORS_INDEX_BYTES as usize + 1)
);
fs::write(root.join("model.safetensors.index.json"), oversized)
.expect("test setup: write oversized index");
let err = parse_index(&root).expect_err("an oversized safetensors index must be rejected");
assert!(
err.to_string().contains("MAX_SAFETENSORS_INDEX_BYTES"),
"wrong guard fired: {err}"
);
}
#[test]
fn parse_index_rejects_weight_map_over_entry_cap() {
let mut entries = String::with_capacity((MAX_WEIGHT_MAP_ENTRIES + 1) * 24);
for i in 0..=MAX_WEIGHT_MAP_ENTRIES {
if i > 0 {
entries.push(',');
}
entries.push_str(&format!(r#""t{i}":"s""#));
}
let json = format!(r#"{{"metadata":{{}},"weight_map":{{{entries}}}}}"#);
let err = serde_json::from_str::<SafetensorsIndex>(&json)
.expect_err("weight_map over MAX_WEIGHT_MAP_ENTRIES must be rejected");
assert!(
err.to_string().contains("MAX_WEIGHT_MAP_ENTRIES"),
"wrong guard fired: {err}"
);
}
#[test]
fn parse_index_accepts_weight_map_at_entry_cap() {
let mut entries = String::with_capacity(MAX_WEIGHT_MAP_ENTRIES * 24);
for i in 0..MAX_WEIGHT_MAP_ENTRIES {
if i > 0 {
entries.push(',');
}
entries.push_str(&format!(r#""t{i}":"s""#));
}
let json = format!(r#"{{"metadata":{{}},"weight_map":{{{entries}}}}}"#);
let index: SafetensorsIndex = serde_json::from_str(&json)
.expect("weight_map at exactly MAX_WEIGHT_MAP_ENTRIES must be accepted");
assert_eq!(index.weight_map.len(), MAX_WEIGHT_MAP_ENTRIES);
}
#[test]
fn open_manifest_entry_once_rejects_traversal_before_opening() {
let root = temp_dir("lattice_manifest_entry_traversal");
let err = open_manifest_entry_once(&root, "../escape.safetensors")
.expect_err("a `../` manifest entry must be rejected");
assert!(
err.to_string().contains("escapes the model directory"),
"expected the lexical traversal guard to fire, got: {err}"
);
}
#[test]
fn load_sharded_rejects_hostile_index_with_traversal_entry() {
let root = temp_dir("lattice_containment_load_sharded");
let escape_target = root.parent().expect("temp dir has a parent").join(format!(
"lattice_containment_load_sharded_secret_{}.safetensors",
std::process::id()
));
write_single_f32_tensor(&escape_target, "tensor.a", &[1.0, 2.0]);
let index_path = root.join("model.safetensors.index.json");
fs::write(
&index_path,
format!(
r#"{{"metadata": {{"total_size": 8.0}}, "weight_map": {{"tensor.a": "../{}"}}}}"#,
escape_target.file_name().unwrap().to_str().unwrap()
),
)
.expect("test setup: write hostile index");
let result = load_sharded(&root);
fs::remove_file(&escape_target).ok();
assert!(
result.is_err(),
"load_sharded must reject a `../` traversal entry in the index manifest"
);
}
}