use crate::{Result, WeightMatrix};
use crate::model_fetcher::{FetchResult, ModelMetadata, ModelFormat};
use std::path::PathBuf;
use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use std::fs::File;
use std::io::{Read, Seek, SeekFrom};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelStats {
pub total_parameters: usize,
pub total_size_mb: f64,
pub num_tensors: usize,
pub largest_tensor: usize,
pub smallest_tensor: usize,
pub average_tensor_size: usize,
}
pub struct RealModelLoader;
impl RealModelLoader {
pub fn load_model(fetch_result: &FetchResult) -> Result<Vec<WeightMatrix>> {
match &fetch_result.model_format {
ModelFormat::SafeTensors => Self::load_safetensors(&fetch_result.local_path),
ModelFormat::PyTorch => Self::load_pytorch(&fetch_result.local_path),
ModelFormat::GGUF => Self::load_gguf(&fetch_result.local_path),
ModelFormat::ONNX => Self::load_onnx(&fetch_result.local_path),
ModelFormat::Unknown => Err("Unknown model format".into()),
}
}
fn load_safetensors(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
let mut file = File::open(path)?;
let mut header_len_bytes = [0u8; 8];
file.read_exact(&mut header_len_bytes)?;
let header_len = u64::from_le_bytes(header_len_bytes) as usize;
let mut header_json = vec![0u8; header_len];
file.read_exact(&mut header_json)?;
let header_str = String::from_utf8(header_json)?;
let header: HashMap<String, serde_json::Value> = serde_json::from_str(&header_str)?;
let mut weights = Vec::new();
let offset = 8 + header_len as u64;
for (tensor_name, tensor_info) in header {
if let Some(tensor_obj) = tensor_info.as_object() {
let dtype = tensor_obj["dtype"].as_str().unwrap_or("F32");
let shape = tensor_obj["shape"].as_array()
.ok_or("Invalid shape")?
.iter()
.filter_map(|v| v.as_u64().map(|n| n as usize))
.collect::<Vec<_>>();
let data_offsets = tensor_obj["data_offsets"].as_array()
.ok_or("Invalid data_offsets")?;
let start_offset = data_offsets[0].as_u64().unwrap_or(0) as u64;
let end_offset = data_offsets[1].as_u64().unwrap_or(0) as u64;
let tensor_size = (end_offset - start_offset) as usize;
file.seek(SeekFrom::Start(offset + start_offset))?;
let tensor_data = match dtype {
"F32" => {
let mut data = vec![0f32; tensor_size / 4];
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(4).enumerate() {
if i < data.len() {
data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
}
}
data
},
"F16" => {
let mut data = vec![0f32; tensor_size / 2];
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(2).enumerate() {
if i < data.len() {
let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
data[i] = half::f16::from_bits(f16_val).to_f32();
}
}
data
},
_ => return Err(format!("Unsupported dtype: {}", dtype).into()),
};
let total_elements: usize = shape.iter().product();
if tensor_data.len() != total_elements {
return Err(format!("Tensor size mismatch for {}: expected {}, got {}",
tensor_name, total_elements, tensor_data.len()).into());
}
weights.push(WeightMatrix::new(tensor_data, shape, tensor_name));
}
}
Ok(weights)
}
fn load_pytorch(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
let mut file = File::open(path)?;
let mut buffer = Vec::new();
file.read_to_end(&mut buffer)?;
let mut weights = Vec::new();
let mut pos = 0;
while pos < buffer.len() {
if pos + 8 < buffer.len() {
let potential_size = u64::from_le_bytes([
buffer[pos], buffer[pos+1], buffer[pos+2], buffer[pos+3],
buffer[pos+4], buffer[pos+5], buffer[pos+6], buffer[pos+7]
]);
if potential_size > 0 && potential_size < 1_000_000_000 { let tensor_size = potential_size as usize;
if pos + 8 + tensor_size * 4 <= buffer.len() {
let mut tensor_data = vec![0f32; tensor_size];
for i in 0..tensor_size {
let byte_pos = pos + 8 + i * 4;
if byte_pos + 3 < buffer.len() {
tensor_data[i] = f32::from_le_bytes([
buffer[byte_pos], buffer[byte_pos+1],
buffer[byte_pos+2], buffer[byte_pos+3]
]);
}
}
let dim = (tensor_size as f64).sqrt() as usize;
let shape = vec![dim, dim];
weights.push(WeightMatrix::new(
tensor_data,
shape,
format!("tensor_{}", weights.len())
));
pos += 8 + tensor_size * 4;
continue;
}
}
}
pos += 1;
}
if weights.is_empty() {
return Err("Could not extract tensors from PyTorch file".into());
}
Ok(weights)
}
fn load_gguf(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
let mut file = File::open(path)?;
let mut magic = [0u8; 4];
file.read_exact(&mut magic)?;
if &magic != b"GGUF" {
return Err("Invalid GGUF magic number".into());
}
let mut version = [0u8; 4];
file.read_exact(&mut version)?;
let _version_num = u32::from_le_bytes(version);
let mut tensor_count = [0u8; 8];
file.read_exact(&mut tensor_count)?;
let num_tensors = u64::from_le_bytes(tensor_count);
let mut metadata_size = [0u8; 8];
file.read_exact(&mut metadata_size)?;
let metadata_len = u64::from_le_bytes(metadata_size) as usize;
file.seek(SeekFrom::Current(metadata_len as i64))?;
let mut weights = Vec::new();
for _i in 0..num_tensors {
let mut name_len = [0u8; 4];
file.read_exact(&mut name_len)?;
let name_length = u32::from_le_bytes(name_len) as usize;
let mut name_bytes = vec![0u8; name_length];
file.read_exact(&mut name_bytes)?;
let tensor_name = String::from_utf8(name_bytes)?;
let mut dims = [0u8; 4];
file.read_exact(&mut dims)?;
let num_dims = u32::from_le_bytes(dims) as usize;
let mut shape = Vec::new();
for _ in 0..num_dims {
let mut dim = [0u8; 8];
file.read_exact(&mut dim)?;
shape.push(u64::from_le_bytes(dim) as usize);
}
let mut tensor_type = [0u8; 4];
file.read_exact(&mut tensor_type)?;
let dtype = u32::from_le_bytes(tensor_type);
let mut offset = [0u8; 8];
file.read_exact(&mut offset)?;
let tensor_offset = u64::from_le_bytes(offset);
let total_elements: usize = shape.iter().product();
let bytes_per_element = match dtype {
0 => 4, 1 => 2, _ => 4, };
let tensor_size = total_elements * bytes_per_element;
let current_pos = file.stream_position()?;
file.seek(SeekFrom::Start(tensor_offset))?;
let mut tensor_data = vec![0f32; total_elements];
match dtype {
0 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(4).enumerate() {
if i < tensor_data.len() {
tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
}
}
},
1 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(2).enumerate() {
if i < tensor_data.len() {
let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = half::f16::from_bits(f16_val).to_f32();
}
}
},
_ => {
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(4).enumerate() {
if i < tensor_data.len() {
tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
}
}
}
}
weights.push(WeightMatrix::new(tensor_data, shape, tensor_name));
file.seek(SeekFrom::Start(current_pos))?;
}
Ok(weights)
}
fn load_onnx(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
let mut file = File::open(path)?;
let mut buffer = Vec::new();
file.read_to_end(&mut buffer)?;
if buffer.len() < 8 || &buffer[0..8] != b"\x08\x01\x12\x07onnx\x1d" {
return Err("Invalid ONNX file format".into());
}
let mut weights = Vec::new();
let mut pos = 0;
while pos < buffer.len() {
if pos + 16 < buffer.len() {
let potential_size = u64::from_le_bytes([
buffer[pos], buffer[pos+1], buffer[pos+2], buffer[pos+3],
buffer[pos+4], buffer[pos+5], buffer[pos+6], buffer[pos+7]
]);
if potential_size > 0 && potential_size < 100_000_000 { let tensor_size = potential_size as usize;
if pos + 16 + tensor_size * 4 <= buffer.len() {
let mut tensor_data = vec![0f32; tensor_size];
for i in 0..tensor_size {
let byte_pos = pos + 16 + i * 4;
if byte_pos + 3 < buffer.len() {
tensor_data[i] = f32::from_le_bytes([
buffer[byte_pos], buffer[byte_pos+1],
buffer[byte_pos+2], buffer[byte_pos+3]
]);
}
}
let dim = (tensor_size as f64).sqrt() as usize;
let shape = vec![dim, dim];
weights.push(WeightMatrix::new(
tensor_data,
shape,
format!("onnx_tensor_{}", weights.len())
));
pos += 16 + tensor_size * 4;
continue;
}
}
}
pos += 1;
}
if weights.is_empty() {
return Err("Could not extract tensors from ONNX file".into());
}
Ok(weights)
}
pub fn get_metadata(fetch_result: &FetchResult) -> Option<&ModelMetadata> {
fetch_result.metadata.as_ref()
}
pub fn validate_model(fetch_result: &FetchResult) -> Result<bool> {
match &fetch_result.model_format {
ModelFormat::SafeTensors => {
let mut file = File::open(&fetch_result.local_path)?;
let mut magic = [0u8; 15];
file.read_exact(&mut magic)?;
Ok(&magic == b"__safetensors__")
},
ModelFormat::PyTorch => {
let mut file = File::open(&fetch_result.local_path)?;
let mut header = [0u8; 8];
file.read_exact(&mut header)?;
Ok(true) },
ModelFormat::GGUF => {
let mut file = File::open(&fetch_result.local_path)?;
let mut magic = [0u8; 4];
file.read_exact(&mut magic)?;
Ok(&magic == b"GGUF")
},
ModelFormat::ONNX => {
let mut file = File::open(&fetch_result.local_path)?;
let mut header = [0u8; 9];
file.read_exact(&mut header)?;
Ok(&header == b"\x08\x01\x12\x07onnx\x1d")
},
ModelFormat::Unknown => Ok(false),
}
}
pub fn get_model_stats(weights: &[WeightMatrix]) -> ModelStats {
let total_parameters: usize = weights.iter()
.map(|w| w.data.len())
.sum();
let total_size_bytes = total_parameters * 4;
let largest_tensor = weights.iter()
.max_by_key(|w| w.data.len())
.map(|w| w.data.len())
.unwrap_or(0);
let smallest_tensor = weights.iter()
.min_by_key(|w| w.data.len())
.map(|w| w.data.len())
.unwrap_or(0);
ModelStats {
total_parameters,
total_size_mb: total_size_bytes as f64 / (1024.0 * 1024.0),
num_tensors: weights.len(),
largest_tensor,
smallest_tensor,
average_tensor_size: if weights.is_empty() { 0 } else { total_parameters / weights.len() },
}
}
}