use crate::{Error, LoadedModel};
use std::path::Path;
pub fn save_as_safetensors(
model: &LoadedModel,
path: &Path,
preserve_metadata: bool,
) -> Result<(), Error> {
use std::fs::File;
use std::io::Write;
let header = format!(
r#"{{"__metadata__":{{"converted_by":"mlmf","conversion_time":"{}"}}}}"#,
chrono::Utc::now().to_rfc3339()
);
let header_len = header.len() as u64;
let header_bytes = header_len.to_le_bytes();
let mut serialized = Vec::new();
serialized.extend_from_slice(&header_bytes);
serialized.extend_from_slice(header.as_bytes());
let mut file = File::create(path)
.map_err(|e| Error::io_error(format!("Failed to create file {}: {}", path.display(), e)))?;
file.write_all(&serialized)
.map_err(|e| Error::io_error(format!("Failed to write SafeTensors data: {}", e)))?;
Ok(())
}
pub fn save_safetensors_with_metadata(
path: &Path,
tensors: &std::collections::HashMap<String, candle_core::Tensor>,
metadata: &std::collections::HashMap<String, String>,
) -> Result<(), Error> {
use std::fs::File;
use std::io::Write;
let mut serialized_data = Vec::new();
let mut header_dict = serde_json::Map::new();
if !metadata.is_empty() {
let metadata_value = serde_json::Value::Object(
metadata
.iter()
.map(|(k, v)| (k.clone(), serde_json::Value::String(v.clone())))
.collect(),
);
header_dict.insert("__metadata__".to_string(), metadata_value);
}
let header = serde_json::Value::Object(header_dict);
let header_string = serde_json::to_string(&header)
.map_err(|e| Error::model_saving(format!("Failed to serialize header: {}", e)))?;
let header_len = header_string.len() as u64;
serialized_data.extend_from_slice(&header_len.to_le_bytes());
serialized_data.extend_from_slice(header_string.as_bytes());
let mut file = File::create(path)
.map_err(|e| Error::io_error(format!("Failed to create file {}: {}", path.display(), e)))?;
file.write_all(&serialized_data)
.map_err(|e| Error::io_error(format!("Failed to write SafeTensors data: {}", e)))?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::TensorInfo;
use std::collections::HashMap;
use tempfile::TempDir;
#[test]
fn test_save_as_safetensors() {
let temp_dir = TempDir::new().unwrap();
let output_path = temp_dir.path().join("test.safetensors");
}
}