Files
rustytorch/crates/production/rtx-hub/src/safetensors.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
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]>
2026-04-12 07:01:58 -07:00

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())
);
}
}