Skip to main content

scirs2_io/ml_framework/
utils.rs

1//! Utility functions for ML framework operations
2#![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
11/// Convert tensor to Python-compatible dictionary
12pub 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
21/// Convert Python dictionary to tensor
22pub 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/// SafeTensors header structure
42#[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
54/// Write SafeTensors header
55pub 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
72/// Read SafeTensors header
73pub 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}