Files
rustytorch/crates/integration/rtx-candle/src/model.rs
T
2026-03-04 00:08:42 +00:00

292 lines
8.1 KiB
Rust

//! Candle model wrapper and utilities
use crate::error::{CandleError, Result};
use crate::session::CandleConfig;
use candle_core::Tensor as CandleTensor;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
use tracing::{debug, info};
/// Model architecture type
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ModelArchitecture {
/// Generic model
Generic,
/// BERT-style encoder
Bert,
/// GPT-style decoder
Gpt,
/// `LLaMA` architecture
Llama,
/// Mistral architecture
Mistral,
/// Phi architecture
Phi,
/// Vision Transformer
ViT,
/// Stable Diffusion
StableDiffusion,
/// Whisper (speech)
Whisper,
}
/// Information about a loaded model
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelInfo {
/// Model name
pub name: String,
/// Model architecture
pub architecture: ModelArchitecture,
/// Number of parameters
pub num_parameters: usize,
/// Model size in bytes
pub size_bytes: usize,
/// Hidden size
pub hidden_size: usize,
/// Number of layers
pub num_layers: usize,
/// Number of attention heads
pub num_heads: usize,
/// Vocabulary size (for language models)
pub vocab_size: Option<usize>,
/// Maximum sequence length
pub max_seq_len: Option<usize>,
/// Additional metadata
pub metadata: HashMap<String, String>,
}
impl Default for ModelInfo {
fn default() -> Self {
Self {
name: "model".to_string(),
architecture: ModelArchitecture::Generic,
num_parameters: 0,
size_bytes: 0,
hidden_size: 768,
num_layers: 12,
num_heads: 12,
vocab_size: None,
max_seq_len: None,
metadata: HashMap::new(),
}
}
}
/// A loaded Candle model ready for inference
pub struct CandleModel {
/// Model information
info: ModelInfo,
/// Model weights
weights: HashMap<String, CandleTensor>,
/// Device the model is on
device: candle_core::Device,
}
impl CandleModel {
/// Load a model from a file path
pub fn load(path: impl AsRef<Path>, config: &CandleConfig) -> Result<Self> {
let path = path.as_ref();
info!("Loading Candle model from: {}", path.display());
if !path.exists() {
return Err(CandleError::model_load(format!(
"Model file not found: {}",
path.display()
)));
}
let device = config.device.to_candle()?;
// Determine file type and load accordingly
let extension = path.extension().and_then(|e| e.to_str()).unwrap_or("");
match extension {
"safetensors" => Self::load_safetensors(path, device),
"bin" | "pt" | "pth" => Self::load_pytorch(path, device),
"gguf" => Self::load_gguf(path, device),
_ => Err(CandleError::model_load(format!(
"Unsupported model format: {extension}"
))),
}
}
/// Load `SafeTensors` format
fn load_safetensors(path: &Path, device: candle_core::Device) -> Result<Self> {
use safetensors::SafeTensors;
let data = std::fs::read(path)?;
let size_bytes = data.len();
// Parse SafeTensors
let tensors = SafeTensors::deserialize(&data)?;
// Convert to Candle tensors
let mut weights = HashMap::new();
let mut num_parameters = 0;
for (name, tensor_view) in tensors.tensors() {
let shape: Vec<usize> = tensor_view.shape().to_vec();
let numel: usize = shape.iter().product();
num_parameters += numel;
// Get dtype and convert
let dtype = match tensor_view.dtype() {
safetensors::Dtype::F32 => candle_core::DType::F32,
safetensors::Dtype::F16 => candle_core::DType::F16,
safetensors::Dtype::BF16 => candle_core::DType::BF16,
safetensors::Dtype::F64 => candle_core::DType::F64,
safetensors::Dtype::I64 => candle_core::DType::I64,
safetensors::Dtype::U32 => candle_core::DType::U32,
_ => candle_core::DType::F32,
};
// Create Candle tensor (simplified - real implementation would handle all dtypes)
if dtype == candle_core::DType::F32 {
let data: &[f32] = bytemuck::cast_slice(tensor_view.data());
let tensor = CandleTensor::from_slice(data, shape.as_slice(), &device)?;
weights.insert(name.clone(), tensor);
}
}
debug!(
"Loaded {} tensors, {} parameters",
weights.len(),
num_parameters
);
let info = ModelInfo {
name: path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("model")
.to_string(),
num_parameters,
size_bytes,
..Default::default()
};
Ok(Self {
info,
weights,
device,
})
}
/// Load `PyTorch` format (placeholder)
fn load_pytorch(path: &Path, device: candle_core::Device) -> Result<Self> {
let size_bytes = std::fs::metadata(path)
.map(|m| m.len() as usize)
.unwrap_or(0);
let info = ModelInfo {
name: path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("model")
.to_string(),
size_bytes,
..Default::default()
};
Ok(Self {
info,
weights: HashMap::new(),
device,
})
}
/// Load GGUF format (for quantized models)
fn load_gguf(path: &Path, device: candle_core::Device) -> Result<Self> {
let size_bytes = std::fs::metadata(path)
.map(|m| m.len() as usize)
.unwrap_or(0);
let info = ModelInfo {
name: path
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("model")
.to_string(),
size_bytes,
..Default::default()
};
Ok(Self {
info,
weights: HashMap::new(),
device,
})
}
/// Get model information
pub fn info(&self) -> &ModelInfo {
&self.info
}
/// Get model name
pub fn name(&self) -> &str {
&self.info.name
}
/// Get number of parameters
pub fn num_parameters(&self) -> usize {
self.info.num_parameters
}
/// Get model size in bytes
pub fn size_bytes(&self) -> usize {
self.info.size_bytes
}
/// Get a weight tensor by name
pub fn get_weight(&self, name: &str) -> Option<&CandleTensor> {
self.weights.get(name)
}
/// Run forward pass
pub fn forward(
&self,
inputs: HashMap<String, CandleTensor>,
) -> Result<HashMap<String, CandleTensor>> {
debug!("Running forward pass with {} inputs", inputs.len());
// Simulated forward pass
// In real implementation, this would execute the model graph
let mut outputs = HashMap::new();
// Create dummy output based on first input
if let Some((_, input)) = inputs.iter().next() {
let batch_size = input.dims().first().copied().unwrap_or(1);
let output_shape = vec![batch_size, 1000]; // Assume classification output
let output = CandleTensor::zeros(
output_shape.as_slice(),
candle_core::DType::F32,
&self.device,
)?;
outputs.insert("output".to_string(), output);
}
Ok(outputs)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_model_info_default() {
let info = ModelInfo::default();
assert_eq!(info.name, "model");
assert_eq!(info.hidden_size, 768);
assert_eq!(info.num_layers, 12);
}
#[test]
fn test_model_architecture() {
let arch = ModelArchitecture::Llama;
assert_eq!(arch, ModelArchitecture::Llama);
}
}