Skip to main content

ohms_adaptq/
real_model_loader.rs

1use crate::{Result, WeightMatrix};
2use crate::model_fetcher::{FetchResult, ModelMetadata, ModelFormat};
3use std::path::PathBuf;
4use std::collections::HashMap;
5use serde::{Deserialize, Serialize};
6use std::fs::File;
7use std::io::{Read, Seek, SeekFrom};
8
9#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct ModelStats {
11    pub total_parameters: usize,
12    pub total_size_mb: f64,
13    pub num_tensors: usize,
14    pub largest_tensor: usize,
15    pub smallest_tensor: usize,
16    pub average_tensor_size: usize,
17}
18
19pub struct RealModelLoader;
20
21impl RealModelLoader {
22    /// Convert BF16 (Brain Float 16) to F32
23    fn bf16_to_f32(bf16_bits: u16) -> f32 {
24        // BF16 has 1 sign bit, 8 exponent bits, 7 mantissa bits
25        let sign = (bf16_bits >> 15) & 0x1;
26        let exponent = (bf16_bits >> 7) & 0xFF;
27        let mantissa = bf16_bits & 0x7F;
28        
29        if exponent == 0 {
30            // Zero or denormalized
31            if mantissa == 0 {
32                return 0.0;
33            } else {
34                // Denormalized - very small numbers
35                let f32_mantissa = mantissa as f32 / 128.0;
36                return if sign == 1 { -f32_mantissa } else { f32_mantissa };
37            }
38        } else if exponent == 0xFF {
39            // Infinity or NaN
40            if mantissa == 0 {
41                return if sign == 1 { f32::NEG_INFINITY } else { f32::INFINITY };
42            } else {
43                return f32::NAN;
44            }
45        } else {
46            // Normalized number
47            let f32_exponent = (exponent as i32 - 127) + 127; // Adjust bias
48            let f32_mantissa = mantissa as u32;
49            
50            let f32_bits = (sign as u32) << 31 | (f32_exponent as u32) << 23 | f32_mantissa << 16;
51            return f32::from_bits(f32_bits);
52        }
53    }
54
55    /// Load model from fetch result and convert to NOVAQ-compatible weights
56    pub fn load_model(fetch_result: &FetchResult) -> Result<Vec<WeightMatrix>> {
57        match &fetch_result.model_format {
58            ModelFormat::SafeTensors => Self::load_safetensors(&fetch_result.local_path),
59            ModelFormat::PyTorch => Self::load_pytorch(&fetch_result.local_path),
60            ModelFormat::GGUF => Self::load_gguf(&fetch_result.local_path),
61            ModelFormat::ONNX => Self::load_onnx(&fetch_result.local_path),
62            ModelFormat::Unknown => Err("Unknown model format".into()),
63        }
64    }
65
66    /// Load SafeTensors format (most common for modern models)
67    fn load_safetensors(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
68        let mut file = File::open(path)?;
69        
70        // Read header length (8 bytes)
71        let mut header_len_bytes = [0u8; 8];
72        file.read_exact(&mut header_len_bytes)?;
73        let header_len = u64::from_le_bytes(header_len_bytes) as usize;
74        
75        // Read JSON header
76        let mut header_json = vec![0u8; header_len];
77        file.read_exact(&mut header_json)?;
78        let header_str = String::from_utf8(header_json)?;
79        let header: HashMap<String, serde_json::Value> = serde_json::from_str(&header_str)?;
80        
81        let mut weights = Vec::new();
82        let offset = 8 + header_len as u64;
83        
84        for (tensor_name, tensor_info) in header {
85            // Skip SafeTensors metadata entries - they contain format info, not tensors
86            if tensor_name == "__metadata__" {
87                continue;
88            }
89            
90            if let Some(tensor_obj) = tensor_info.as_object() {
91                let dtype = tensor_obj.get("dtype")
92                    .and_then(|v| v.as_str())
93                    .unwrap_or("F32");
94                let shape = if let Some(shape_array) = tensor_obj.get("shape").and_then(|v| v.as_array()) {
95                    shape_array.iter()
96                        .filter_map(|v| {
97                            // Try as u64 first, then as i64, then as f64
98                            v.as_u64()
99                                .map(|n| n as usize)
100                                .or_else(|| v.as_i64().map(|n| n as usize))
101                                .or_else(|| v.as_f64().map(|n| n as usize))
102                        })
103                        .collect::<Vec<_>>()
104                } else {
105                    return Err(format!("Invalid shape for tensor '{}': shape field missing or not an array", tensor_name).into());
106                };
107                
108                if shape.is_empty() {
109                    return Err(format!("Empty shape for tensor '{}'", tensor_name).into());
110                }
111                
112                let data_offsets = tensor_obj.get("data_offsets")
113                    .and_then(|v| v.as_array())
114                    .ok_or_else(|| format!("Invalid data_offsets for tensor '{}': field missing or not an array", tensor_name))?;
115                
116                if data_offsets.len() != 2 {
117                    return Err(format!("Invalid data_offsets for tensor '{}': expected 2 elements, got {}", tensor_name, data_offsets.len()).into());
118                }
119                
120                let start_offset = data_offsets[0].as_u64()
121                    .ok_or_else(|| format!("Invalid start offset for tensor '{}': not a valid number", tensor_name))? as u64;
122                let end_offset = data_offsets[1].as_u64()
123                    .ok_or_else(|| format!("Invalid end offset for tensor '{}': not a valid number", tensor_name))? as u64;
124                let tensor_size = (end_offset - start_offset) as usize;
125                
126                // Seek to tensor data
127                file.seek(SeekFrom::Start(offset + start_offset))?;
128                
129                // Read tensor data based on dtype - Universal LLM model compatibility
130                let tensor_data = match dtype {
131                    // Standard floating point formats
132                    "F32" | "FLOAT32" | "float32" => {
133                        let mut data = vec![0f32; tensor_size / 4];
134                        let mut bytes = vec![0u8; tensor_size];
135                        file.read_exact(&mut bytes)?;
136                        for (i, chunk) in bytes.chunks(4).enumerate() {
137                            if i < data.len() {
138                                data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
139                            }
140                        }
141                        data
142                    },
143                    "F64" | "FLOAT64" | "float64" => {
144                        let mut data = vec![0f32; tensor_size / 8];
145                        let mut bytes = vec![0u8; tensor_size];
146                        file.read_exact(&mut bytes)?;
147                        for (i, chunk) in bytes.chunks(8).enumerate() {
148                            if i < data.len() {
149                                let f64_val = f64::from_le_bytes([
150                                    chunk[0], chunk[1], chunk[2], chunk[3],
151                                    chunk[4], chunk[5], chunk[6], chunk[7]
152                                ]);
153                                data[i] = f64_val as f32; // Convert to f32
154                            }
155                        }
156                        data
157                    },
158                    // Half precision formats
159                    "F16" | "FLOAT16" | "float16" | "HALF" => {
160                        let mut data = vec![0f32; tensor_size / 2];
161                        let mut bytes = vec![0u8; tensor_size];
162                        file.read_exact(&mut bytes)?;
163                        for (i, chunk) in bytes.chunks(2).enumerate() {
164                            if i < data.len() {
165                                let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
166                                data[i] = half::f16::from_bits(f16_val).to_f32();
167                            }
168                        }
169                        data
170                    },
171                    // Brain Float 16 - Critical for Phi-3, GPT-4, Claude models
172                    "BF16" | "bfloat16" | "BFLOAT16" | "brain_float16" => {
173                        let mut data = vec![0f32; tensor_size / 2];
174                        let mut bytes = vec![0u8; tensor_size];
175                        file.read_exact(&mut bytes)?;
176                        for (i, chunk) in bytes.chunks(2).enumerate() {
177                            if i < data.len() {
178                                let bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
179                                data[i] = Self::bf16_to_f32(bf16_val);
180                            }
181                        }
182                        data
183                    },
184                    // Integer formats - for quantized models (GPTQ, AWQ, etc.)
185                    "I8" | "INT8" | "int8" => {
186                        let mut data = vec![0f32; tensor_size];
187                        let mut bytes = vec![0u8; tensor_size];
188                        file.read_exact(&mut bytes)?;
189                        for (i, &byte) in bytes.iter().enumerate() {
190                            if i < data.len() {
191                                data[i] = (byte as i8) as f32; // Convert signed int8 to float
192                            }
193                        }
194                        data
195                    },
196                    "U8" | "UINT8" | "uint8" => {
197                        let mut data = vec![0f32; tensor_size];
198                        let mut bytes = vec![0u8; tensor_size];
199                        file.read_exact(&mut bytes)?;
200                        for (i, &byte) in bytes.iter().enumerate() {
201                            if i < data.len() {
202                                data[i] = byte as f32; // Convert unsigned int8 to float
203                            }
204                        }
205                        data
206                    },
207                    "I16" | "INT16" | "int16" => {
208                        let mut data = vec![0f32; tensor_size / 2];
209                        let mut bytes = vec![0u8; tensor_size];
210                        file.read_exact(&mut bytes)?;
211                        for (i, chunk) in bytes.chunks(2).enumerate() {
212                            if i < data.len() {
213                                let i16_val = i16::from_le_bytes([chunk[0], chunk[1]]);
214                                data[i] = i16_val as f32;
215                            }
216                        }
217                        data
218                    },
219                    "I32" | "INT32" | "int32" => {
220                        let mut data = vec![0f32; tensor_size / 4];
221                        let mut bytes = vec![0u8; tensor_size];
222                        file.read_exact(&mut bytes)?;
223                        for (i, chunk) in bytes.chunks(4).enumerate() {
224                            if i < data.len() {
225                                let i32_val = i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
226                                data[i] = i32_val as f32;
227                            }
228                        }
229                        data
230                    },
231                    // Boolean for mask tensors
232                    "BOOL" | "bool" => {
233                        let mut data = vec![0f32; tensor_size];
234                        let mut bytes = vec![0u8; tensor_size];
235                        file.read_exact(&mut bytes)?;
236                        for (i, &byte) in bytes.iter().enumerate() {
237                            if i < data.len() {
238                                data[i] = if byte != 0 { 1.0 } else { 0.0 };
239                            }
240                        }
241                        data
242                    },
243                    _ => return Err(format!("Unsupported dtype: {} - NOVAQ supports F32, F64, F16, BF16, I8, U8, I16, I32, BOOL for universal LLM compatibility", dtype).into()),
244                };
245                
246                // Reshape tensor data
247                let total_elements: usize = shape.iter().product();
248                if tensor_data.len() != total_elements {
249                    return Err(format!("Tensor size mismatch for {}: expected {}, got {}", 
250                        tensor_name, total_elements, tensor_data.len()).into());
251                }
252                
253                weights.push(WeightMatrix::new(tensor_data, shape, tensor_name));
254            }
255        }
256        
257        Ok(weights)
258    }
259
260    /// Load PyTorch format (.bin, .pt, .pth files)
261    fn load_pytorch(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
262        let mut file = File::open(path)?;
263        let mut buffer = Vec::new();
264        file.read_to_end(&mut buffer)?;
265        
266        // Try to parse PyTorch format
267        // Note: This is a simplified implementation. Full PyTorch support would require
268        // proper pickle protocol parsing or Python interop
269        let mut weights = Vec::new();
270        
271        // Look for common PyTorch patterns
272        // Check for ZIP archive structure (newer PyTorch format)
273        if buffer.len() > 4 && &buffer[0..4] == b"PK\x03\x04" {
274            // This is a ZIP-based PyTorch file (modern format)
275            return Err("ZIP-based PyTorch files require specialized parsing. Please convert to SafeTensors format for BF16 support.".into());
276        }
277        
278        // Look for tensor data patterns in binary PyTorch files
279        let mut pos = 0;
280        while pos < buffer.len() - 20 {
281            // Look for potential tensor headers with data type information
282            if let Some((tensor_data, tensor_shape, tensor_name, new_pos)) = Self::try_parse_pytorch_tensor(&buffer, pos)? {
283                weights.push(WeightMatrix::new(tensor_data, tensor_shape, tensor_name));
284                pos = new_pos;
285            } else {
286                pos += 1;
287            }
288        }
289        
290        if weights.is_empty() {
291            return Err("Could not extract tensors from PyTorch file. For BF16 models, consider converting to SafeTensors format.".into());
292        }
293        
294        Ok(weights)
295    }
296
297    /// Try to parse a single tensor from PyTorch binary data
298    fn try_parse_pytorch_tensor(buffer: &[u8], start_pos: usize) -> Result<Option<(Vec<f32>, Vec<usize>, String, usize)>> {
299        if start_pos + 20 >= buffer.len() {
300            return Ok(None);
301        }
302        
303        // Look for tensor markers (this is heuristic-based)
304        let potential_size = u64::from_le_bytes([
305            buffer[start_pos], buffer[start_pos+1], buffer[start_pos+2], buffer[start_pos+3],
306            buffer[start_pos+4], buffer[start_pos+5], buffer[start_pos+6], buffer[start_pos+7]
307        ]);
308        
309        if potential_size == 0 || potential_size > 1_000_000_000 {
310            return Ok(None);
311        }
312        
313        let tensor_elements = potential_size as usize;
314        
315        // Try to detect data type from the next bytes (heuristic)
316        let dtype_hint = buffer[start_pos + 8];
317        let (bytes_per_element, dtype_name) = match dtype_hint {
318            1 => (2, "F16"),    // F16 hint
319            2 => (2, "BF16"),   // BF16 hint
320            4 => (4, "F32"),    // F32 hint
321            _ => (4, "F32"),    // Default to F32
322        };
323        
324        let tensor_bytes = tensor_elements * bytes_per_element;
325        let data_start = start_pos + 12; // Skip header
326        
327        if data_start + tensor_bytes > buffer.len() {
328            return Ok(None);
329        }
330        
331        // Read tensor data based on detected type
332        let mut tensor_data = vec![0f32; tensor_elements];
333        match dtype_name {
334            "F32" => {
335                for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(4).enumerate() {
336                    if i < tensor_data.len() && chunk.len() >= 4 {
337                        tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
338                    }
339                }
340            },
341            "F16" => {
342                for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
343                    if i < tensor_data.len() && chunk.len() >= 2 {
344                        let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
345                        tensor_data[i] = half::f16::from_bits(f16_val).to_f32();
346                    }
347                }
348            },
349            "BF16" => {
350                for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
351                    if i < tensor_data.len() && chunk.len() >= 2 {
352                        let bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
353                        tensor_data[i] = Self::bf16_to_f32(bf16_val);
354                    }
355                }
356            },
357            _ => return Ok(None),
358        }
359        
360        // Create a reasonable shape (simplified)
361        let tensor_shape = if tensor_elements <= 1024 {
362            vec![tensor_elements] // 1D for small tensors
363        } else {
364            // Try to create a 2D shape
365            let dim = (tensor_elements as f64).sqrt() as usize;
366            if dim * dim == tensor_elements {
367                vec![dim, dim]
368            } else {
369                // Find factors
370                let mut factors = Vec::new();
371                let mut n = tensor_elements;
372                let mut d = 2;
373                while d * d <= n {
374                    if n % d == 0 {
375                        factors.push(d);
376                        n /= d;
377                    } else {
378                        d += 1;
379                    }
380                }
381                if n > 1 {
382                    factors.push(n);
383                }
384                
385                if factors.len() >= 2 {
386                    vec![factors[0] * factors[1], tensor_elements / (factors[0] * factors[1])]
387                } else {
388                    vec![tensor_elements]
389                }
390            }
391        };
392        
393        let tensor_name = format!("pytorch_tensor_{}", start_pos);
394        let next_pos = data_start + tensor_bytes;
395        
396        Ok(Some((tensor_data, tensor_shape, tensor_name, next_pos)))
397    }
398
399    /// Load GGUF format (Ollama models)
400    fn load_gguf(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
401        let mut file = File::open(path)?;
402        
403        // Read GGUF header
404        let mut magic = [0u8; 4];
405        file.read_exact(&mut magic)?;
406        if &magic != b"GGUF" {
407            return Err("Invalid GGUF magic number".into());
408        }
409        
410        let mut version = [0u8; 4];
411        file.read_exact(&mut version)?;
412        let _version_num = u32::from_le_bytes(version);
413        
414        let mut tensor_count = [0u8; 8];
415        file.read_exact(&mut tensor_count)?;
416        let num_tensors = u64::from_le_bytes(tensor_count);
417        
418        let mut metadata_size = [0u8; 8];
419        file.read_exact(&mut metadata_size)?;
420        let metadata_len = u64::from_le_bytes(metadata_size) as usize;
421        
422        // Skip metadata for now (we'll implement proper parsing later)
423        file.seek(SeekFrom::Current(metadata_len as i64))?;
424        
425        let mut weights = Vec::new();
426        
427        // Read tensors
428        for _i in 0..num_tensors {
429            // Read tensor name length
430            let mut name_len = [0u8; 4];
431            file.read_exact(&mut name_len)?;
432            let name_length = u32::from_le_bytes(name_len) as usize;
433            
434            // Read tensor name
435            let mut name_bytes = vec![0u8; name_length];
436            file.read_exact(&mut name_bytes)?;
437            let tensor_name = String::from_utf8(name_bytes)?;
438            
439            // Read tensor dimensions
440            let mut dims = [0u8; 4];
441            file.read_exact(&mut dims)?;
442            let num_dims = u32::from_le_bytes(dims) as usize;
443            
444            let mut shape = Vec::new();
445            for _ in 0..num_dims {
446                let mut dim = [0u8; 8];
447                file.read_exact(&mut dim)?;
448                shape.push(u64::from_le_bytes(dim) as usize);
449            }
450            
451            // Read tensor type
452            let mut tensor_type = [0u8; 4];
453            file.read_exact(&mut tensor_type)?;
454            let dtype = u32::from_le_bytes(tensor_type);
455            
456            // Read tensor offset
457            let mut offset = [0u8; 8];
458            file.read_exact(&mut offset)?;
459            let tensor_offset = u64::from_le_bytes(offset);
460            
461            // Calculate tensor size - Universal GGUF data type support
462            let total_elements: usize = shape.iter().product();
463            let bytes_per_element = match dtype {
464                0 => 4,  // F32
465                1 => 2,  // F16
466                2 => 2,  // BF16 
467                3 => 1,  // I8
468                4 => 1,  // U8
469                5 => 2,  // I16
470                6 => 2,  // U16
471                7 => 4,  // I32
472                8 => 4,  // U32
473                9 => 8,  // F64
474                10 => 8, // I64
475                11 => 8, // U64
476                12 => 1, // BOOL
477                _ => 4,  // Default to F32
478            };
479            let tensor_size = total_elements * bytes_per_element;
480            
481            // Store current position
482            let current_pos = file.stream_position()?;
483            
484            // Seek to tensor data
485            file.seek(SeekFrom::Start(tensor_offset))?;
486            
487            // Read tensor data - Universal GGUF data type support
488            let mut tensor_data = vec![0f32; total_elements];
489            match dtype {
490                0 => { // F32
491                    let mut bytes = vec![0u8; tensor_size];
492                    file.read_exact(&mut bytes)?;
493                    for (i, chunk) in bytes.chunks(4).enumerate() {
494                        if i < tensor_data.len() {
495                            tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
496                        }
497                    }
498                },
499                1 => { // F16
500                    let mut bytes = vec![0u8; tensor_size];
501                    file.read_exact(&mut bytes)?;
502                    for (i, chunk) in bytes.chunks(2).enumerate() {
503                        if i < tensor_data.len() {
504                            let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
505                            tensor_data[i] = half::f16::from_bits(f16_val).to_f32();
506                        }
507                    }
508                },
509                2 => { // BF16
510                    let mut bytes = vec![0u8; tensor_size];
511                    file.read_exact(&mut bytes)?;
512                    for (i, chunk) in bytes.chunks(2).enumerate() {
513                        if i < tensor_data.len() {
514                            let bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
515                            tensor_data[i] = Self::bf16_to_f32(bf16_val);
516                        }
517                    }
518                },
519                3 => { // I8
520                    let mut bytes = vec![0u8; tensor_size];
521                    file.read_exact(&mut bytes)?;
522                    for (i, &byte) in bytes.iter().enumerate() {
523                        if i < tensor_data.len() {
524                            tensor_data[i] = (byte as i8) as f32;
525                        }
526                    }
527                },
528                4 => { // U8
529                    let mut bytes = vec![0u8; tensor_size];
530                    file.read_exact(&mut bytes)?;
531                    for (i, &byte) in bytes.iter().enumerate() {
532                        if i < tensor_data.len() {
533                            tensor_data[i] = byte as f32;
534                        }
535                    }
536                },
537                5 => { // I16
538                    let mut bytes = vec![0u8; tensor_size];
539                    file.read_exact(&mut bytes)?;
540                    for (i, chunk) in bytes.chunks(2).enumerate() {
541                        if i < tensor_data.len() {
542                            let i16_val = i16::from_le_bytes([chunk[0], chunk[1]]);
543                            tensor_data[i] = i16_val as f32;
544                        }
545                    }
546                },
547                6 => { // U16
548                    let mut bytes = vec![0u8; tensor_size];
549                    file.read_exact(&mut bytes)?;
550                    for (i, chunk) in bytes.chunks(2).enumerate() {
551                        if i < tensor_data.len() {
552                            let u16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
553                            tensor_data[i] = u16_val as f32;
554                        }
555                    }
556                },
557                7 => { // I32
558                    let mut bytes = vec![0u8; tensor_size];
559                    file.read_exact(&mut bytes)?;
560                    for (i, chunk) in bytes.chunks(4).enumerate() {
561                        if i < tensor_data.len() {
562                            let i32_val = i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
563                            tensor_data[i] = i32_val as f32;
564                        }
565                    }
566                },
567                8 => { // U32
568                    let mut bytes = vec![0u8; tensor_size];
569                    file.read_exact(&mut bytes)?;
570                    for (i, chunk) in bytes.chunks(4).enumerate() {
571                        if i < tensor_data.len() {
572                            let u32_val = u32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
573                            tensor_data[i] = u32_val as f32;
574                        }
575                    }
576                },
577                9 => { // F64
578                    let mut bytes = vec![0u8; tensor_size];
579                    file.read_exact(&mut bytes)?;
580                    for (i, chunk) in bytes.chunks(8).enumerate() {
581                        if i < tensor_data.len() {
582                            let f64_val = f64::from_le_bytes([
583                                chunk[0], chunk[1], chunk[2], chunk[3],
584                                chunk[4], chunk[5], chunk[6], chunk[7]
585                            ]);
586                            tensor_data[i] = f64_val as f32;
587                        }
588                    }
589                },
590                10 => { // I64
591                    let mut bytes = vec![0u8; tensor_size];
592                    file.read_exact(&mut bytes)?;
593                    for (i, chunk) in bytes.chunks(8).enumerate() {
594                        if i < tensor_data.len() {
595                            let i64_val = i64::from_le_bytes([
596                                chunk[0], chunk[1], chunk[2], chunk[3],
597                                chunk[4], chunk[5], chunk[6], chunk[7]
598                            ]);
599                            tensor_data[i] = i64_val as f32;
600                        }
601                    }
602                },
603                11 => { // U64
604                    let mut bytes = vec![0u8; tensor_size];
605                    file.read_exact(&mut bytes)?;
606                    for (i, chunk) in bytes.chunks(8).enumerate() {
607                        if i < tensor_data.len() {
608                            let u64_val = u64::from_le_bytes([
609                                chunk[0], chunk[1], chunk[2], chunk[3],
610                                chunk[4], chunk[5], chunk[6], chunk[7]
611                            ]);
612                            tensor_data[i] = u64_val as f32;
613                        }
614                    }
615                },
616                12 => { // BOOL
617                    let mut bytes = vec![0u8; tensor_size];
618                    file.read_exact(&mut bytes)?;
619                    for (i, &byte) in bytes.iter().enumerate() {
620                        if i < tensor_data.len() {
621                            tensor_data[i] = if byte != 0 { 1.0 } else { 0.0 };
622                        }
623                    }
624                },
625                _ => {
626                    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());
627                }
628            }
629            
630            weights.push(WeightMatrix::new(tensor_data, shape, tensor_name));
631            
632            // Return to metadata position
633            file.seek(SeekFrom::Start(current_pos))?;
634        }
635        
636        Ok(weights)
637    }
638
639    /// Load ONNX format
640    fn load_onnx(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
641        let mut file = File::open(path)?;
642        let mut buffer = Vec::new();
643        file.read_to_end(&mut buffer)?;
644        
645        // Check ONNX magic number (simplified check)
646        if buffer.len() < 8 {
647            return Err("ONNX file too small".into());
648        }
649        
650        // Note: This is a simplified ONNX parser for demonstration
651        // Real implementation would use proper ONNX protobuf parsing
652        let mut weights = Vec::new();
653        
654        // Look for weight tensors in the ONNX file with BF16 support
655        let mut pos = 0;
656        while pos < buffer.len() - 20 {
657            if let Some((tensor_data, tensor_shape, tensor_name, new_pos)) = Self::try_parse_onnx_tensor(&buffer, pos)? {
658                weights.push(WeightMatrix::new(tensor_data, tensor_shape, tensor_name));
659                pos = new_pos;
660            } else {
661                pos += 1;
662            }
663        }
664        
665        if weights.is_empty() {
666            return Err("Could not extract tensors from ONNX file. For BF16 models, consider converting to SafeTensors format.".into());
667        }
668        
669        Ok(weights)
670    }
671
672    /// Try to parse a single tensor from ONNX binary data
673    fn try_parse_onnx_tensor(buffer: &[u8], start_pos: usize) -> Result<Option<(Vec<f32>, Vec<usize>, String, usize)>> {
674        if start_pos + 20 >= buffer.len() {
675            return Ok(None);
676        }
677        
678        // Look for potential ONNX tensor patterns
679        let potential_size = u64::from_le_bytes([
680            buffer[start_pos], buffer[start_pos+1], buffer[start_pos+2], buffer[start_pos+3],
681            buffer[start_pos+4], buffer[start_pos+5], buffer[start_pos+6], buffer[start_pos+7]
682        ]);
683        
684        if potential_size == 0 || potential_size > 100_000_000 {
685            return Ok(None);
686        }
687        
688        let tensor_elements = potential_size as usize;
689        
690        // Try to detect ONNX data type from header (simplified)
691        let dtype_marker = buffer[start_pos + 8];
692        let (bytes_per_element, dtype_name) = match dtype_marker {
693            1 => (4, "F32"),    // ONNX FLOAT type
694            10 => (2, "F16"),   // ONNX FLOAT16 type
695            16 => (2, "BF16"),  // ONNX BFLOAT16 type (if present)
696            _ => {
697                // Try to infer from data patterns
698                if tensor_elements * 2 + start_pos + 16 < buffer.len() {
699                    (2, "F16") // Assume F16
700                } else if tensor_elements * 4 + start_pos + 16 < buffer.len() {
701                    (4, "F32") // Assume F32
702                } else {
703                    return Ok(None);
704                }
705            },
706        };
707        
708        let tensor_bytes = tensor_elements * bytes_per_element;
709        let data_start = start_pos + 16; // Skip ONNX header
710        
711        if data_start + tensor_bytes > buffer.len() {
712            return Ok(None);
713        }
714        
715        // Read tensor data based on detected type
716        let mut tensor_data = vec![0f32; tensor_elements];
717        match dtype_name {
718            "F32" => {
719                for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(4).enumerate() {
720                    if i < tensor_data.len() && chunk.len() >= 4 {
721                        tensor_data[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
722                    }
723                }
724            },
725            "F16" => {
726                for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
727                    if i < tensor_data.len() && chunk.len() >= 2 {
728                        let f16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
729                        tensor_data[i] = half::f16::from_bits(f16_val).to_f32();
730                    }
731                }
732            },
733            "BF16" => {
734                for (i, chunk) in buffer[data_start..data_start + tensor_bytes].chunks(2).enumerate() {
735                    if i < tensor_data.len() && chunk.len() >= 2 {
736                        let bf16_val = u16::from_le_bytes([chunk[0], chunk[1]]);
737                        tensor_data[i] = Self::bf16_to_f32(bf16_val);
738                    }
739                }
740            },
741            _ => return Ok(None),
742        }
743        
744        // Create tensor shape (simplified approach for ONNX)
745        let tensor_shape = Self::infer_tensor_shape(tensor_elements);
746        let tensor_name = format!("onnx_tensor_{}", start_pos);
747        let next_pos = data_start + tensor_bytes;
748        
749        Ok(Some((tensor_data, tensor_shape, tensor_name, next_pos)))
750    }
751
752    /// Infer reasonable tensor shape from number of elements
753    fn infer_tensor_shape(elements: usize) -> Vec<usize> {
754        if elements <= 1024 {
755            vec![elements] // 1D for small tensors
756        } else {
757            // Try to create a reasonable 2D shape
758            let sqrt_elements = (elements as f64).sqrt() as usize;
759            if sqrt_elements * sqrt_elements == elements {
760                vec![sqrt_elements, sqrt_elements]
761            } else {
762                // Find factors for better shape
763                let mut best_factor = 1;
764                for i in 2..=((elements as f64).sqrt() as usize) {
765                    if elements % i == 0 {
766                        best_factor = i;
767                    }
768                }
769                vec![best_factor, elements / best_factor]
770            }
771        }
772    }
773
774    /// Get model metadata from fetch result
775    pub fn get_metadata(fetch_result: &FetchResult) -> Option<&ModelMetadata> {
776        fetch_result.metadata.as_ref()
777    }
778
779    /// Validate model format and compatibility
780    pub fn validate_model(fetch_result: &FetchResult) -> Result<bool> {
781        match &fetch_result.model_format {
782            ModelFormat::SafeTensors => {
783                // Validate SafeTensors header
784                let mut file = File::open(&fetch_result.local_path)?;
785                let mut magic = [0u8; 15];
786                file.read_exact(&mut magic)?;
787                Ok(&magic == b"__safetensors__")
788            },
789            ModelFormat::PyTorch => {
790                // Basic PyTorch validation
791                let mut file = File::open(&fetch_result.local_path)?;
792                let mut header = [0u8; 8];
793                file.read_exact(&mut header)?;
794                // Check for PyTorch magic (simplified)
795                Ok(true) // Assume valid for now
796            },
797            ModelFormat::GGUF => {
798                // Validate GGUF header
799                let mut file = File::open(&fetch_result.local_path)?;
800                let mut magic = [0u8; 4];
801                file.read_exact(&mut magic)?;
802                Ok(&magic == b"GGUF")
803            },
804            ModelFormat::ONNX => {
805                // Validate ONNX header
806                let mut file = File::open(&fetch_result.local_path)?;
807                let mut header = [0u8; 9];
808                file.read_exact(&mut header)?;
809                Ok(&header == b"\x08\x01\x12\x07onnx\x1d")
810            },
811            ModelFormat::Unknown => Ok(false),
812        }
813    }
814
815    /// Get model statistics
816    pub fn get_model_stats(weights: &[WeightMatrix]) -> ModelStats {
817        let total_parameters: usize = weights.iter()
818            .map(|w| w.data.len())
819            .sum();
820        
821        let total_size_bytes = total_parameters * 4; // f32 = 4 bytes
822        
823        let largest_tensor = weights.iter()
824            .max_by_key(|w| w.data.len())
825            .map(|w| w.data.len())
826            .unwrap_or(0);
827        
828        let smallest_tensor = weights.iter()
829            .min_by_key(|w| w.data.len())
830            .map(|w| w.data.len())
831            .unwrap_or(0);
832        
833        ModelStats {
834            total_parameters,
835            total_size_mb: total_size_bytes as f64 / (1024.0 * 1024.0),
836            num_tensors: weights.len(),
837            largest_tensor,
838            smallest_tensor,
839            average_tensor_size: if weights.is_empty() { 0 } else { total_parameters / weights.len() },
840        }
841    }
842}