| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| use candle::{DType, Device, Tensor, Module}; |
| use std::io::Write; |
|
|
| fn main() -> candle::Result<()> { |
| let model_id = std::env::args() |
| .find(|arg| arg.starts_with("--model=")) |
| .map(|s| s[8..].to_string()) |
| .unwrap_or_else(|| "igorls/gemma-4-12B-it-heretic-v1".to_string()); |
|
|
| let prompt = std::env::args() |
| .find(|arg| arg.starts_with("--prompt=")) |
| .map(|s| s[9..].to_string()) |
| .unwrap_or_else(|| "What is the meaning of life?".to_string()); |
|
|
| let n_tokens: usize = std::env::args() |
| .find(|arg| arg.starts_with("--n-tokens=")) |
| .and_then(|s| s[11..].parse().ok()) |
| .unwrap_or(128); |
|
|
| println!("NVFP4 Inference Example"); |
| println!("Model: {model_id}"); |
| println!("Prompt: {prompt}"); |
| println!("Max tokens: {n_tokens}"); |
| println!(); |
|
|
| |
| let device = Device::cuda_if_available(0)?; |
| println!("Device: {:?}", device); |
|
|
| |
| println!("\n=== NVFP4 Quantization Test ==="); |
| let test_data = Tensor::randn(0f32, 1f32, (256, 512), &device)?; |
| println!("Input shape: {:?}", test_data.shape()); |
|
|
| |
| use candle::quantized::{GgmlDType, QTensor}; |
| let qtensor = QTensor::quantize(&test_data, GgmlDType::NVFP4)?; |
| println!("Quantized dtype: {:?}", qtensor.dtype()); |
| println!("Quantized shape: {:?}", qtensor.shape()); |
|
|
| |
| let dequant = qtensor.dequantize(&device)?; |
| println!("Dequantized shape: {:?}", dequant.shape()); |
|
|
| |
| let diff = test_data.broadcast_sub(&dequant)?; |
| let max_err = diff.max_all()?.to_scalar::<f32>()?; |
| let mean_err = diff.mean_all()?.to_scalar::<f32>()?; |
| let ref_mean = test_data.abs()?.mean_all()?.to_scalar::<f32>()?; |
| println!("Max error: {:.6}", max_err); |
| println!("Mean error: {:.6}", mean_err); |
| println!("Relative error: {:.6}", mean_err / ref_mean); |
|
|
| |
| println!("\n=== NVFP4 MatMul Test ==="); |
| let x = Tensor::randn(0f32, 1f32, (1, 512), &device)?; |
| let w = Tensor::randn(0f32, 0.5f32, (256, 512), &device)?; |
| let w_q = QTensor::quantize(&w, GgmlDType::NVFP4)?; |
|
|
| |
| let ref_out = x.matmul(&w.t()?)?; |
| |
| |
| use candle::quantized::QMatMul; |
| let qmatmul = QMatMul::from_arc(std::sync::Arc::new(w_q))?; |
| let nvfp4_out = qmatmul.forward(&x)?; |
|
|
| let matmul_diff = ref_out.broadcast_sub(&nvfp4_out)?; |
| let matmul_max_err = matmul_diff.max_all()?.to_scalar::<f32>()?; |
| let matmul_mean_err = matmul_diff.mean_all()?.to_scalar::<f32>()?; |
| let matmul_ref_mean = ref_out.abs()?.mean_all()?.to_scalar::<f32>()?; |
| println!("MatMul max error: {:.6}", matmul_max_err); |
| println!("MatMul mean error: {:.6}", matmul_mean_err); |
| println!("MatMul relative error: {:.6}", matmul_mean_err / matmul_ref_mean); |
|
|
| println!("\n=== NVFP4 Memory Analysis ==="); |
| let n = 4096usize; |
| let k = 4096usize; |
| let fp32_bytes = n * k * 4; |
| let nvfp4_bytes = n * (k / 16) * 9; |
| println!("FP32 weight: {:.1} MB", fp32_bytes as f64 / 1024.0 / 1024.0); |
| println!("NVFP4 weight: {:.1} MB", nvfp4_bytes as f64 / 1024.0 / 1024.0); |
| println!("Compression: {:.2}x", fp32_bytes as f64 / nvfp4_bytes as f64); |
|
|
| println!("\n=== Done ==="); |
| Ok(()) |
| } |
|
|