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 {
fn bf16_to_f32(bf16_bits: u16) -> f32 {
let sign = (bf16_bits >> 15) & 0x1;
let exponent = (bf16_bits >> 7) & 0xFF;
let mantissa = bf16_bits & 0x7F;
if exponent == 0 {
if mantissa == 0 {
return 0.0;
} else {
let f32_mantissa = mantissa as f32 / 128.0;
return if sign == 1 { -f32_mantissa } else { f32_mantissa };
}
} else if exponent == 0xFF {
if mantissa == 0 {
return if sign == 1 { f32::NEG_INFINITY } else { f32::INFINITY };
} else {
return f32::NAN;
}
} else {
let f32_exponent = (exponent as i32 - 127) + 127; let f32_mantissa = mantissa as u32;
let f32_bits = (sign as u32) << 31 | (f32_exponent as u32) << 23 | f32_mantissa << 16;
return f32::from_bits(f32_bits);
}
}
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.get("dtype")
.and_then(|v| v.as_str())
.unwrap_or("F32");
let shape = tensor_obj.get("shape")
.and_then(|v| v.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.get("data_offsets")
.and_then(|v| v.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" | "FLOAT32" | "float32" => {
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
},
"F64" | "FLOAT64" | "float64" => {
let mut data = vec![0f32; tensor_size / 8];
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(8).enumerate() {
if i < data.len() {
let f64_val = f64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3],
chunk[4], chunk[5], chunk[6], chunk[7]
]);
data[i] = f64_val as f32; }
}
data
},
"F16" | "FLOAT16" | "float16" | "HALF" => {
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
},
"BF16" | "bfloat16" | "BFLOAT16" | "brain_float16" => {
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 bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
data[i] = Self::bf16_to_f32(bf16_val);
}
}
data
},
"I8" | "INT8" | "int8" => {
let mut data = vec![0f32; tensor_size];
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, &byte) in bytes.iter().enumerate() {
if i < data.len() {
data[i] = (byte as i8) as f32; }
}
data
},
"U8" | "UINT8" | "uint8" => {
let mut data = vec![0f32; tensor_size];
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, &byte) in bytes.iter().enumerate() {
if i < data.len() {
data[i] = byte as f32; }
}
data
},
"I16" | "INT16" | "int16" => {
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 i16_val = i16::from_le_bytes([chunk[0], chunk[1]]);
data[i] = i16_val as f32;
}
}
data
},
"I32" | "INT32" | "int32" => {
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() {
let i32_val = i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
data[i] = i32_val as f32;
}
}
data
},
"BOOL" | "bool" => {
let mut data = vec![0f32; tensor_size];
let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, &byte) in bytes.iter().enumerate() {
if i < data.len() {
data[i] = if byte != 0 { 1.0 } else { 0.0 };
}
}
data
},
_ => return Err(format!("Unsupported dtype: {} - NOVAQ supports F32, F64, F16, BF16, I8, U8, I16, I32, BOOL for universal LLM compatibility", 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();
if buffer.len() > 4 && &buffer[0..4] == b"PK\x03\x04" {
return Err("ZIP-based PyTorch files require specialized parsing. Please convert to SafeTensors format for BF16 support.".into());
}
let mut pos = 0;
while pos < buffer.len() - 20 {
if let Some((tensor_data, tensor_shape, tensor_name, new_pos)) = Self::try_parse_pytorch_tensor(&buffer, pos)? {
weights.push(WeightMatrix::new(tensor_data, tensor_shape, tensor_name));
pos = new_pos;
} else {
pos += 1;
}
}
if weights.is_empty() {
return Err("Could not extract tensors from PyTorch file. For BF16 models, consider converting to SafeTensors format.".into());
}
Ok(weights)
}
fn try_parse_pytorch_tensor(buffer: &[u8], start_pos: usize) -> Result<Option<(Vec<f32>, Vec<usize>, String, usize)>> {
if start_pos + 20 >= buffer.len() {
return Ok(None);
}
let potential_size = u64::from_le_bytes([
buffer[start_pos], buffer[start_pos+1], buffer[start_pos+2], buffer[start_pos+3],
buffer[start_pos+4], buffer[start_pos+5], buffer[start_pos+6], buffer[start_pos+7]
]);
if potential_size == 0 || potential_size > 1_000_000_000 {
return Ok(None);
}
let tensor_elements = potential_size as usize;
let dtype_hint = buffer[start_pos + 8];
let (bytes_per_element, dtype_name) = match dtype_hint {
1 => (2, "F16"), 2 => (2, "BF16"), 4 => (4, "F32"), _ => (4, "F32"), };
let tensor_bytes = tensor_elements * bytes_per_element;
let data_start = start_pos + 12;
if data_start + tensor_bytes > buffer.len() {
return Ok(None);
}
let mut tensor_data = vec![0f32; tensor_elements];
match dtype_name {
"F32" => {
for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(4).enumerate() {
if i < tensor_data.len() && chunk.len() >= 4 {
tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
}
}
},
"F16" => {
for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
if i < tensor_data.len() && chunk.len() >= 2 {
let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = half::f16::from_bits(f16_val).to_f32();
}
}
},
"BF16" => {
for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
if i < tensor_data.len() && chunk.len() >= 2 {
let bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = Self::bf16_to_f32(bf16_val);
}
}
},
_ => return Ok(None),
}
let tensor_shape = if tensor_elements <= 1024 {
vec![tensor_elements] } else {
let dim = (tensor_elements as f64).sqrt() as usize;
if dim * dim == tensor_elements {
vec![dim, dim]
} else {
let mut factors = Vec::new();
let mut n = tensor_elements;
let mut d = 2;
while d * d <= n {
if n % d == 0 {
factors.push(d);
n /= d;
} else {
d += 1;
}
}
if n > 1 {
factors.push(n);
}
if factors.len() >= 2 {
vec![factors[0] * factors[1], tensor_elements / (factors[0] * factors[1])]
} else {
vec![tensor_elements]
}
}
};
let tensor_name = format!("pytorch_tensor_{}", start_pos);
let next_pos = data_start + tensor_bytes;
Ok(Some((tensor_data, tensor_shape, tensor_name, next_pos)))
}
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, 2 => 2, 3 => 1, 4 => 1, 5 => 2, 6 => 2, 7 => 4, 8 => 4, 9 => 8, 10 => 8, 11 => 8, 12 => 1, _ => 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();
}
}
},
2 => { 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 bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = Self::bf16_to_f32(bf16_val);
}
}
},
3 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, &byte) in bytes.iter().enumerate() {
if i < tensor_data.len() {
tensor_data[i] = (byte as i8) as f32;
}
}
},
4 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, &byte) in bytes.iter().enumerate() {
if i < tensor_data.len() {
tensor_data[i] = byte as f32;
}
}
},
5 => { 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 i16_val = i16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = i16_val as f32;
}
}
},
6 => { 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 u16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = u16_val as f32;
}
}
},
7 => { 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() {
let i32_val = i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
tensor_data[i] = i32_val as f32;
}
}
},
8 => { 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() {
let u32_val = u32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
tensor_data[i] = u32_val as f32;
}
}
},
9 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(8).enumerate() {
if i < tensor_data.len() {
let f64_val = f64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3],
chunk[4], chunk[5], chunk[6], chunk[7]
]);
tensor_data[i] = f64_val as f32;
}
}
},
10 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(8).enumerate() {
if i < tensor_data.len() {
let i64_val = i64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3],
chunk[4], chunk[5], chunk[6], chunk[7]
]);
tensor_data[i] = i64_val as f32;
}
}
},
11 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, chunk) in bytes.chunks(8).enumerate() {
if i < tensor_data.len() {
let u64_val = u64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3],
chunk[4], chunk[5], chunk[6], chunk[7]
]);
tensor_data[i] = u64_val as f32;
}
}
},
12 => { let mut bytes = vec![0u8; tensor_size];
file.read_exact(&mut bytes)?;
for (i, &byte) in bytes.iter().enumerate() {
if i < tensor_data.len() {
tensor_data[i] = if byte != 0 { 1.0 } else { 0.0 };
}
}
},
_ => {
return Err(format!("Unsupported GGUF tensor dtype: {} - NOVAQ supports F32(0), F16(1), BF16(2), I8(3), U8(4), I16(5), U16(6), I32(7), U32(8), F64(9), I64(10), U64(11), BOOL(12)", dtype).into());
}
}
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 {
return Err("ONNX file too small".into());
}
let mut weights = Vec::new();
let mut pos = 0;
while pos < buffer.len() - 20 {
if let Some((tensor_data, tensor_shape, tensor_name, new_pos)) = Self::try_parse_onnx_tensor(&buffer, pos)? {
weights.push(WeightMatrix::new(tensor_data, tensor_shape, tensor_name));
pos = new_pos;
} else {
pos += 1;
}
}
if weights.is_empty() {
return Err("Could not extract tensors from ONNX file. For BF16 models, consider converting to SafeTensors format.".into());
}
Ok(weights)
}
fn try_parse_onnx_tensor(buffer: &[u8], start_pos: usize) -> Result<Option<(Vec<f32>, Vec<usize>, String, usize)>> {
if start_pos + 20 >= buffer.len() {
return Ok(None);
}
let potential_size = u64::from_le_bytes([
buffer[start_pos], buffer[start_pos+1], buffer[start_pos+2], buffer[start_pos+3],
buffer[start_pos+4], buffer[start_pos+5], buffer[start_pos+6], buffer[start_pos+7]
]);
if potential_size == 0 || potential_size > 100_000_000 {
return Ok(None);
}
let tensor_elements = potential_size as usize;
let dtype_marker = buffer[start_pos + 8];
let (bytes_per_element, dtype_name) = match dtype_marker {
1 => (4, "F32"), 10 => (2, "F16"), 16 => (2, "BF16"), _ => {
if tensor_elements * 2 + start_pos + 16 < buffer.len() {
(2, "F16") } else if tensor_elements * 4 + start_pos + 16 < buffer.len() {
(4, "F32") } else {
return Ok(None);
}
},
};
let tensor_bytes = tensor_elements * bytes_per_element;
let data_start = start_pos + 16;
if data_start + tensor_bytes > buffer.len() {
return Ok(None);
}
let mut tensor_data = vec![0f32; tensor_elements];
match dtype_name {
"F32" => {
for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(4).enumerate() {
if i < tensor_data.len() && chunk.len() >= 4 {
tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
}
}
},
"F16" => {
for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
if i < tensor_data.len() && chunk.len() >= 2 {
let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = half::f16::from_bits(f16_val).to_f32();
}
}
},
"BF16" => {
for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
if i < tensor_data.len() && chunk.len() >= 2 {
let bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
tensor_data[i] = Self::bf16_to_f32(bf16_val);
}
}
},
_ => return Ok(None),
}
let tensor_shape = Self::infer_tensor_shape(tensor_elements);
let tensor_name = format!("onnx_tensor_{}", start_pos);
let next_pos = data_start + tensor_bytes;
Ok(Some((tensor_data, tensor_shape, tensor_name, next_pos)))
}
fn infer_tensor_shape(elements: usize) -> Vec<usize> {
if elements <= 1024 {
vec![elements] } else {
let sqrt_elements = (elements as f64).sqrt() as usize;
if sqrt_elements * sqrt_elements == elements {
vec![sqrt_elements, sqrt_elements]
} else {
let mut best_factor = 1;
for i in 2..=((elements as f64).sqrt() as usize) {
if elements % i == 0 {
best_factor = i;
}
}
vec![best_factor, elements / best_factor]
}
}
}
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() },
}
}
}