Consistent formatting pass: line wrapping, import sorting, trailing whitespace removal, let-chain indentation, merged derive attributes, and unsafe block reformatting. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
598 lines
18 KiB
Rust
598 lines
18 KiB
Rust
//! SafeTensors file format loading and manipulation.
|
|
//!
|
|
//! SafeTensors is a simple, safe file format for storing tensors. This module provides
|
|
//! comprehensive support for loading, saving, and manipulating SafeTensors files,
|
|
//! enabling seamless integration with HuggingFace model weights.
|
|
|
|
use crate::{HubError, HubResult};
|
|
use serde::{Deserialize, Serialize};
|
|
use std::collections::HashMap;
|
|
use std::path::Path;
|
|
use tokio::fs;
|
|
|
|
/// SafeTensors file header containing tensor metadata.
|
|
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
|
pub struct SafeTensorsHeader {
|
|
/// Tensor metadata indexed by tensor name
|
|
#[serde(flatten)]
|
|
pub tensors: HashMap<String, TensorInfo>,
|
|
/// Optional metadata
|
|
#[serde(rename = "__metadata__", skip_serializing_if = "Option::is_none")]
|
|
pub metadata: Option<HashMap<String, String>>,
|
|
}
|
|
|
|
/// Information about a single tensor in the SafeTensors file.
|
|
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
|
pub struct TensorInfo {
|
|
/// Data type of the tensor
|
|
pub dtype: SafeTensorsDType,
|
|
/// Shape of the tensor
|
|
pub shape: Vec<usize>,
|
|
/// Byte offsets [start, end) in the data section
|
|
pub data_offsets: [usize; 2],
|
|
}
|
|
|
|
/// Data types supported by SafeTensors.
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
|
#[serde(rename_all = "UPPERCASE")]
|
|
pub enum SafeTensorsDType {
|
|
/// Boolean
|
|
Bool,
|
|
/// Unsigned 8-bit integer
|
|
U8,
|
|
/// Signed 8-bit integer
|
|
I8,
|
|
/// Signed 16-bit integer
|
|
I16,
|
|
/// Signed 32-bit integer
|
|
I32,
|
|
/// Signed 64-bit integer
|
|
I64,
|
|
/// 16-bit floating point (half precision)
|
|
F16,
|
|
/// Brain floating point (16-bit)
|
|
#[serde(rename = "BF16")]
|
|
BF16,
|
|
/// 32-bit floating point
|
|
F32,
|
|
/// 64-bit floating point
|
|
F64,
|
|
/// FP8 E4M3 format (8-bit floating point)
|
|
#[serde(rename = "F8_E4M3")]
|
|
F8E4M3,
|
|
/// FP8 E5M2 format (8-bit floating point)
|
|
#[serde(rename = "F8_E5M2")]
|
|
F8E5M2,
|
|
}
|
|
|
|
impl SafeTensorsDType {
|
|
/// Get the size in bytes for this dtype.
|
|
pub fn size_bytes(&self) -> usize {
|
|
match self {
|
|
SafeTensorsDType::Bool | SafeTensorsDType::U8 | SafeTensorsDType::I8 => 1,
|
|
SafeTensorsDType::F8E4M3 | SafeTensorsDType::F8E5M2 => 1,
|
|
SafeTensorsDType::I16 | SafeTensorsDType::F16 | SafeTensorsDType::BF16 => 2,
|
|
SafeTensorsDType::I32 | SafeTensorsDType::F32 => 4,
|
|
SafeTensorsDType::I64 | SafeTensorsDType::F64 => 8,
|
|
}
|
|
}
|
|
|
|
/// Convert from string representation.
|
|
pub fn from_str(s: &str) -> Option<Self> {
|
|
match s.to_uppercase().as_str() {
|
|
"BOOL" => Some(SafeTensorsDType::Bool),
|
|
"U8" => Some(SafeTensorsDType::U8),
|
|
"I8" => Some(SafeTensorsDType::I8),
|
|
"I16" => Some(SafeTensorsDType::I16),
|
|
"I32" => Some(SafeTensorsDType::I32),
|
|
"I64" => Some(SafeTensorsDType::I64),
|
|
"F16" => Some(SafeTensorsDType::F16),
|
|
"BF16" => Some(SafeTensorsDType::BF16),
|
|
"F32" => Some(SafeTensorsDType::F32),
|
|
"F64" => Some(SafeTensorsDType::F64),
|
|
"F8_E4M3" => Some(SafeTensorsDType::F8E4M3),
|
|
"F8_E5M2" => Some(SafeTensorsDType::F8E5M2),
|
|
_ => None,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// A loaded SafeTensors file.
|
|
#[derive(Debug)]
|
|
pub struct SafeTensors {
|
|
/// Header containing tensor metadata
|
|
pub header: SafeTensorsHeader,
|
|
/// Raw tensor data
|
|
data: Vec<u8>,
|
|
/// Header size in bytes
|
|
header_size: usize,
|
|
}
|
|
|
|
impl SafeTensors {
|
|
/// Load a SafeTensors file from path.
|
|
pub async fn load(path: impl AsRef<Path>) -> HubResult<Self> {
|
|
let data = fs::read(path.as_ref()).await?;
|
|
Self::from_bytes(&data)
|
|
}
|
|
|
|
/// Load a SafeTensors file from bytes.
|
|
pub fn from_bytes(bytes: &[u8]) -> HubResult<Self> {
|
|
if bytes.len() < 8 {
|
|
return Err(HubError::InvalidPackage {
|
|
reason: "SafeTensors file too small (< 8 bytes)".to_string(),
|
|
});
|
|
}
|
|
|
|
// Read header size (8-byte little-endian)
|
|
let header_size = u64::from_le_bytes([
|
|
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
|
|
]) as usize;
|
|
|
|
if header_size > 100_000_000 {
|
|
return Err(HubError::InvalidPackage {
|
|
reason: format!("SafeTensors header too large: {} bytes", header_size),
|
|
});
|
|
}
|
|
|
|
if bytes.len() < 8 + header_size {
|
|
return Err(HubError::InvalidPackage {
|
|
reason: format!(
|
|
"SafeTensors file truncated: expected {} bytes, got {}",
|
|
8 + header_size,
|
|
bytes.len()
|
|
),
|
|
});
|
|
}
|
|
|
|
// Parse header JSON
|
|
let header_bytes = &bytes[8..8 + header_size];
|
|
let header: SafeTensorsHeader =
|
|
serde_json::from_slice(header_bytes).map_err(|e| HubError::InvalidPackage {
|
|
reason: format!("Invalid SafeTensors header: {}", e),
|
|
})?;
|
|
|
|
// Extract tensor data
|
|
let data = bytes[8 + header_size..].to_vec();
|
|
|
|
Ok(Self {
|
|
header,
|
|
data,
|
|
header_size,
|
|
})
|
|
}
|
|
|
|
/// Get the names of all tensors in the file.
|
|
pub fn tensor_names(&self) -> Vec<&str> {
|
|
self.header
|
|
.tensors
|
|
.keys()
|
|
.map(std::string::String::as_str)
|
|
.collect()
|
|
}
|
|
|
|
/// Get information about a specific tensor.
|
|
pub fn tensor_info(&self, name: &str) -> Option<&TensorInfo> {
|
|
self.header.tensors.get(name)
|
|
}
|
|
|
|
/// Get the raw bytes for a specific tensor.
|
|
pub fn tensor_data(&self, name: &str) -> Option<&[u8]> {
|
|
self.header.tensors.get(name).map(|info| {
|
|
let start = info.data_offsets[0];
|
|
let end = info.data_offsets[1];
|
|
&self.data[start..end]
|
|
})
|
|
}
|
|
|
|
/// Get all tensor data as a map.
|
|
pub fn all_tensors(&self) -> HashMap<&str, (&TensorInfo, &[u8])> {
|
|
self.header
|
|
.tensors
|
|
.iter()
|
|
.map(|(name, info)| {
|
|
let start = info.data_offsets[0];
|
|
let end = info.data_offsets[1];
|
|
(name.as_str(), (info, &self.data[start..end]))
|
|
})
|
|
.collect()
|
|
}
|
|
|
|
/// Get file metadata.
|
|
pub fn metadata(&self) -> Option<&HashMap<String, String>> {
|
|
self.header.metadata.as_ref()
|
|
}
|
|
|
|
/// Get total file size.
|
|
pub fn file_size(&self) -> usize {
|
|
8 + self.header_size + self.data.len()
|
|
}
|
|
|
|
/// Get total number of tensors.
|
|
pub fn num_tensors(&self) -> usize {
|
|
self.header.tensors.len()
|
|
}
|
|
|
|
/// Validate tensor data integrity.
|
|
pub fn validate(&self) -> HubResult<()> {
|
|
for (name, info) in &self.header.tensors {
|
|
let start = info.data_offsets[0];
|
|
let end = info.data_offsets[1];
|
|
|
|
if start > end {
|
|
return Err(HubError::ValidationFailed {
|
|
details: format!("Tensor '{}' has invalid offsets: {} > {}", name, start, end),
|
|
});
|
|
}
|
|
|
|
if end > self.data.len() {
|
|
return Err(HubError::ValidationFailed {
|
|
details: format!(
|
|
"Tensor '{}' data exceeds file bounds: {} > {}",
|
|
name,
|
|
end,
|
|
self.data.len()
|
|
),
|
|
});
|
|
}
|
|
|
|
// Verify data size matches shape
|
|
let expected_size: usize =
|
|
info.shape.iter().product::<usize>() * info.dtype.size_bytes();
|
|
let actual_size = end - start;
|
|
|
|
if expected_size != actual_size {
|
|
return Err(HubError::ValidationFailed {
|
|
details: format!(
|
|
"Tensor '{}' size mismatch: expected {} bytes, got {}",
|
|
name, expected_size, actual_size
|
|
),
|
|
});
|
|
}
|
|
}
|
|
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
/// Builder for creating SafeTensors files.
|
|
pub struct SafeTensorsBuilder {
|
|
tensors: Vec<(String, TensorInfo, Vec<u8>)>,
|
|
metadata: Option<HashMap<String, String>>,
|
|
}
|
|
|
|
impl SafeTensorsBuilder {
|
|
/// Create a new SafeTensors builder.
|
|
pub fn new() -> Self {
|
|
Self {
|
|
tensors: Vec::new(),
|
|
metadata: None,
|
|
}
|
|
}
|
|
|
|
/// Add a tensor to the file.
|
|
pub fn add_tensor(
|
|
mut self,
|
|
name: impl Into<String>,
|
|
dtype: SafeTensorsDType,
|
|
shape: Vec<usize>,
|
|
data: Vec<u8>,
|
|
) -> Self {
|
|
let name = name.into();
|
|
let info = TensorInfo {
|
|
dtype,
|
|
shape,
|
|
data_offsets: [0, 0], // Will be computed during build
|
|
};
|
|
self.tensors.push((name, info, data));
|
|
self
|
|
}
|
|
|
|
/// Add metadata to the file.
|
|
pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
|
|
self.metadata
|
|
.get_or_insert_with(HashMap::new)
|
|
.insert(key.into(), value.into());
|
|
self
|
|
}
|
|
|
|
/// Build the SafeTensors file.
|
|
pub fn build(mut self) -> HubResult<Vec<u8>> {
|
|
// Calculate offsets
|
|
let mut current_offset = 0usize;
|
|
for (_, info, data) in &mut self.tensors {
|
|
info.data_offsets = [current_offset, current_offset + data.len()];
|
|
current_offset += data.len();
|
|
}
|
|
|
|
// Build header
|
|
let mut header_tensors: HashMap<String, TensorInfo> = HashMap::new();
|
|
for (name, info, _) in &self.tensors {
|
|
header_tensors.insert(name.clone(), info.clone());
|
|
}
|
|
|
|
let header = SafeTensorsHeader {
|
|
tensors: header_tensors,
|
|
metadata: self.metadata,
|
|
};
|
|
|
|
let header_json = serde_json::to_vec(&header)?;
|
|
let header_size = header_json.len() as u64;
|
|
|
|
// Build file
|
|
let mut result = Vec::with_capacity(8 + header_json.len() + current_offset);
|
|
result.extend_from_slice(&header_size.to_le_bytes());
|
|
result.extend_from_slice(&header_json);
|
|
|
|
for (_, _, data) in self.tensors {
|
|
result.extend_from_slice(&data);
|
|
}
|
|
|
|
Ok(result)
|
|
}
|
|
}
|
|
|
|
impl Default for SafeTensorsBuilder {
|
|
fn default() -> Self {
|
|
Self::new()
|
|
}
|
|
}
|
|
|
|
/// Load multiple SafeTensors files (sharded models).
|
|
pub struct ShardedSafeTensors {
|
|
/// List of loaded shards
|
|
shards: Vec<SafeTensors>,
|
|
/// Index file data (if present)
|
|
index: Option<SafeTensorsIndex>,
|
|
}
|
|
|
|
/// SafeTensors index file (model.safetensors.index.json).
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct SafeTensorsIndex {
|
|
/// Metadata about the sharded model
|
|
pub metadata: Option<HashMap<String, serde_json::Value>>,
|
|
/// Map from tensor name to shard filename
|
|
pub weight_map: HashMap<String, String>,
|
|
}
|
|
|
|
impl ShardedSafeTensors {
|
|
/// Load a sharded SafeTensors model from a directory.
|
|
pub async fn load_from_dir(dir: impl AsRef<Path>) -> HubResult<Self> {
|
|
let dir = dir.as_ref();
|
|
|
|
// Check for index file
|
|
let index_path = dir.join("model.safetensors.index.json");
|
|
let index = if index_path.exists() {
|
|
let index_data = fs::read_to_string(&index_path).await?;
|
|
Some(serde_json::from_str::<SafeTensorsIndex>(&index_data)?)
|
|
} else {
|
|
None
|
|
};
|
|
|
|
// Find all .safetensors files
|
|
let mut shard_files: Vec<std::path::PathBuf> = Vec::new();
|
|
|
|
let mut entries = fs::read_dir(dir).await?;
|
|
while let Some(entry) = entries.next_entry().await? {
|
|
let path = entry.path();
|
|
if let Some(ext) = path.extension() {
|
|
if ext == "safetensors" {
|
|
shard_files.push(path);
|
|
}
|
|
}
|
|
}
|
|
|
|
// Sort shard files for consistent ordering
|
|
shard_files.sort();
|
|
|
|
// Load all shards
|
|
let mut shards = Vec::with_capacity(shard_files.len());
|
|
for shard_path in shard_files {
|
|
let shard = SafeTensors::load(&shard_path).await?;
|
|
shards.push(shard);
|
|
}
|
|
|
|
Ok(Self { shards, index })
|
|
}
|
|
|
|
/// Get all tensor names across all shards.
|
|
pub fn tensor_names(&self) -> Vec<&str> {
|
|
let mut names: Vec<&str> = self.shards.iter().flat_map(|s| s.tensor_names()).collect();
|
|
names.sort_unstable();
|
|
names.dedup();
|
|
names
|
|
}
|
|
|
|
/// Get tensor data by name (searches all shards).
|
|
pub fn tensor_data(&self, name: &str) -> Option<(&TensorInfo, &[u8])> {
|
|
for shard in &self.shards {
|
|
if let Some(info) = shard.tensor_info(name) {
|
|
let data = shard.tensor_data(name)?;
|
|
return Some((info, data));
|
|
}
|
|
}
|
|
None
|
|
}
|
|
|
|
/// Get the index file data.
|
|
pub fn index(&self) -> Option<&SafeTensorsIndex> {
|
|
self.index.as_ref()
|
|
}
|
|
|
|
/// Get number of shards.
|
|
pub fn num_shards(&self) -> usize {
|
|
self.shards.len()
|
|
}
|
|
|
|
/// Get total number of tensors across all shards.
|
|
pub fn num_tensors(&self) -> usize {
|
|
self.shards.iter().map(SafeTensors::num_tensors).sum()
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_dtype_size() {
|
|
assert_eq!(SafeTensorsDType::Bool.size_bytes(), 1);
|
|
assert_eq!(SafeTensorsDType::F16.size_bytes(), 2);
|
|
assert_eq!(SafeTensorsDType::BF16.size_bytes(), 2);
|
|
assert_eq!(SafeTensorsDType::F32.size_bytes(), 4);
|
|
assert_eq!(SafeTensorsDType::F64.size_bytes(), 8);
|
|
assert_eq!(SafeTensorsDType::F8E4M3.size_bytes(), 1);
|
|
assert_eq!(SafeTensorsDType::F8E5M2.size_bytes(), 1);
|
|
}
|
|
|
|
#[test]
|
|
fn test_dtype_from_str() {
|
|
assert_eq!(
|
|
SafeTensorsDType::from_str("F32"),
|
|
Some(SafeTensorsDType::F32)
|
|
);
|
|
assert_eq!(
|
|
SafeTensorsDType::from_str("f16"),
|
|
Some(SafeTensorsDType::F16)
|
|
);
|
|
assert_eq!(
|
|
SafeTensorsDType::from_str("BF16"),
|
|
Some(SafeTensorsDType::BF16)
|
|
);
|
|
assert_eq!(
|
|
SafeTensorsDType::from_str("F8_E4M3"),
|
|
Some(SafeTensorsDType::F8E4M3)
|
|
);
|
|
assert_eq!(SafeTensorsDType::from_str("invalid"), None);
|
|
}
|
|
|
|
#[test]
|
|
fn test_safetensors_builder() {
|
|
let tensor_data = vec![0u8; 16]; // 4 floats
|
|
|
|
let file = SafeTensorsBuilder::new()
|
|
.add_tensor("test.weight", SafeTensorsDType::F32, vec![4], tensor_data)
|
|
.with_metadata("format", "pt")
|
|
.build()
|
|
.unwrap();
|
|
|
|
// Verify file structure
|
|
assert!(file.len() > 8);
|
|
|
|
// Load back
|
|
let loaded = SafeTensors::from_bytes(&file).unwrap();
|
|
assert_eq!(loaded.num_tensors(), 1);
|
|
assert!(loaded.tensor_names().contains(&"test.weight"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_safetensors_from_bytes() {
|
|
// Create a simple SafeTensors file
|
|
let mut header: HashMap<String, TensorInfo> = HashMap::new();
|
|
header.insert(
|
|
"layer.weight".to_string(),
|
|
TensorInfo {
|
|
dtype: SafeTensorsDType::F32,
|
|
shape: vec![2, 2],
|
|
data_offsets: [0, 16],
|
|
},
|
|
);
|
|
|
|
let header_struct = SafeTensorsHeader {
|
|
tensors: header,
|
|
metadata: Some({
|
|
let mut m = HashMap::new();
|
|
m.insert("format".to_string(), "pt".to_string());
|
|
m
|
|
}),
|
|
};
|
|
|
|
let header_json = serde_json::to_vec(&header_struct).unwrap();
|
|
let header_size = header_json.len() as u64;
|
|
|
|
let mut file_data = Vec::new();
|
|
file_data.extend_from_slice(&header_size.to_le_bytes());
|
|
file_data.extend_from_slice(&header_json);
|
|
file_data.extend_from_slice(&[0u8; 16]); // Tensor data
|
|
|
|
let loaded = SafeTensors::from_bytes(&file_data).unwrap();
|
|
|
|
assert_eq!(loaded.num_tensors(), 1);
|
|
assert!(loaded.tensor_names().contains(&"layer.weight"));
|
|
|
|
let info = loaded.tensor_info("layer.weight").unwrap();
|
|
assert_eq!(info.dtype, SafeTensorsDType::F32);
|
|
assert_eq!(info.shape, vec![2, 2]);
|
|
|
|
let data = loaded.tensor_data("layer.weight").unwrap();
|
|
assert_eq!(data.len(), 16);
|
|
|
|
assert!(loaded.validate().is_ok());
|
|
}
|
|
|
|
#[test]
|
|
fn test_safetensors_validation() {
|
|
// Create a file with mismatched size
|
|
let mut header: HashMap<String, TensorInfo> = HashMap::new();
|
|
header.insert(
|
|
"test".to_string(),
|
|
TensorInfo {
|
|
dtype: SafeTensorsDType::F32,
|
|
shape: vec![4, 4], // 16 floats = 64 bytes expected
|
|
data_offsets: [0, 16], // Only 16 bytes
|
|
},
|
|
);
|
|
|
|
let header_struct = SafeTensorsHeader {
|
|
tensors: header,
|
|
metadata: None,
|
|
};
|
|
|
|
let header_json = serde_json::to_vec(&header_struct).unwrap();
|
|
let header_size = header_json.len() as u64;
|
|
|
|
let mut file_data = Vec::new();
|
|
file_data.extend_from_slice(&header_size.to_le_bytes());
|
|
file_data.extend_from_slice(&header_json);
|
|
file_data.extend_from_slice(&[0u8; 16]);
|
|
|
|
let loaded = SafeTensors::from_bytes(&file_data).unwrap();
|
|
assert!(loaded.validate().is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_safetensors_invalid_file() {
|
|
// Too small
|
|
assert!(SafeTensors::from_bytes(&[0u8; 4]).is_err());
|
|
|
|
// Header size too large
|
|
let mut bad_header_size = Vec::new();
|
|
bad_header_size.extend_from_slice(&(200_000_000u64).to_le_bytes());
|
|
assert!(SafeTensors::from_bytes(&bad_header_size).is_err());
|
|
|
|
// Truncated file
|
|
let mut truncated = Vec::new();
|
|
truncated.extend_from_slice(&(100u64).to_le_bytes());
|
|
truncated.extend_from_slice(&[0u8; 50]); // Not enough for header
|
|
assert!(SafeTensors::from_bytes(&truncated).is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_safetensors_index_deserialization() {
|
|
let index_json = r#"{
|
|
"metadata": {"total_size": 12345},
|
|
"weight_map": {
|
|
"layer.0.weight": "model-00001-of-00002.safetensors",
|
|
"layer.1.weight": "model-00002-of-00002.safetensors"
|
|
}
|
|
}"#;
|
|
|
|
let index: SafeTensorsIndex = serde_json::from_str(index_json).unwrap();
|
|
assert!(index.metadata.is_some());
|
|
assert_eq!(index.weight_map.len(), 2);
|
|
assert_eq!(
|
|
index.weight_map.get("layer.0.weight"),
|
|
Some(&"model-00001-of-00002.safetensors".to_string())
|
|
);
|
|
}
|
|
}
|