use crate::{
config::ModelConfig,
error::{Error, Result},
loader::{LoadOptions, LoadedModel},
name_mapping::Architecture,
progress::{ProgressEvent, ProgressFn},
smart_mapping::SmartTensorNameMapper,
validation,
};
use candle_core::{DType, Device, Shape, Tensor};
use std::{collections::HashMap, fs, path::Path};
#[cfg(feature = "onnx")]
use prost::Message;
#[cfg(feature = "onnx")]
pub struct ONNXLoader {
progress_fn: Option<ProgressFn>,
device: Device,
dtype: DType,
}
#[cfg(feature = "onnx")]
pub struct ONNXLoadOptions {
pub device: Device,
pub dtype: DType,
pub validate_shapes: bool,
pub use_f16: bool,
pub progress: Option<Box<dyn Fn(ProgressEvent) + Send + Sync>>,
}
impl Default for ONNXLoadOptions {
fn default() -> Self {
Self {
device: Device::Cpu,
dtype: DType::F32,
validate_shapes: true,
use_f16: false,
progress: None,
}
}
}
#[cfg(feature = "onnx")]
#[derive(Debug, Clone)]
pub struct ONNXModelInfo {
pub model_version: i64,
pub producer_name: String,
pub producer_version: String,
pub domain: String,
pub doc_string: String,
pub graph_name: String,
pub num_nodes: usize,
pub inputs: Vec<String>,
pub outputs: Vec<String>,
pub architecture: Architecture,
}
#[cfg(feature = "onnx")]
impl ONNXLoader {
pub fn new(options: ONNXLoadOptions) -> Self {
Self {
progress_fn: options.progress,
device: options.device,
dtype: options.dtype,
}
}
pub fn load_from_path(&self, path: &Path, load_options: &LoadOptions) -> Result<LoadedModel> {
let model_bytes = fs::read(path).map_err(|e| {
Error::model_loading(format!("Failed to read ONNX file {:?}: {}", path, e))
})?;
self.load_from_bytes(&model_bytes, load_options)
}
pub fn load_from_bytes(&self, data: &[u8], load_options: &LoadOptions) -> Result<LoadedModel> {
if let Some(ref progress) = self.progress_fn {
progress(ProgressEvent::LoadingFile {
file: "onnx_model.onnx".into(),
format: "ONNX".to_string(),
});
}
let onnx_model = self.parse_onnx_model(data)?;
let model_info = self.extract_model_info(&onnx_model)?;
if let Some(ref progress) = self.progress_fn {
progress(ProgressEvent::Status {
message: format!("Parsed ONNX model: {} nodes", model_info.num_nodes),
});
}
let tensors = self.extract_tensors(&onnx_model, &model_info)?;
if let Some(ref progress) = self.progress_fn {
progress(ProgressEvent::LoadingTensorsFromFiles {
count: tensors.len(),
format: "ONNX".to_string(),
});
}
let tensor_names: Vec<String> = tensors.keys().cloned().collect();
let mut name_mapper = SmartTensorNameMapper::from_tensor_names(&tensor_names)?;
if let Some(oracle) = load_options.smart_mapping_oracle.as_ref() {
}
let config = self.infer_model_config(&model_info, &tensors, &name_mapper)?;
if let Some(ref progress) = self.progress_fn {
progress(ProgressEvent::DetectingArchitecture);
}
validation::validate_memory_requirements(&config, self.dtype)?;
let converted_tensors = self.convert_tensors(tensors, &self.device, self.dtype)?;
if let Some(ref progress) = self.progress_fn {
progress(ProgressEvent::Complete {
tensor_count: converted_tensors.len(),
format: "ONNX".to_string(),
});
}
use candle_nn::VarBuilder;
let var_builder =
VarBuilder::from_tensors(converted_tensors.clone(), self.dtype, &self.device);
Ok(LoadedModel {
var_builder,
config,
name_mapper,
raw_tensors: converted_tensors,
quantized_tensors: None,
metadata: crate::metadata::ModelMetadata::new(),
tensor_info: HashMap::new(),
quantization_info: None,
provenance: crate::metadata::ModelProvenance::new(),
})
}
fn parse_onnx_model(&self, data: &[u8]) -> Result<onnx_proto::ModelProto> {
onnx_proto::ModelProto::decode(data)
.map_err(|e| Error::invalid_format(format!("Failed to parse ONNX protobuf: {}", e)))
}
fn extract_model_info(&self, model: &onnx_proto::ModelProto) -> Result<ONNXModelInfo> {
let graph = model
.graph
.as_ref()
.ok_or_else(|| Error::invalid_format("ONNX model missing computational graph"))?;
let inputs: Vec<String> = graph
.input
.iter()
.filter_map(|input| input.name.clone())
.collect();
let outputs: Vec<String> = graph
.output
.iter()
.filter_map(|output| output.name.clone())
.collect();
let weight_names: Vec<String> = graph
.initializer_tensor
.iter()
.filter_map(|init| init.name.clone())
.collect();
let temp_name_mapper = SmartTensorNameMapper::from_tensor_names(&weight_names)
.unwrap_or_else(|_| SmartTensorNameMapper::new());
let architecture = temp_name_mapper
.architecture()
.copied()
.unwrap_or(Architecture::Unknown);
Ok(ONNXModelInfo {
model_version: model.model_version.unwrap_or(0),
producer_name: model
.producer_name
.clone()
.unwrap_or_else(|| "Unknown".to_string()),
producer_version: model
.producer_version
.clone()
.unwrap_or_else(|| "Unknown".to_string()),
domain: model.domain.clone().unwrap_or_else(|| "".to_string()),
doc_string: model.doc_string.clone().unwrap_or_else(|| "".to_string()),
graph_name: graph.name.clone(),
num_nodes: graph.node.len(),
inputs,
outputs,
architecture,
})
}
fn extract_tensors(
&self,
model: &onnx_proto::ModelProto,
_info: &ONNXModelInfo,
) -> Result<HashMap<String, Tensor>> {
let graph = model.graph.as_ref().unwrap(); let mut tensors = HashMap::new();
for (i, initializer) in graph.initializer_tensor.iter().enumerate() {
if let Some(ref progress) = self.progress_fn {
progress(ProgressEvent::LoadingTensors {
current: i + 1,
total: graph.initializer_tensor.len(),
file_name: initializer.name.clone(),
});
}
let tensor = self.convert_onnx_tensor(initializer)?;
if let Some(name) = &initializer.name {
tensors.insert(name.clone(), tensor);
}
}
Ok(tensors)
}
fn convert_onnx_tensor(&self, tensor_proto: &onnx_proto::TensorProto) -> Result<Tensor> {
let dims: Vec<usize> = tensor_proto.dims.iter().map(|&d| d as usize).collect();
let shape = Shape::from_dims(&dims);
match tensor_proto.data_type {
Some(1) => {
let data = if !tensor_proto.float_data.is_empty() {
tensor_proto.float_data.clone()
} else if let Some(ref raw_data) = tensor_proto.raw_data {
if !raw_data.is_empty() {
if raw_data.len() % 4 != 0 {
return Err(Error::invalid_format("Invalid f32 raw data length"));
}
raw_data
.chunks_exact(4)
.map(|chunk| {
f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]])
})
.collect()
} else {
return Err(Error::invalid_format("ONNX tensor missing float data"));
}
} else {
return Err(Error::invalid_format("ONNX tensor missing float data"));
};
Tensor::from_vec(data, shape, &self.device).map_err(|e| {
Error::model_loading(format!("Failed to create f32 tensor: {}", e))
})
}
Some(10) => {
let data: Vec<f32> = if !tensor_proto.int32_data.is_empty() {
tensor_proto
.int32_data
.iter()
.map(|&i| half::f16::from_bits(i as u16).to_f32())
.collect()
} else if let Some(ref raw_data) = tensor_proto.raw_data {
if !raw_data.is_empty() {
if raw_data.len() % 2 != 0 {
return Err(Error::invalid_format("Invalid f16 raw data length"));
}
raw_data
.chunks_exact(2)
.map(|chunk| {
let bits = u16::from_le_bytes([chunk[0], chunk[1]]);
half::f16::from_bits(bits).to_f32()
})
.collect()
} else {
return Err(Error::invalid_format("ONNX tensor missing f16 data"));
}
} else {
return Err(Error::invalid_format("ONNX tensor missing f16 data"));
};
Tensor::from_vec(data, shape, &self.device).map_err(|e| {
Error::model_loading(format!("Failed to create f16 tensor: {}", e))
})
}
Some(6) => {
let data = if !tensor_proto.int32_data.is_empty() {
tensor_proto.int32_data.clone()
} else if let Some(ref raw_data) = tensor_proto.raw_data {
if !raw_data.is_empty() {
if raw_data.len() % 4 != 0 {
return Err(Error::invalid_format("Invalid int32 raw data length"));
}
raw_data
.chunks_exact(4)
.map(|chunk| {
i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]])
})
.collect()
} else {
return Err(Error::invalid_format("ONNX tensor missing int32 data"));
}
} else {
return Err(Error::invalid_format("ONNX tensor missing int32 data"));
};
let float_data: Vec<f32> = data.into_iter().map(|i| i as f32).collect();
Tensor::from_vec(float_data, shape, &self.device).map_err(|e| {
Error::model_loading(format!("Failed to create int32 tensor: {}", e))
})
}
Some(7) => {
let data = if !tensor_proto.int64_data.is_empty() {
tensor_proto.int64_data.clone()
} else if let Some(ref raw_data) = tensor_proto.raw_data {
if !raw_data.is_empty() {
if raw_data.len() % 8 != 0 {
return Err(Error::invalid_format("Invalid int64 raw data length"));
}
raw_data
.chunks_exact(8)
.map(|chunk| {
i64::from_le_bytes([
chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5],
chunk[6], chunk[7],
])
})
.collect()
} else {
return Err(Error::invalid_format("ONNX tensor missing int64 data"));
}
} else {
return Err(Error::invalid_format("ONNX tensor missing int64 data"));
};
let float_data: Vec<f32> = data.into_iter().map(|i| i as f32).collect();
Tensor::from_vec(float_data, shape, &self.device).map_err(|e| {
Error::model_loading(format!("Failed to create int64 tensor: {}", e))
})
}
_ => Err(Error::invalid_format(format!(
"Unsupported ONNX tensor data type: {:?}",
tensor_proto.data_type
))),
}
}
fn infer_model_config(
&self,
info: &ONNXModelInfo,
tensors: &HashMap<String, Tensor>,
name_mapper: &SmartTensorNameMapper,
) -> Result<ModelConfig> {
let mut vocab_size = 50257; let mut hidden_size = 768;
let mut num_layers = 12;
let mut num_heads = 12;
let mut intermediate_size = 3072;
let mut max_pos_embeddings = 2048;
for (name, tensor) in tensors {
let shape = tensor.shape();
if name.contains("embed") && name.contains("weight") {
if shape.rank() == 2 {
vocab_size = shape.dims()[0];
hidden_size = shape.dims()[1];
}
}
if name.contains("attn") && name.contains("weight") {
if shape.rank() == 2 && shape.dims()[0] == shape.dims()[1] {
hidden_size = shape.dims()[0];
}
}
if let Some(layer_num) = extract_layer_number(name) {
num_layers = num_layers.max(layer_num + 1);
}
if (name.contains("mlp") || name.contains("ffn")) && name.contains("weight") {
if shape.rank() == 2 {
let dim0 = shape.dims()[0];
let dim1 = shape.dims()[1];
if dim0 > hidden_size || dim1 > hidden_size {
intermediate_size = dim0.max(dim1);
}
}
}
}
num_heads = if hidden_size % 64 == 0 {
hidden_size / 64
} else if hidden_size % 32 == 0 {
hidden_size / 32
} else {
(hidden_size / 64).max(1)
};
Ok(ModelConfig {
vocab_size,
hidden_size,
num_attention_heads: num_heads,
num_hidden_layers: num_layers,
intermediate_size,
max_position_embeddings: max_pos_embeddings,
dropout: 0.1,
layer_norm_eps: 1e-5,
attention_dropout: 0.1,
activation_function: "gelu".to_string(),
rope_theta: 10000.0,
tie_word_embeddings: false,
architecture: info.architecture,
raw_config: serde_json::Value::Null,
})
}
fn convert_tensors(
&self,
tensors: HashMap<String, Tensor>,
device: &Device,
dtype: DType,
) -> Result<HashMap<String, Tensor>> {
let mut converted = HashMap::with_capacity(tensors.len());
for (name, tensor) in tensors {
let converted_tensor = tensor
.to_device(device)
.map_err(|e| {
Error::model_loading(format!("Failed to move tensor to device: {}", e))
})?
.to_dtype(dtype)
.map_err(|e| {
Error::model_loading(format!("Failed to convert tensor dtype: {}", e))
})?;
converted.insert(name, converted_tensor);
}
Ok(converted)
}
}
fn extract_layer_number(name: &str) -> Option<usize> {
for part in name.split('.') {
if let Ok(num) = part.parse::<usize>() {
return Some(num);
}
}
None
}
#[cfg(feature = "onnx")]
pub fn load_onnx<P: AsRef<Path>>(path: P, mut options: LoadOptions) -> Result<LoadedModel> {
let onnx_options = ONNXLoadOptions {
device: options.device.clone(),
dtype: options.dtype,
validate_shapes: true,
use_f16: matches!(options.dtype, DType::F16 | DType::BF16),
progress: options.progress.take(),
};
let loader = ONNXLoader::new(onnx_options);
loader.load_from_path(path.as_ref(), &options)
}
#[cfg(feature = "onnx")]
pub mod onnx_proto {
include!(concat!(env!("OUT_DIR"), "/onnx.rs"));
}
#[cfg(not(feature = "onnx"))]
pub fn load_onnx<P: AsRef<Path>>(
_path: P,
_options: crate::loader::LoadOptions,
) -> Result<crate::loader::LoadedModel> {
Err(Error::invalid_format(
"ONNX support not enabled. Enable the 'onnx' feature to load ONNX models.",
))
}
#[cfg(not(feature = "onnx"))]
pub struct ONNXLoader;
#[cfg(not(feature = "onnx"))]
#[derive(Debug, Clone)]
pub struct ONNXLoadOptions {
pub device: Device,
pub dtype: DType,
}
#[cfg(not(feature = "onnx"))]
impl Default for ONNXLoadOptions {
fn default() -> Self {
Self {
device: Device::Cpu,
dtype: DType::F32,
}
}
}