//! 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, /// Optional metadata #[serde(rename = "__metadata__", skip_serializing_if = "Option::is_none")] pub metadata: Option>, } /// 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, /// 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 { 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, /// Header size in bytes header_size: usize, } impl SafeTensors { /// Load a SafeTensors file from path. pub async fn load(path: impl AsRef) -> HubResult { 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 { 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> { 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::() * 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)>, metadata: Option>, } 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, dtype: SafeTensorsDType, shape: Vec, data: Vec, ) -> 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, value: impl Into) -> 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> { // 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 = 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, /// Index file data (if present) index: Option, } /// SafeTensors index file (model.safetensors.index.json). #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SafeTensorsIndex { /// Metadata about the sharded model pub metadata: Option>, /// Map from tensor name to shard filename pub weight_map: HashMap, } impl ShardedSafeTensors { /// Load a sharded SafeTensors model from a directory. pub async fn load_from_dir(dir: impl AsRef) -> HubResult { 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::(&index_data)?) } else { None }; // Find all .safetensors files let mut shard_files: Vec = 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 = 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 = 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()) ); } }