use candle_core::{DType, Device};
use mlmf::{loader::load_safetensors, LoadOptions};
fn main() -> Result<(), Box<dyn std::error::Error>> {
let device = Device::cuda_if_available(0).unwrap_or(Device::Cpu);
let dtype = match device {
Device::Cuda(_) => DType::F16, _ => DType::F32, };
println!("🚀 Loading LLaMA model...");
println!("📱 Device: {:?}", device);
println!("🔢 Data type: {:?}", dtype);
let options = LoadOptions::new(device, dtype)
.with_progress() .without_mmap();
let model_path = std::env::args()
.nth(1)
.unwrap_or_else(|| "./models/llama-7b".to_string());
let loaded = load_safetensors(&model_path, options)?;
println!("\n✅ Model loaded successfully!");
println!(
"🏗️ Architecture: {}",
loaded
.name_mapper
.architecture()
.map(|arch| arch.name())
.unwrap_or("Unknown")
);
println!("📊 Configuration: {}", loaded.config.summary());
println!("🧮 Total tensors: {}", loaded.raw_tensors.len());
println!("🗺️ Mapped tensors: {}", loaded.name_mapper.len());
println!("\n📝 Example tensor name mappings:");
let mut count = 0;
for (hf_name, mapped_name) in loaded.name_mapper.iter() {
if count >= 5 {
println!(" ... and {} more", loaded.name_mapper.len() - count);
break;
}
println!(" {} → {}", hf_name, mapped_name);
count += 1;
}
println!("\n🔍 Tensor inspection:");
if let Some(tensor) = loaded.get_tensor("wte.weight") {
println!(
" Token embeddings: {:?} {:?}",
tensor.dims(),
tensor.dtype()
);
}
if let Some(tensor) = loaded.get_tensor("h.0.attn.q_proj.weight") {
println!(
" Query projection (layer 0): {:?} {:?}",
tensor.dims(),
tensor.dtype()
);
}
let memory_estimate =
mlmf::validation::estimate_memory_usage(&loaded.config, dtype, Some(1), None);
println!("\n💾 {}", memory_estimate.summary());
Ok(())
}