use std::collections::HashMap;
use std::fs::File;
use std::ops::Range;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicUsize, Ordering};
use memmap2::Mmap;
use onnx_runtime_ir::{DataType, TensorData, ValueId, WeightRef};
use crate::proto::onnx::{ModelProto, TensorProto, tensor_proto};
use crate::{LoaderError, pathsafe::guarded_join};
#[derive(Debug, Default)]
pub struct WeightStore {
pub weights: HashMap<ValueId, WeightRef>,
mmaps: HashMap<PathBuf, MappedFile>,
}
#[derive(Debug)]
struct MappedFile {
id: usize,
mmap: Mmap,
}
static NEXT_MAPPING_ID: AtomicUsize = AtomicUsize::new(1);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ExpertQuantization {
pub bits: usize,
pub block_size: usize,
pub blocks_per_row: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ExpertStorageOrder {
ExpertMajor,
Interleaved,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ExpertTensorLayout {
pub version: u32,
pub experts: usize,
pub rows_per_expert: usize,
pub storage_elements_per_row: usize,
pub order: ExpertStorageOrder,
pub quantization: Option<ExpertQuantization>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum NonPageableReason {
InlineTensor,
UnsupportedLayoutVersion(u32),
NotExpertMajor,
ShapeMismatch {
expected: Vec<usize>,
actual: Vec<usize>,
},
InvalidQuantization(String),
Range(String),
ExternalLengthMismatch {
expected: usize,
actual: usize,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Pageability {
Pageable,
NonPageable(NonPageableReason),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ExpertWeightRegion {
pub expert: usize,
pub offset: usize,
pub len: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WeightRegionCatalog {
path: Option<PathBuf>,
tensor_offset: usize,
tensor_len: usize,
dtype: DataType,
layout: ExpertTensorLayout,
regions: Vec<ExpertWeightRegion>,
pageability: Pageability,
}
impl WeightRegionCatalog {
pub fn classify(weight: &WeightRef, layout: ExpertTensorLayout) -> Self {
let dtype = weight.dtype();
let dims = weight.dims();
let (path, tensor_offset, tensor_len) = match weight {
WeightRef::Inline(tensor) => {
return Self::non_pageable(
None,
0,
tensor.data.len(),
dtype,
layout,
NonPageableReason::InlineTensor,
);
}
WeightRef::External {
path,
offset,
length,
..
} => (Some(path.clone()), *offset, *length),
};
if layout.version != 1 {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout.clone(),
NonPageableReason::UnsupportedLayoutVersion(layout.version),
);
}
if layout.order != ExpertStorageOrder::ExpertMajor {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::NotExpertMajor,
);
}
let expected_shape = vec![
layout.experts,
layout.rows_per_expert,
layout.storage_elements_per_row,
];
if dims != expected_shape {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::ShapeMismatch {
expected: expected_shape,
actual: dims.to_vec(),
},
);
}
if let Some(quantization) = layout.quantization
&& (!matches!(quantization.bits, 1 | 2 | 4 | 8)
|| quantization.block_size == 0
|| quantization.blocks_per_row == 0)
{
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::InvalidQuantization(format!(
"bits={}, block_size={}, blocks_per_row={}",
quantization.bits, quantization.block_size, quantization.blocks_per_row
)),
);
}
let elements_per_expert = match checked_product(
&[layout.rows_per_expert, layout.storage_elements_per_row],
"per-expert element count",
) {
Ok(value) => value,
Err(error) => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(error.to_string()),
);
}
};
let bytes_per_expert =
match checked_storage_byte_count(dtype, elements_per_expert, "per-expert byte count") {
Ok(value) => value,
Err(error) => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(error.to_string()),
);
}
};
let expected_len = match checked_byte_count(
layout.experts,
bytes_per_expert,
"expert tensor byte count",
) {
Ok(value) => value,
Err(error) => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(error.to_string()),
);
}
};
if expected_len != tensor_len {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::ExternalLengthMismatch {
expected: expected_len,
actual: tensor_len,
},
);
}
let tensor_end = match tensor_offset.checked_add(tensor_len) {
Some(end) if end <= isize::MAX as usize => end,
Some(_) => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(
"expert tensor absolute endpoint exceeds isize::MAX".into(),
),
);
}
None => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range("expert tensor absolute endpoint overflow".into()),
);
}
};
let mut regions = Vec::new();
if let Err(error) = regions.try_reserve_exact(layout.experts) {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(format!("expert region allocation failed: {error}")),
);
}
for expert in 0..layout.experts {
let relative = match checked_range(expert, bytes_per_expert, "expert byte range") {
Ok(value) => value,
Err(error) => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(error.to_string()),
);
}
};
let offset = match tensor_offset.checked_add(relative.start) {
Some(value) => value,
None => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range("expert absolute offset overflow".into()),
);
}
};
let _end = match offset.checked_add(bytes_per_expert) {
Some(end) if end <= isize::MAX as usize && end <= tensor_end => end,
Some(_) => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range(
"expert absolute endpoint exceeds validated tensor range".into(),
),
);
}
None => {
return Self::non_pageable(
path,
tensor_offset,
tensor_len,
dtype,
layout,
NonPageableReason::Range("expert absolute endpoint overflow".into()),
);
}
};
regions.push(ExpertWeightRegion {
expert,
offset,
len: bytes_per_expert,
});
}
Self {
path,
tensor_offset,
tensor_len,
dtype,
layout,
regions,
pageability: Pageability::Pageable,
}
}
pub fn for_mapped_tensor_view(
dtype: DataType,
dims: &[usize],
tensor_len: usize,
layout: ExpertTensorLayout,
) -> Self {
let synthetic = WeightRef::External {
path: PathBuf::new(),
offset: 0,
length: tensor_len,
dtype,
dims: dims.to_vec(),
};
let mut catalog = Self::classify(&synthetic, layout);
catalog.path = None;
catalog
}
fn non_pageable(
path: Option<PathBuf>,
tensor_offset: usize,
tensor_len: usize,
dtype: DataType,
layout: ExpertTensorLayout,
reason: NonPageableReason,
) -> Self {
Self {
path,
tensor_offset,
tensor_len,
dtype,
layout,
regions: Vec::new(),
pageability: Pageability::NonPageable(reason),
}
}
pub fn pageability(&self) -> &Pageability {
&self.pageability
}
pub fn is_pageable(&self) -> bool {
matches!(self.pageability, Pageability::Pageable)
}
pub fn layout(&self) -> &ExpertTensorLayout {
&self.layout
}
pub fn dtype(&self) -> DataType {
self.dtype
}
pub fn mapped_bytes(&self) -> usize {
if self.is_pageable() {
self.tensor_len
} else {
0
}
}
pub fn region(&self, expert: usize) -> Option<&ExpertWeightRegion> {
self.regions.get(expert)
}
pub fn relative_range(&self, expert: usize) -> Option<Range<usize>> {
let region = self.region(expert)?;
let start = region.offset.checked_sub(self.tensor_offset)?;
let end = start.checked_add(region.len)?;
Some(start..end)
}
pub fn expert_bytes<'a>(&self, store: &'a WeightStore, expert: usize) -> Option<&'a [u8]> {
let path = self.path.as_ref()?;
let region = self.region(expert)?;
store.external_bytes(path, region.offset, region.len)
}
pub fn tensor_offset(&self) -> usize {
self.tensor_offset
}
}
#[derive(Clone, Debug, thiserror::Error, PartialEq, Eq)]
#[error("{0}")]
pub struct WeightRangeError(String);
pub fn checked_product(factors: &[usize], context: &str) -> Result<usize, WeightRangeError> {
let mut product = 1usize;
let mut has_zero = false;
for &factor in factors {
if factor == 0 {
has_zero = true;
} else {
product = product
.checked_mul(factor)
.ok_or_else(|| WeightRangeError(format!("{context} overflow")))?;
}
}
Ok(if has_zero { 0 } else { product })
}
pub fn checked_byte_count(
elements: usize,
element_size: usize,
context: &str,
) -> Result<usize, WeightRangeError> {
let bytes = elements
.checked_mul(element_size)
.ok_or_else(|| WeightRangeError(format!("{context} overflow")))?;
if bytes > isize::MAX as usize {
return Err(WeightRangeError(format!("{context} exceeds isize::MAX")));
}
Ok(bytes)
}
pub fn checked_storage_byte_count(
dtype: DataType,
elements: usize,
context: &str,
) -> Result<usize, WeightRangeError> {
let bytes = dtype
.checked_storage_bytes(elements)
.ok_or_else(|| WeightRangeError(format!("{context} overflow")))?;
if bytes > isize::MAX as usize {
return Err(WeightRangeError(format!("{context} exceeds isize::MAX")));
}
Ok(bytes)
}
pub fn checked_range(
index: usize,
width: usize,
context: &str,
) -> Result<Range<usize>, WeightRangeError> {
let start = index
.checked_mul(width)
.ok_or_else(|| WeightRangeError(format!("{context} start offset overflow")))?;
let end = start
.checked_add(width)
.ok_or_else(|| WeightRangeError(format!("{context} end offset overflow")))?;
if end > isize::MAX as usize {
return Err(WeightRangeError(format!("{context} exceeds isize::MAX")));
}
Ok(start..end)
}
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.mmap.get(*offset..offset.checked_add(*length)?)
}
}
}
pub fn external_mmap_provenance(&self, weight: &WeightRef) -> Option<(usize, usize, usize)> {
let WeightRef::External {
path,
offset,
length,
..
} = weight
else {
return None;
};
let mmap = self.mmaps.get(path)?;
let end = offset.checked_add(*length)?;
if end > mmap.mmap.len() || end > isize::MAX as usize {
return None;
}
Some((mmap.id, *offset, *length))
}
fn external_bytes(&self, path: &Path, offset: usize, length: usize) -> Option<&[u8]> {
let mmap = self.mmaps.get(path)?;
mmap.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()))?;
let id = NEXT_MAPPING_ID
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |id| id.checked_add(1))
.map_err(|_| LoaderError::Mmap("external mmap identity space exhausted".into()))?;
self.mmaps
.insert(path.to_path_buf(), MappedFile { id, 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.mmap.len()) {
return Err(LoaderError::Mmap(format!(
"external initializer {:?}: window [{offset}, {:?}) exceeds file {} ({} bytes)",
init.name,
end,
path.display(),
mmap.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::Undefined => {
return Err(LoaderError::UnsupportedDataType {
raw: 0,
context: format!("tensor {:?}", proto.name),
});
}
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::Complex64 => proto
.float_data
.iter()
.flat_map(|v| v.to_le_bytes())
.collect(),
DataType::Complex128 => 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()
}
DataType::Float8E4M3FN
| DataType::Float8E4M3FNUZ
| DataType::Float8E5M2
| DataType::Float8E5M2FNUZ
| DataType::Float8E8M0
| DataType::Int4
| DataType::Uint4
| DataType::Float4E2M1
| DataType::Int2
| DataType::Uint2 => proto.int32_data.iter().map(|v| *v as u8).collect(),
DataType::String => unreachable!("STRING tensors returned above"),
};
Ok(td)
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::{SystemTime, UNIX_EPOCH};
#[test]
fn float8_typed_data_preserves_each_byte() {
let proto = TensorProto {
data_type: DataType::Float8E4M3FN.to_onnx(),
dims: vec![3],
int32_data: vec![0x01, 0x7f, 0xff],
..Default::default()
};
let data =
tensor_data_from_proto(&proto, DataType::Float8E4M3FN, &[3]).expect("tensor data");
assert_eq!(data.data, [0x01, 0x7f, 0xff]);
}
#[test]
fn four_bit_typed_data_preserves_packed_nibbles() {
let proto = TensorProto {
data_type: DataType::Int4.to_onnx(),
dims: vec![3],
int32_data: vec![0x21, 0x03],
..Default::default()
};
let data = tensor_data_from_proto(&proto, DataType::Int4, &[3]).expect("tensor data");
assert_eq!(data.data, [0x21, 0x03]);
}
#[test]
fn two_bit_typed_data_preserves_four_packed_elements_per_byte() {
let proto = TensorProto {
data_type: DataType::Int2.to_onnx(),
dims: vec![5],
int32_data: vec![0b11_10_01_00, 0b0000_0001],
..Default::default()
};
let data = tensor_data_from_proto(&proto, DataType::Int2, &[5]).expect("tensor data");
assert_eq!(data.data, [0b11_10_01_00, 0b0000_0001]);
assert_eq!(data.data.len(), DataType::Int2.storage_bytes(5));
}
#[test]
fn float8e8m0_typed_data_preserves_each_byte() {
let proto = TensorProto {
data_type: DataType::Float8E8M0.to_onnx(),
dims: vec![2],
int32_data: vec![0x7f, 0xff],
..Default::default()
};
let data = tensor_data_from_proto(&proto, DataType::Float8E8M0, &[2]).expect("tensor data");
assert_eq!(data.data, [0x7f, 0xff]);
}
fn external_weight(path: PathBuf, length: usize, dims: Vec<usize>) -> WeightRef {
WeightRef::External {
path,
offset: 16,
length,
dtype: DataType::Uint8,
dims,
}
}
fn expert_layout(order: ExpertStorageOrder) -> ExpertTensorLayout {
ExpertTensorLayout {
version: 1,
experts: 3,
rows_per_expert: 2,
storage_elements_per_row: 4,
order,
quantization: Some(ExpertQuantization {
bits: 4,
block_size: 16,
blocks_per_row: 1,
}),
}
}
#[test]
fn expert_major_external_tensor_catalogs_contiguous_ranges() {
let stamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("clock")
.as_nanos();
let path = std::env::current_dir()
.expect("cwd")
.join("target")
.join(format!(
"weight-region-catalog-{}-{stamp}.bin",
std::process::id()
));
std::fs::create_dir_all(path.parent().expect("parent")).expect("create target");
let mut bytes = vec![0u8; 40];
for (index, byte) in bytes.iter_mut().enumerate() {
*byte = index as u8;
}
std::fs::write(&path, &bytes).expect("write external data");
let weight = external_weight(path.clone(), 24, vec![3, 2, 4]);
let catalog =
WeightRegionCatalog::classify(&weight, expert_layout(ExpertStorageOrder::ExpertMajor));
assert_eq!(catalog.pageability(), &Pageability::Pageable);
assert_eq!(catalog.mapped_bytes(), 24);
assert_eq!(
catalog.region(1),
Some(&ExpertWeightRegion {
expert: 1,
offset: 24,
len: 8,
})
);
let mut store = WeightStore::new();
store.map_external(&path).expect("map external data");
assert_eq!(catalog.expert_bytes(&store, 1), Some(&bytes[24..32]));
drop(store);
std::fs::remove_file(path).expect("remove external data");
}
#[test]
fn interleaved_external_tensor_is_non_pageable_without_error() {
let weight = external_weight(PathBuf::from("weights.bin"), 24, vec![3, 2, 4]);
let catalog =
WeightRegionCatalog::classify(&weight, expert_layout(ExpertStorageOrder::Interleaved));
assert_eq!(
catalog.pageability(),
&Pageability::NonPageable(NonPageableReason::NotExpertMajor)
);
assert!(catalog.region(0).is_none());
}
#[test]
fn catalog_range_math_rejects_zero_masked_overflow_and_isize_excess() {
let overflow = ExpertTensorLayout {
version: 1,
experts: 0,
rows_per_expert: usize::MAX,
storage_elements_per_row: 2,
order: ExpertStorageOrder::ExpertMajor,
quantization: None,
};
let weight = external_weight(PathBuf::from("weights.bin"), 0, vec![0, usize::MAX, 2]);
assert!(matches!(
WeightRegionCatalog::classify(&weight, overflow).pageability(),
Pageability::NonPageable(NonPageableReason::Range(message))
if message.contains("overflow")
));
assert!(
checked_range(0, isize::MAX as usize + 1, "range")
.unwrap_err()
.to_string()
.contains("isize::MAX")
);
let endpoint = WeightRef::External {
path: PathBuf::from("weights.bin"),
offset: isize::MAX as usize,
length: 1,
dtype: DataType::Uint8,
dims: vec![1, 1, 1],
};
let layout = ExpertTensorLayout {
version: 1,
experts: 1,
rows_per_expert: 1,
storage_elements_per_row: 1,
order: ExpertStorageOrder::ExpertMajor,
quantization: None,
};
assert!(matches!(
WeightRegionCatalog::classify(&endpoint, layout).pageability(),
Pageability::NonPageable(NonPageableReason::Range(message))
if message.contains("endpoint")
));
}
}