use std::collections::HashMap;
use std::fs::File;
use std::path::{Path, PathBuf};
use memmap2::Mmap;
use onnx_runtime_ir::{DataType, TensorData, ValueId, WeightRef};
use crate::proto::onnx::{tensor_proto, ModelProto, TensorProto};
use crate::{pathsafe::guarded_join, LoaderError};
#[derive(Debug, Default)]
pub struct WeightStore {
pub weights: HashMap<ValueId, WeightRef>,
mmaps: HashMap<PathBuf, Mmap>,
}
impl WeightStore {
pub fn new() -> Self {
Self::default()
}
pub fn map_external(&mut self, path: impl AsRef<Path>) -> Result<(), LoaderError> {
self.mmap_file(path.as_ref())
}
pub fn bytes<'a>(&'a self, weight: &'a WeightRef) -> Option<&'a [u8]> {
match weight {
WeightRef::Inline(t) => Some(&t.data),
WeightRef::External {
path,
offset,
length,
..
} => {
let mmap = self.mmaps.get(path)?;
mmap.get(*offset..offset.checked_add(*length)?)
}
}
}
fn mmap_file(&mut self, path: &Path) -> Result<(), LoaderError> {
if self.mmaps.contains_key(path) {
return Ok(());
}
let file = File::open(path).map_err(|_| LoaderError::ExternalDataNotFound {
path: path.to_path_buf(),
})?;
let mmap = unsafe { Mmap::map(&file) }.map_err(|e| LoaderError::Mmap(e.to_string()))?;
self.mmaps.insert(path.to_path_buf(), mmap);
Ok(())
}
}
pub fn load_weights(
model: &ModelProto,
model_dir: &Path,
name_map: &HashMap<String, ValueId>,
) -> Result<WeightStore, LoaderError> {
let mut store = WeightStore::new();
let Some(graph) = model.graph.as_ref() else {
return Ok(store);
};
for init in &graph.initializer {
let Some(&vid) = name_map.get(&init.name) else {
continue;
};
let weight = resolve_initializer(&mut store, init, model_dir)?;
store.weights.insert(vid, weight);
}
Ok(store)
}
fn resolve_initializer(
store: &mut WeightStore,
init: &TensorProto,
model_dir: &Path,
) -> Result<WeightRef, LoaderError> {
let dtype = DataType::from_onnx(init.data_type).ok_or_else(|| {
LoaderError::UnsupportedDataType {
raw: init.data_type,
context: format!("initializer {:?}", init.name),
}
})?;
let dims: Vec<usize> = init.dims.iter().map(|&d| d.max(0) as usize).collect();
if init.data_location == tensor_proto::DataLocation::External as i32 {
let mut location = None;
let mut offset: usize = 0;
let mut length: Option<usize> = None;
for kv in &init.external_data {
match kv.key.as_str() {
"location" => location = Some(kv.value.clone()),
"offset" => offset = kv.value.parse().unwrap_or(0),
"length" => length = kv.value.parse().ok(),
_ => {}
}
}
let location = location.ok_or_else(|| {
LoaderError::GraphBuild(format!(
"external initializer {:?} missing 'location'",
init.name
))
})?;
let path = resolve_external_path(model_dir, &location)?;
store.mmap_file(&path)?;
let numel: usize = dims.iter().product();
let length = length.unwrap_or_else(|| dtype.storage_bytes(numel));
if let Some(mmap) = store.mmaps.get(&path) {
let end = offset.checked_add(length);
if end.is_none_or(|e| e > mmap.len()) {
return Err(LoaderError::Mmap(format!(
"external initializer {:?}: window [{offset}, {:?}) exceeds file {} ({} bytes)",
init.name,
end,
path.display(),
mmap.len()
)));
}
}
Ok(WeightRef::External {
path,
offset,
length,
dtype,
dims,
})
} else {
let data = tensor_data_from_proto(init, dtype, &dims)?;
Ok(WeightRef::Inline(data))
}
}
fn resolve_external_path(model_dir: &Path, location: &str) -> Result<PathBuf, LoaderError> {
guarded_join(model_dir, location).map_err(|reason| LoaderError::ExternalDataPath {
path: location.to_string(),
reason,
})
}
pub(crate) fn tensor_data_from_proto(
proto: &TensorProto,
dtype: DataType,
dims: &[usize],
) -> Result<TensorData, LoaderError> {
let mut td = TensorData::from_raw(dtype, dims.to_vec(), Vec::new());
if !proto.name.is_empty() {
td.name = Some(proto.name.clone());
}
if dtype == DataType::String {
td.strings = proto
.string_data
.iter()
.map(|b| String::from_utf8_lossy(b).into_owned())
.collect();
return Ok(td);
}
if !proto.raw_data.is_empty() {
td.data = proto.raw_data.clone();
return Ok(td);
}
td.data = match dtype {
DataType::Float32 => proto
.float_data
.iter()
.flat_map(|v| v.to_le_bytes())
.collect(),
DataType::Float64 => proto
.double_data
.iter()
.flat_map(|v| v.to_le_bytes())
.collect(),
DataType::Int64 => proto
.int64_data
.iter()
.flat_map(|v| v.to_le_bytes())
.collect(),
DataType::Uint64 | DataType::Uint32 => proto
.uint64_data
.iter()
.flat_map(|v| match dtype {
DataType::Uint32 => (*v as u32).to_le_bytes().to_vec(),
_ => v.to_le_bytes().to_vec(),
})
.collect(),
DataType::Int32 => proto
.int32_data
.iter()
.flat_map(|v| v.to_le_bytes())
.collect(),
DataType::Int16 | DataType::Uint16 | DataType::Float16 | DataType::BFloat16 => proto
.int32_data
.iter()
.flat_map(|v| (*v as u16).to_le_bytes())
.collect(),
DataType::Int8 | DataType::Uint8 | DataType::Bool => {
proto.int32_data.iter().map(|v| *v as u8).collect()
}
_ => Vec::new(),
};
Ok(td)
}