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 fn bf16_to_f32(bf16_bits: u16) -> f32 {
24 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 if mantissa == 0 {
32 return 0.0;
33 } else {
34 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 if mantissa == 0 {
41 return if sign == 1 { f32::NEG_INFINITY } else { f32::INFINITY };
42 } else {
43 return f32::NAN;
44 }
45 } else {
46 let f32_exponent = (exponent as i32 - 127) + 127; 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 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 fn load_safetensors(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
68 let mut file = File::open(path)?;
69
70 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 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 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 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 file.seek(SeekFrom::Start(offset + start_offset))?;
128
129 let tensor_data = match dtype {
131 "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; }
155 }
156 data
157 },
158 "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 "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 "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; }
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; }
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 "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 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 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 let mut weights = Vec::new();
270
271 if buffer.len() > 4 && &buffer[0..4] == b"PK\x03\x04" {
274 return Err("ZIP-based PyTorch files require specialized parsing. Please convert to SafeTensors format for BF16 support.".into());
276 }
277
278 let mut pos = 0;
280 while pos < buffer.len() - 20 {
281 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 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 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 let dtype_hint = buffer[start_pos + 8];
317 let (bytes_per_element, dtype_name) = match dtype_hint {
318 1 => (2, "F16"), 2 => (2, "BF16"), 4 => (4, "F32"), _ => (4, "F32"), };
323
324 let tensor_bytes = tensor_elements * bytes_per_element;
325 let data_start = start_pos + 12; if data_start + tensor_bytes > buffer.len() {
328 return Ok(None);
329 }
330
331 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 let tensor_shape = if tensor_elements <= 1024 {
362 vec![tensor_elements] } else {
364 let dim = (tensor_elements as f64).sqrt() as usize;
366 if dim * dim == tensor_elements {
367 vec![dim, dim]
368 } else {
369 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 fn load_gguf(path: &PathBuf) -> Result<Vec<WeightMatrix>> {
401 let mut file = File::open(path)?;
402
403 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 file.seek(SeekFrom::Current(metadata_len as i64))?;
424
425 let mut weights = Vec::new();
426
427 for _i in 0..num_tensors {
429 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 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 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 let mut tensor_type = [0u8; 4];
453 file.read_exact(&mut tensor_type)?;
454 let dtype = u32::from_le_bytes(tensor_type);
455
456 let mut offset = [0u8; 8];
458 file.read_exact(&mut offset)?;
459 let tensor_offset = u64::from_le_bytes(offset);
460
461 let total_elements: usize = shape.iter().product();
463 let bytes_per_element = match dtype {
464 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, };
479 let tensor_size = total_elements * bytes_per_element;
480
481 let current_pos = file.stream_position()?;
483
484 file.seek(SeekFrom::Start(tensor_offset))?;
486
487 let mut tensor_data = vec![0f32; total_elements];
489 match dtype {
490 0 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 => { 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 file.seek(SeekFrom::Start(current_pos))?;
634 }
635
636 Ok(weights)
637 }
638
639 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 if buffer.len() < 8 {
647 return Err("ONNX file too small".into());
648 }
649
650 let mut weights = Vec::new();
653
654 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 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 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 let dtype_marker = buffer[start_pos + 8];
692 let (bytes_per_element, dtype_name) = match dtype_marker {
693 1 => (4, "F32"), 10 => (2, "F16"), 16 => (2, "BF16"), _ => {
697 if tensor_elements * 2 + start_pos + 16 < buffer.len() {
699 (2, "F16") } else if tensor_elements * 4 + start_pos + 16 < buffer.len() {
701 (4, "F32") } 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; if data_start + tensor_bytes > buffer.len() {
712 return Ok(None);
713 }
714
715 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 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 fn infer_tensor_shape(elements: usize) -> Vec<usize> {
754 if elements <= 1024 {
755 vec![elements] } else {
757 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 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 pub fn get_metadata(fetch_result: &FetchResult) -> Option<&ModelMetadata> {
776 fetch_result.metadata.as_ref()
777 }
778
779 pub fn validate_model(fetch_result: &FetchResult) -> Result<bool> {
781 match &fetch_result.model_format {
782 ModelFormat::SafeTensors => {
783 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 let mut file = File::open(&fetch_result.local_path)?;
792 let mut header = [0u8; 8];
793 file.read_exact(&mut header)?;
794 Ok(true) },
797 ModelFormat::GGUF => {
798 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 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 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; 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}