Initial commit
This commit is contained in:
@@ -0,0 +1,593 @@
|
||||
//! 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())
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user