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);
impl WeightStore {
#[must_use]
pub fn mapped_external_bytes(&self) -> u64 {
self.mmaps
.values()
.map(|mapped| mapped.mmap.len() as u64)
.fold(0u64, u64::saturating_add)
}
#[must_use]
pub fn unreferenced_external_bytes(&self, referenced: &ReferencedWeightBytes) -> u64 {
self.mapped_external_bytes()
.saturating_sub(referenced.external)
}
}
#[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 path(&self) -> Option<&Path> {
self.path.as_deref()
}
pub fn tensor_offset(&self) -> usize {
self.tensor_offset
}
pub fn tensor_len(&self) -> usize {
self.tensor_len
}
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 qmoe_expert_tensor_layout(
bits: usize,
block_size: usize,
blocks_per_row: usize,
dims: &[usize],
) -> Option<ExpertTensorLayout> {
if dims.len() != 3 {
return None;
}
Some(ExpertTensorLayout {
version: 1,
experts: dims[0],
rows_per_expert: dims[1],
storage_elements_per_row: dims[2],
order: ExpertStorageOrder::ExpertMajor,
quantization: Some(ExpertQuantization {
bits,
block_size,
blocks_per_row,
}),
})
}
#[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))
}
pub fn mmap_region_bytes(&self, mapping_id: usize, offset: usize, len: usize) -> Option<&[u8]> {
let mmap = self.mmaps.values().find(|mapped| mapped.id == mapping_id)?;
mmap.mmap.get(offset..offset.checked_add(len)?)
}
pub fn mmap_full_bytes(&self, mapping_id: usize) -> Option<&[u8]> {
let mmap = self.mmaps.values().find(|mapped| mapped.id == mapping_id)?;
Some(&mmap.mmap[..])
}
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(())
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct ReferencedWeightBytes {
pub inline: u64,
pub external: u64,
}
impl ReferencedWeightBytes {
#[must_use]
pub fn total(&self) -> u64 {
self.inline.saturating_add(self.external)
}
}
#[must_use]
pub fn referenced_weight_bytes(model: &ModelProto) -> ReferencedWeightBytes {
let mut totals = ReferencedWeightBytes::default();
let Some(graph) = model.graph.as_ref() else {
return totals;
};
for init in &graph.initializer {
let dims: Vec<usize> = init.dims.iter().map(|&d| d.max(0) as usize).collect();
let declared_geometry = DataType::from_onnx(init.data_type).and_then(|dtype| {
TensorData::from_raw(dtype, dims, Vec::new()).checked_expected_bytes()
});
if init.data_location == tensor_proto::DataLocation::External as i32 {
let declared_length = init
.external_data
.iter()
.find(|kv| kv.key == "length")
.and_then(|kv| kv.value.parse::<u64>().ok());
let bytes = declared_length
.or_else(|| declared_geometry.map(|bytes| bytes as u64))
.unwrap_or(0);
totals.external = totals.external.saturating_add(bytes);
} else {
let bytes = if init.raw_data.is_empty() {
declared_geometry.map_or(0, |bytes| bytes as u64)
} else {
init.raw_data.len() as u64
};
totals.inline = totals.inline.saturating_add(bytes);
}
}
totals
}
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 expected_length = TensorData::from_raw(dtype, dims.clone(), Vec::new())
.checked_expected_bytes()
.ok_or_else(|| {
LoaderError::GraphBuild(format!(
"external initializer {:?} geometry overflows for shape {dims:?} and dtype {dtype:?}",
init.name
))
})?;
let length = length.unwrap_or(expected_length);
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 external_initializer_rejects_geometry_overflow() {
let stamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("clock")
.as_nanos();
let model_dir = std::env::current_dir().expect("cwd").join("target");
std::fs::create_dir_all(&model_dir).expect("create target");
let file_name = format!(
"overflowing-external-weight-{}-{stamp}.bin",
std::process::id()
);
let path = model_dir.join(&file_name);
std::fs::write(&path, [0u8]).expect("write external data");
let initializer = TensorProto {
name: "huge".to_string(),
data_type: DataType::Float32.to_onnx(),
dims: vec![i64::MAX, 3],
external_data: vec![crate::proto::onnx::StringStringEntryProto {
key: "location".to_string(),
value: file_name,
}],
data_location: tensor_proto::DataLocation::External as i32,
..Default::default()
};
let mut store = WeightStore::new();
let error = resolve_initializer(&mut store, &initializer, &model_dir)
.expect_err("overflowing external tensor geometry must be rejected");
assert!(
matches!(error, LoaderError::GraphBuild(message) if message.contains("geometry overflows"))
);
drop(store);
std::fs::remove_file(path).expect("remove external data");
}
#[test]
fn unreferenced_external_bytes_reports_an_orphaned_blob_prefix() {
let stamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("clock")
.as_nanos();
let model_dir = std::env::current_dir().expect("cwd").join("target");
std::fs::create_dir_all(&model_dir).expect("create target");
let file_name = format!("orphaned-prefix-blob-{}-{stamp}.bin", std::process::id());
let path = model_dir.join(&file_name);
const BLOB_LEN: usize = 256;
const LIVE_LEN: usize = 128;
std::fs::write(&path, vec![0u8; BLOB_LEN]).expect("write external data");
let initializer = TensorProto {
name: "live".to_string(),
data_type: DataType::Uint8.to_onnx(),
dims: vec![LIVE_LEN as i64],
external_data: vec![
crate::proto::onnx::StringStringEntryProto {
key: "location".to_string(),
value: file_name,
},
crate::proto::onnx::StringStringEntryProto {
key: "offset".to_string(),
value: LIVE_LEN.to_string(),
},
crate::proto::onnx::StringStringEntryProto {
key: "length".to_string(),
value: LIVE_LEN.to_string(),
},
],
data_location: tensor_proto::DataLocation::External as i32,
..Default::default()
};
let mut store = WeightStore::new();
resolve_initializer(&mut store, &initializer, &model_dir).expect("resolve live tensor");
let model = ModelProto {
graph: Some(crate::proto::onnx::GraphProto {
initializer: vec![initializer],
..Default::default()
}),
..Default::default()
};
let referenced = referenced_weight_bytes(&model);
assert_eq!(referenced.external, LIVE_LEN as u64);
assert_eq!(store.mapped_external_bytes(), BLOB_LEN as u64);
assert_eq!(
store.unreferenced_external_bytes(&referenced),
(BLOB_LEN - LIVE_LEN) as u64,
"the orphaned prefix must be reported, not rounded away"
);
drop(store);
std::fs::remove_file(path).expect("remove external data");
}
#[test]
fn unreferenced_external_bytes_is_zero_for_a_fully_referenced_blob() {
let stamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("clock")
.as_nanos();
let model_dir = std::env::current_dir().expect("cwd").join("target");
std::fs::create_dir_all(&model_dir).expect("create target");
let file_name = format!("packed-blob-{}-{stamp}.bin", std::process::id());
let path = model_dir.join(&file_name);
const BLOB_LEN: usize = 128;
std::fs::write(&path, vec![0u8; BLOB_LEN]).expect("write external data");
let initializer = TensorProto {
name: "live".to_string(),
data_type: DataType::Uint8.to_onnx(),
dims: vec![BLOB_LEN as i64],
external_data: vec![
crate::proto::onnx::StringStringEntryProto {
key: "location".to_string(),
value: file_name,
},
crate::proto::onnx::StringStringEntryProto {
key: "offset".to_string(),
value: "0".to_string(),
},
crate::proto::onnx::StringStringEntryProto {
key: "length".to_string(),
value: BLOB_LEN.to_string(),
},
],
data_location: tensor_proto::DataLocation::External as i32,
..Default::default()
};
let mut store = WeightStore::new();
resolve_initializer(&mut store, &initializer, &model_dir).expect("resolve live tensor");
let model = ModelProto {
graph: Some(crate::proto::onnx::GraphProto {
initializer: vec![initializer],
..Default::default()
}),
..Default::default()
};
assert_eq!(
store.unreferenced_external_bytes(&referenced_weight_bytes(&model)),
0
);
drop(store);
std::fs::remove_file(path).expect("remove external data");
}
#[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")
));
}
#[test]
fn qmoe_expert_tensor_layout_rejects_non_rank3_dims() {
assert_eq!(qmoe_expert_tensor_layout(4, 32, 2, &[8, 16]), None);
assert_eq!(qmoe_expert_tensor_layout(4, 32, 2, &[8, 16, 4, 2]), None);
assert_eq!(qmoe_expert_tensor_layout(4, 32, 2, &[]), None);
}
#[test]
fn qmoe_expert_tensor_layout_populates_expert_major_fields_for_rank3_dims() {
let layout = qmoe_expert_tensor_layout(4, 32, 3, &[8, 16, 24])
.expect("rank-3 dims must derive a layout");
assert_eq!(layout.version, 1);
assert_eq!(layout.experts, 8);
assert_eq!(layout.rows_per_expert, 16);
assert_eq!(layout.storage_elements_per_row, 24);
assert_eq!(layout.order, ExpertStorageOrder::ExpertMajor);
assert_eq!(
layout.quantization,
Some(ExpertQuantization {
bits: 4,
block_size: 32,
blocks_per_row: 3,
})
);
}
#[test]
fn qmoe_expert_tensor_layout_classifies_as_pageable_and_partitions_the_bank() {
let stamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("clock")
.as_nanos();
let dir = std::env::current_dir().expect("cwd").join("target");
std::fs::create_dir_all(&dir).expect("create target");
let path = dir.join(format!(
"qmoe-expert-region-catalog-{}-{stamp}.bin",
std::process::id()
));
let tensor_len = 4 * 16 * 24;
std::fs::write(&path, vec![0u8; tensor_len]).expect("write external data");
let weight = WeightRef::External {
path,
offset: 0,
length: tensor_len,
dtype: DataType::Uint8,
dims: vec![4, 16, 24],
};
let layout = qmoe_expert_tensor_layout(4, 32, 3, weight.dims())
.expect("rank-3 dims must derive a layout");
let catalog = WeightRegionCatalog::classify(&weight, layout);
assert!(catalog.is_pageable());
let mut expected_offset = 0usize;
let per_expert_len = 16 * 24;
for expert in 0..4 {
let range = catalog
.relative_range(expert)
.unwrap_or_else(|| panic!("expert {expert} must have a region"));
assert_eq!(range.start, expected_offset);
assert_eq!(range.end, expected_offset + per_expert_len);
expected_offset = range.end;
}
assert_eq!(expected_offset, tensor_len);
assert!(catalog.region(4).is_none());
std::fs::remove_file(weight_path_for_test(&catalog)).ok();
}
fn weight_path_for_test(catalog: &WeightRegionCatalog) -> PathBuf {
catalog.path.clone().unwrap_or_default()
}
#[test]
fn qmoe_expert_tensor_layout_rejects_inline_tensor_with_reason() {
let inline = WeightRef::Inline(TensorData::from_raw(
DataType::Uint8,
vec![2, 4, 8],
vec![0u8; 2 * 4 * 8],
));
let layout = qmoe_expert_tensor_layout(4, 32, 1, inline.dims())
.expect("rank-3 dims must derive a layout");
let catalog = WeightRegionCatalog::classify(&inline, layout);
assert!(!catalog.is_pageable());
assert_eq!(
catalog.pageability(),
&Pageability::NonPageable(NonPageableReason::InlineTensor)
);
}
}