scirs2_io/ml_framework/
utils.rs1#![allow(dead_code)]
3
4use crate::error::{IoError, Result};
5use crate::ml_framework::types::{DataType, MLTensor};
6use byteorder::{LittleEndian, ReadBytesExt, WriteBytesExt};
7use scirs2_core::ndarray::{ArrayD, IxDyn};
8use std::fs::File;
9use std::io::Read;
10
11pub fn tensor_to_python_dict(tensor: &MLTensor) -> Result<serde_json::Value> {
13 Ok(serde_json::json!({
14 "data": tensor.data.as_slice().expect("Operation failed").to_vec(),
15 "shape": tensor.metadata.shape,
16 "dtype": format!("{:?}", tensor.metadata.dtype),
17 "requires_grad": tensor.metadata.requires_grad,
18 }))
19}
20
21pub fn python_dict_to_tensor(dict: &serde_json::Value) -> Result<MLTensor> {
23 let shape: Vec<usize> = serde_json::from_value(dict["shape"].clone())
24 .map_err(|e| IoError::SerializationError(e.to_string()))?;
25
26 let data: Vec<f32> = serde_json::from_value(dict["data"].clone())
27 .map_err(|e| IoError::SerializationError(e.to_string()))?;
28
29 let array =
30 ArrayD::from_shape_vec(IxDyn(&shape), data).map_err(|e| IoError::Other(e.to_string()))?;
31
32 let mut tensor = MLTensor::new(array, None);
33
34 if let Some(requires_grad) = dict.get("requires_grad").and_then(|v| v.as_bool()) {
35 tensor.metadata.requires_grad = requires_grad;
36 }
37
38 Ok(tensor)
39}
40
41#[derive(Debug, Clone)]
43pub struct SafeTensorsHeader {
44 pub tensors: std::collections::HashMap<String, TensorInfo>,
45}
46
47#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
48pub struct TensorInfo {
49 pub dtype: DataType,
50 pub shape: Vec<usize>,
51 pub data_offsets: (usize, usize),
52}
53
54pub fn write_safetensors_header<W: std::io::Write>(
56 writer: &mut W,
57 header: &SafeTensorsHeader,
58) -> Result<()> {
59 let header_json = serde_json::to_string(&header.tensors)
60 .map_err(|e| IoError::SerializationError(e.to_string()))?;
61
62 writer
63 .write_u64::<LittleEndian>(header_json.len() as u64)
64 .map_err(IoError::Io)?;
65 writer
66 .write_all(header_json.as_bytes())
67 .map_err(IoError::Io)?;
68
69 Ok(())
70}
71
72pub fn read_safetensors_header<R: std::io::Read>(reader: &mut R) -> Result<SafeTensorsHeader> {
74 let header_size = reader.read_u64::<LittleEndian>().map_err(IoError::Io)?;
75 let mut header_bytes = vec![0u8; header_size as usize];
76 reader.read_exact(&mut header_bytes).map_err(IoError::Io)?;
77
78 let header_str = String::from_utf8(header_bytes).map_err(|e| IoError::Other(e.to_string()))?;
79
80 let tensors: std::collections::HashMap<String, TensorInfo> = serde_json::from_str(&header_str)
81 .map_err(|e| IoError::SerializationError(e.to_string()))?;
82
83 Ok(SafeTensorsHeader { tensors })
84}