use crate::{
ApplyResult, IdentityAdapter, ModuleAdapter, ModuleSnapshot, ModuleStore, PathFilter,
TensorSnapshot,
};
#[cfg(feature = "std")]
use crate::KeyRemapper;
use alloc::boxed::Box;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec;
use alloc::vec::Vec;
use burn_core::module::ParamId;
use burn_tensor::backend::Backend;
use burn_tensor::{DType, TensorData};
use core::fmt;
use core::ops::Deref;
use hashbrown::HashMap;
#[cfg(target_has_atomic = "ptr")]
use alloc::sync::Arc;
#[cfg(not(target_has_atomic = "ptr"))]
type Arc<T> = Box<T>;
#[derive(Debug)]
pub enum SafetensorsStoreError {
Safetensors(safetensors::SafeTensorError),
#[cfg(feature = "std")]
Io(std::io::Error),
TensorNotFound(String),
ValidationFailed(String),
Other(String),
}
impl fmt::Display for SafetensorsStoreError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Safetensors(e) => write!(f, "SafeTensors error: {}", e),
#[cfg(feature = "std")]
Self::Io(e) => write!(f, "I/O error: {}", e),
Self::TensorNotFound(name) => write!(f, "Tensor not found: {}", name),
Self::ValidationFailed(msg) => write!(f, "Validation failed: {}", msg),
Self::Other(msg) => write!(f, "{}", msg),
}
}
}
impl core::error::Error for SafetensorsStoreError {}
impl From<safetensors::SafeTensorError> for SafetensorsStoreError {
fn from(e: safetensors::SafeTensorError) -> Self {
SafetensorsStoreError::Safetensors(e)
}
}
#[cfg(feature = "std")]
impl From<std::io::Error> for SafetensorsStoreError {
fn from(e: std::io::Error) -> Self {
SafetensorsStoreError::Io(e)
}
}
pub enum SafetensorsStore {
#[cfg(feature = "std")]
File(FileStore),
Memory(MemoryStore),
}
impl Default for SafetensorsStore {
fn default() -> Self {
Self::from_bytes(None)
}
}
impl SafetensorsStore {
pub fn default_metadata() -> HashMap<String, String> {
let mut metadata = HashMap::new();
metadata.insert("format".to_string(), "safetensors".to_string());
metadata.insert("producer".to_string(), "burn".to_string());
metadata.insert("version".to_string(), env!("CARGO_PKG_VERSION").to_string());
metadata
}
#[cfg(feature = "std")]
pub fn from_file(path: impl Into<std::path::PathBuf>) -> Self {
Self::File(FileStore {
path: path.into(),
filter: PathFilter::new(),
remapper: KeyRemapper::new(),
metadata: Self::default_metadata(),
validate: true,
allow_partial: false,
overwrite: false,
from_adapter: Box::new(IdentityAdapter),
to_adapter: Box::new(IdentityAdapter),
})
}
pub fn from_bytes(bytes: Option<Vec<u8>>) -> Self {
Self::Memory(MemoryStore {
data: bytes.map(Arc::new),
filter: PathFilter::new(),
#[cfg(feature = "std")]
remapper: KeyRemapper::new(),
metadata: Self::default_metadata(),
validate: true,
allow_partial: false,
from_adapter: Box::new(IdentityAdapter),
to_adapter: Box::new(IdentityAdapter),
})
}
pub fn filter(mut self, filter: PathFilter) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = filter,
Self::Memory(p) => p.filter = filter,
}
self
}
#[cfg(feature = "std")]
pub fn with_regex<S: AsRef<str>>(mut self, pattern: S) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().with_regex(pattern),
Self::Memory(p) => p.filter = p.filter.clone().with_regex(pattern),
}
self
}
#[cfg(feature = "std")]
pub fn with_regexes<I, S>(mut self, patterns: I) -> Self
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().with_regexes(patterns),
Self::Memory(p) => p.filter = p.filter.clone().with_regexes(patterns),
}
self
}
pub fn with_full_path<S: Into<String>>(mut self, path: S) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().with_full_path(path),
Self::Memory(p) => p.filter = p.filter.clone().with_full_path(path),
}
self
}
pub fn with_full_paths<I, S>(mut self, paths: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().with_full_paths(paths),
Self::Memory(p) => p.filter = p.filter.clone().with_full_paths(paths),
}
self
}
pub fn with_predicate(mut self, predicate: fn(&str, &str) -> bool) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().with_predicate(predicate),
Self::Memory(p) => p.filter = p.filter.clone().with_predicate(predicate),
}
self
}
pub fn with_predicates<I>(mut self, predicates: I) -> Self
where
I: IntoIterator<Item = fn(&str, &str) -> bool>,
{
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().with_predicates(predicates),
Self::Memory(p) => p.filter = p.filter.clone().with_predicates(predicates),
}
self
}
pub fn match_all(mut self) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.filter = p.filter.clone().match_all(),
Self::Memory(p) => p.filter = p.filter.clone().match_all(),
}
self
}
#[cfg(feature = "std")]
pub fn remap(mut self, remapper: KeyRemapper) -> Self {
match &mut self {
Self::File(p) => p.remapper = remapper,
Self::Memory(p) => p.remapper = remapper,
}
self
}
#[cfg(feature = "std")]
pub fn with_key_remapping(
mut self,
from_pattern: impl AsRef<str>,
to_pattern: impl Into<String>,
) -> Self {
match &mut self {
Self::File(p) => {
p.remapper = p
.remapper
.clone()
.add_pattern(from_pattern, to_pattern)
.expect("Invalid regex pattern");
}
Self::Memory(p) => {
p.remapper = p
.remapper
.clone()
.add_pattern(from_pattern, to_pattern)
.expect("Invalid regex pattern");
}
}
self
}
pub fn metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
let key = key.into();
let value = value.into();
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => {
p.metadata.insert(key, value);
}
Self::Memory(p) => {
p.metadata.insert(key, value);
}
}
self
}
pub fn clear_metadata(mut self) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => {
p.metadata.clear();
}
Self::Memory(p) => {
p.metadata.clear();
}
}
self
}
pub fn validate(mut self, validate: bool) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.validate = validate,
Self::Memory(p) => p.validate = validate,
}
self
}
pub fn allow_partial(mut self, allow: bool) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.allow_partial = allow,
Self::Memory(p) => p.allow_partial = allow,
}
self
}
#[cfg(feature = "std")]
pub fn overwrite(mut self, overwrite: bool) -> Self {
match &mut self {
Self::File(p) => p.overwrite = overwrite,
Self::Memory(_) => {
}
}
self
}
pub fn with_from_adapter(mut self, adapter: impl ModuleAdapter + 'static) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.from_adapter = Box::new(adapter),
Self::Memory(p) => p.from_adapter = Box::new(adapter),
}
self
}
pub fn with_to_adapter(mut self, adapter: impl ModuleAdapter + 'static) -> Self {
match &mut self {
#[cfg(feature = "std")]
Self::File(p) => p.to_adapter = Box::new(adapter),
Self::Memory(p) => p.to_adapter = Box::new(adapter),
}
self
}
pub fn get_bytes(&self) -> Result<Vec<u8>, SafetensorsStoreError> {
match self {
#[cfg(feature = "std")]
Self::File(_) => Err(SafetensorsStoreError::Other(
"Cannot get bytes from file-based store".to_string(),
)),
Self::Memory(p) => p
.data()
.map(|arc| arc.as_ref().clone())
.ok_or_else(|| SafetensorsStoreError::Other("No data available".to_string())),
}
}
}
#[cfg(feature = "std")]
pub struct FileStore {
path: std::path::PathBuf,
filter: PathFilter,
remapper: KeyRemapper,
metadata: HashMap<String, String>,
validate: bool,
allow_partial: bool,
overwrite: bool,
from_adapter: Box<dyn ModuleAdapter>,
to_adapter: Box<dyn ModuleAdapter>,
}
pub struct MemoryStore {
data: Option<Arc<Vec<u8>>>,
filter: PathFilter,
#[cfg(feature = "std")]
remapper: KeyRemapper,
metadata: HashMap<String, String>,
validate: bool,
allow_partial: bool,
from_adapter: Box<dyn ModuleAdapter>,
to_adapter: Box<dyn ModuleAdapter>,
}
impl Default for MemoryStore {
fn default() -> Self {
Self {
data: None,
filter: PathFilter::new(),
#[cfg(feature = "std")]
remapper: KeyRemapper::new(),
metadata: HashMap::new(),
validate: true,
allow_partial: false,
from_adapter: Box::new(IdentityAdapter),
to_adapter: Box::new(IdentityAdapter),
}
}
}
impl MemoryStore {
#[cfg(test)]
pub(crate) fn data(&self) -> Option<Arc<Vec<u8>>> {
self.data.clone()
}
#[cfg(not(test))]
fn data(&self) -> Option<Arc<Vec<u8>>> {
self.data.clone()
}
#[cfg(test)]
pub(crate) fn set_data(&mut self, data: Vec<u8>) {
self.data = Some(Arc::new(data));
}
}
struct TensorSnapshotAdapter(TensorSnapshot);
impl safetensors::View for TensorSnapshotAdapter {
fn dtype(&self) -> safetensors::Dtype {
dtype_to_safetensors(self.0.dtype).unwrap_or(safetensors::Dtype::F32)
}
fn shape(&self) -> &[usize] {
&self.0.shape
}
fn data(&self) -> alloc::borrow::Cow<'_, [u8]> {
let data = self
.0
.to_data()
.unwrap_or_else(|e| panic!("Failed to get tensor data: {:?}", e));
alloc::borrow::Cow::Owned(data.bytes.deref().to_vec())
}
fn data_len(&self) -> usize {
self.0.data_len()
}
}
impl ModuleStore for SafetensorsStore {
type Error = SafetensorsStoreError;
fn collect_from<B: Backend, M: ModuleSnapshot<B>>(
&mut self,
module: &M,
) -> Result<(), Self::Error> {
let to_adapter = match self {
#[cfg(feature = "std")]
Self::File(p) => p.to_adapter.clone(),
Self::Memory(p) => p.to_adapter.clone(),
};
let mut snapshots = module.collect(None, Some(to_adapter));
snapshots = apply_filter(snapshots, self.get_filter());
#[cfg(feature = "std")]
{
snapshots = apply_remapping(snapshots, self.get_remapper());
}
let metadata = self.get_metadata().clone();
#[cfg(feature = "std")]
let std_metadata: std::collections::HashMap<String, String> = metadata
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
match self {
#[cfg(feature = "std")]
Self::File(p) => {
if p.path.exists() && !p.overwrite {
return Err(SafetensorsStoreError::Other(format!(
"File already exists: {}. Use .overwrite(true) to overwrite.",
p.path.display()
)));
}
let tensors = snapshots_to_safetensors(snapshots)?;
safetensors::serialize_to_file(tensors, Some(std_metadata), &p.path)?;
Ok(())
}
Self::Memory(p) => {
let tensors = snapshots_to_safetensors(snapshots)?;
#[cfg(feature = "std")]
let data = safetensors::serialize(tensors, Some(std_metadata))?;
#[cfg(not(feature = "std"))]
let data = safetensors::serialize(tensors, Some(metadata))?;
p.data = Some(Arc::new(data));
Ok(())
}
}
}
fn apply_to<B: Backend, M: ModuleSnapshot<B>>(
&mut self,
module: &mut M,
) -> Result<ApplyResult, Self::Error> {
#[allow(unused_mut)]
let mut snapshots = match self {
#[cfg(feature = "std")]
Self::File(p) => {
safetensors_to_snapshots_lazy_file(&p.path)?
}
Self::Memory(p) => {
let data_arc = p
.data
.clone()
.ok_or_else(|| SafetensorsStoreError::Other("No data loaded".to_string()))?;
safetensors_to_snapshots_lazy(data_arc)?
}
};
#[cfg(feature = "std")]
{
snapshots = match self {
Self::File(p) => apply_remapping(snapshots, &p.remapper),
Self::Memory(p) => apply_remapping(snapshots, &p.remapper),
};
}
let adapter: Box<dyn ModuleAdapter> = match self {
#[cfg(feature = "std")]
Self::File(p) => p.from_adapter.clone(),
Self::Memory(p) => p.from_adapter.clone(),
};
let result = module.apply(snapshots, None, Some(adapter));
if self.get_validate() && !result.errors.is_empty() {
return Err(SafetensorsStoreError::ValidationFailed(format!(
"Import errors: {:?}",
result.errors
)));
}
if !self.get_allow_partial() && !result.missing.is_empty() {
return Err(SafetensorsStoreError::TensorNotFound(format!(
"Missing tensors: {:?}",
result.missing
)));
}
Ok(result)
}
}
impl SafetensorsStore {
fn get_filter(&self) -> &PathFilter {
match self {
#[cfg(feature = "std")]
Self::File(p) => &p.filter,
Self::Memory(p) => &p.filter,
}
}
#[cfg(feature = "std")]
fn get_remapper(&self) -> &KeyRemapper {
match self {
Self::File(p) => &p.remapper,
Self::Memory(p) => &p.remapper,
}
}
fn get_metadata(&self) -> &HashMap<String, String> {
match self {
#[cfg(feature = "std")]
Self::File(p) => &p.metadata,
Self::Memory(p) => &p.metadata,
}
}
fn get_validate(&self) -> bool {
match self {
#[cfg(feature = "std")]
Self::File(p) => p.validate,
Self::Memory(p) => p.validate,
}
}
fn get_allow_partial(&self) -> bool {
match self {
#[cfg(feature = "std")]
Self::File(p) => p.allow_partial,
Self::Memory(p) => p.allow_partial,
}
}
}
fn apply_filter(mut snapshots: Vec<TensorSnapshot>, filter: &PathFilter) -> Vec<TensorSnapshot> {
if filter.is_empty() {
return snapshots;
}
snapshots.retain(|snapshot| {
let path = snapshot.full_path();
filter.matches(&path)
});
snapshots
}
#[cfg(feature = "std")]
fn apply_remapping(snapshots: Vec<TensorSnapshot>, remapper: &KeyRemapper) -> Vec<TensorSnapshot> {
if remapper.is_empty() {
return snapshots;
}
let (remapped, _) = remapper.remap(snapshots);
remapped
}
fn snapshots_to_safetensors(
snapshots: Vec<TensorSnapshot>,
) -> Result<Vec<(String, TensorSnapshotAdapter)>, SafetensorsStoreError> {
let mut tensors = Vec::new();
for snapshot in snapshots {
let name = snapshot.full_path();
tensors.push((name, TensorSnapshotAdapter(snapshot)));
}
Ok(tensors)
}
fn safetensors_to_snapshots_lazy(
data_arc: Arc<Vec<u8>>,
) -> Result<Vec<TensorSnapshot>, SafetensorsStoreError> {
let tensors = safetensors::SafeTensors::deserialize(&data_arc)?;
let mut snapshots = Vec::new();
for (name, tensor_snapshot) in tensors.tensors() {
let dtype = safetensor_dtype_to_burn(tensor_snapshot.dtype())?;
let shape = tensor_snapshot.shape().to_vec();
let path_parts: Vec<String> = name.split('.').map(|s| s.to_string()).collect();
#[cfg(target_has_atomic = "ptr")]
let data_clone = Arc::clone(&data_arc);
#[cfg(not(target_has_atomic = "ptr"))]
let data_clone = data_arc.clone();
let name_clone = name.to_string();
let data_fn = alloc::rc::Rc::new(move || {
let tensors = safetensors::SafeTensors::deserialize(&data_clone).map_err(|e| {
crate::TensorSnapshotError::IoError(format!(
"Failed to re-deserialize safetensors: {}",
e
))
})?;
let tensor = tensors.tensor(&name_clone).map_err(|e| {
crate::TensorSnapshotError::DataError(format!(
"Tensor '{}' not found: {}",
name_clone, e
))
})?;
let bytes = burn_tensor::Bytes::from_bytes_vec(tensor.data().to_vec());
Ok(TensorData {
bytes,
shape: tensor.shape().to_vec(),
dtype: safetensor_dtype_to_burn(tensor.dtype())
.map_err(|_| crate::TensorSnapshotError::DataError("Invalid dtype".into()))?,
})
});
let snapshot = TensorSnapshot::from_closure(
data_fn,
dtype,
shape,
path_parts,
vec![], ParamId::new(),
);
snapshots.push(snapshot);
}
Ok(snapshots)
}
#[cfg(feature = "std")]
fn safetensors_to_snapshots_lazy_file(
path: &std::path::Path,
) -> Result<Vec<TensorSnapshot>, SafetensorsStoreError> {
use memmap2::MmapOptions;
let file = std::fs::File::open(path)?;
let mmap = unsafe { MmapOptions::new().map(&file)? };
let mmap_arc = Arc::new(mmap);
let tensors = safetensors::SafeTensors::deserialize(&mmap_arc)?;
let mut snapshots = Vec::new();
for (name, tensor_snapshot) in tensors.tensors() {
let dtype = safetensor_dtype_to_burn(tensor_snapshot.dtype())?;
let shape = tensor_snapshot.shape().to_vec();
let path_parts: Vec<String> = name.split('.').map(|s| s.to_string()).collect();
let mmap_clone = Arc::clone(&mmap_arc);
let name_clone = name.to_string();
let data_fn = alloc::rc::Rc::new(move || {
let tensors = safetensors::SafeTensors::deserialize(&mmap_clone).map_err(|e| {
crate::TensorSnapshotError::IoError(format!("Failed to deserialize: {}", e))
})?;
let tensor = tensors.tensor(&name_clone).map_err(|e| {
crate::TensorSnapshotError::DataError(format!(
"Tensor '{}' not found: {}",
name_clone, e
))
})?;
Ok(TensorData {
bytes: burn_tensor::Bytes::from_bytes_vec(tensor.data().to_vec()),
shape: tensor.shape().to_vec(),
dtype: safetensor_dtype_to_burn(tensor.dtype())
.map_err(|_| crate::TensorSnapshotError::DataError("Invalid dtype".into()))?,
})
});
let snapshot = TensorSnapshot::from_closure(
data_fn,
dtype,
shape,
path_parts,
vec![], ParamId::new(),
);
snapshots.push(snapshot);
}
Ok(snapshots)
}
fn safetensor_dtype_to_burn(dtype: safetensors::Dtype) -> Result<DType, SafetensorsStoreError> {
use safetensors::Dtype;
match dtype {
Dtype::F64 => Ok(DType::F64),
Dtype::F32 => Ok(DType::F32),
Dtype::F16 => Ok(DType::F16),
Dtype::BF16 => Ok(DType::BF16),
Dtype::I64 => Ok(DType::I64),
Dtype::I32 => Ok(DType::I32),
Dtype::I16 => Ok(DType::I16),
Dtype::I8 => Ok(DType::I8),
Dtype::U64 => Ok(DType::U64),
Dtype::U32 => Ok(DType::U32),
Dtype::U8 => Ok(DType::U8),
Dtype::BOOL => Ok(DType::Bool),
_ => Err(SafetensorsStoreError::Other(format!(
"Unsupported dtype: {:?}",
dtype
))),
}
}
fn dtype_to_safetensors(dtype: DType) -> Result<safetensors::Dtype, SafetensorsStoreError> {
use safetensors::Dtype;
match dtype {
DType::F64 => Ok(Dtype::F64),
DType::F32 | DType::Flex32 => Ok(Dtype::F32), DType::F16 => Ok(Dtype::F16),
DType::BF16 => Ok(Dtype::BF16),
DType::I64 => Ok(Dtype::I64),
DType::I32 => Ok(Dtype::I32),
DType::I16 => Ok(Dtype::I16),
DType::I8 => Ok(Dtype::I8),
DType::U64 => Ok(Dtype::U64),
DType::U32 => Ok(Dtype::U32),
DType::U16 => Err(SafetensorsStoreError::Other(
"U16 dtype not yet supported in safetensors".to_string(),
)),
DType::U8 => Ok(Dtype::U8),
DType::Bool => Ok(Dtype::BOOL),
DType::QFloat(_) => Err(SafetensorsStoreError::Other(
"Quantized tensors not yet supported in safetensors".to_string(),
)),
}
}