use alloc::collections::BTreeMap;
use alloc::string::String;
use alloc::vec::Vec;
use burn_std::DType;
use byteorder::{ByteOrder, LittleEndian};
use serde::{Deserialize, Serialize};
pub const MAGIC_NUMBER: u32 = 0x4255524E;
pub const FORMAT_VERSION: u16 = 0x0001;
pub const MAGIC_SIZE: usize = 4;
pub const VERSION_SIZE: usize = 2;
pub const METADATA_SIZE_FIELD_SIZE: usize = 4;
pub const HEADER_SIZE: usize = MAGIC_SIZE + VERSION_SIZE + METADATA_SIZE_FIELD_SIZE;
pub const TENSOR_ALIGNMENT: u64 = 256;
#[inline]
pub fn aligned_data_section_start(metadata_size: usize) -> usize {
let unaligned_start = (HEADER_SIZE + metadata_size) as u64;
(unaligned_start.div_ceil(TENSOR_ALIGNMENT) * TENSOR_ALIGNMENT) as usize
}
pub const MAX_METADATA_SIZE: u32 = 100 * 1024 * 1024;
#[cfg(target_pointer_width = "32")]
pub const MAX_TENSOR_SIZE: usize = 2 * 1024 * 1024 * 1024;
#[cfg(not(target_pointer_width = "32"))]
pub const MAX_TENSOR_SIZE: usize = 10 * 1024 * 1024 * 1024;
pub const MAX_TENSOR_COUNT: usize = 100_000;
pub const MAX_CBOR_RECURSION_DEPTH: usize = 128;
#[cfg(feature = "std")]
pub const MAX_FILE_SIZE: u64 = 100 * 1024 * 1024 * 1024;
pub const fn magic_range() -> core::ops::Range<usize> {
let start = 0;
let end = start + MAGIC_SIZE;
start..end
}
pub const fn version_range() -> core::ops::Range<usize> {
let start = MAGIC_SIZE;
let end = start + VERSION_SIZE;
start..end
}
pub const fn metadata_size_range() -> core::ops::Range<usize> {
let start = MAGIC_SIZE + VERSION_SIZE;
let end = start + METADATA_SIZE_FIELD_SIZE;
start..end
}
const _: () = assert!(MAGIC_SIZE + VERSION_SIZE + METADATA_SIZE_FIELD_SIZE == HEADER_SIZE);
#[derive(Debug, Clone, Copy)]
pub struct Header {
pub magic: u32,
pub version: u16,
pub metadata_size: u32,
}
impl Header {
pub fn new(metadata_size: u32) -> Self {
Self {
magic: MAGIC_NUMBER,
version: FORMAT_VERSION,
metadata_size,
}
}
pub fn into_bytes(self) -> [u8; HEADER_SIZE] {
let mut bytes = [0u8; HEADER_SIZE];
LittleEndian::write_u32(&mut bytes[magic_range()], self.magic);
LittleEndian::write_u16(&mut bytes[version_range()], self.version);
LittleEndian::write_u32(&mut bytes[metadata_size_range()], self.metadata_size);
bytes
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, Error> {
if bytes.len() < HEADER_SIZE {
return Err(Error::InvalidHeader);
}
let magic = LittleEndian::read_u32(&bytes[magic_range()]);
if magic != MAGIC_NUMBER {
return Err(Error::InvalidMagicNumber);
}
let version = LittleEndian::read_u16(&bytes[version_range()]);
let metadata_size = LittleEndian::read_u32(&bytes[metadata_size_range()]);
Ok(Self {
magic,
version,
metadata_size,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub enum Scalar {
Int(i64),
UInt(u64),
Float(f64),
Bool(bool),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ScalarConversionError;
impl core::fmt::Display for ScalarConversionError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "scalar value does not fit the requested type")
}
}
impl core::error::Error for ScalarConversionError {}
macro_rules! impl_scalar_int {
($($t:ty => $variant:ident),* $(,)?) => {
$(
impl From<$t> for Scalar {
fn from(value: $t) -> Self {
Scalar::$variant(value as _)
}
}
impl TryFrom<Scalar> for $t {
type Error = ScalarConversionError;
fn try_from(scalar: Scalar) -> Result<Self, Self::Error> {
match scalar {
Scalar::Int(v) => v.try_into().map_err(|_| ScalarConversionError),
Scalar::UInt(v) => v.try_into().map_err(|_| ScalarConversionError),
_ => Err(ScalarConversionError),
}
}
}
)*
};
}
impl_scalar_int!(
i8 => Int, i16 => Int, i32 => Int, i64 => Int, isize => Int,
u8 => UInt, u16 => UInt, u32 => UInt, u64 => UInt, usize => UInt,
);
impl From<f64> for Scalar {
fn from(value: f64) -> Self {
Scalar::Float(value)
}
}
impl From<f32> for Scalar {
fn from(value: f32) -> Self {
Scalar::Float(value as f64)
}
}
impl From<bool> for Scalar {
fn from(value: bool) -> Self {
Scalar::Bool(value)
}
}
impl TryFrom<Scalar> for f64 {
type Error = ScalarConversionError;
fn try_from(scalar: Scalar) -> Result<Self, Self::Error> {
match scalar {
Scalar::Float(v) => Ok(v),
Scalar::Int(v) => Ok(v as f64),
Scalar::UInt(v) => Ok(v as f64),
_ => Err(ScalarConversionError),
}
}
}
impl TryFrom<Scalar> for f32 {
type Error = ScalarConversionError;
fn try_from(scalar: Scalar) -> Result<Self, Self::Error> {
match scalar {
Scalar::Float(v) => Ok(v as f32),
Scalar::Int(v) => Ok(v as f32),
Scalar::UInt(v) => Ok(v as f32),
_ => Err(ScalarConversionError),
}
}
}
impl TryFrom<Scalar> for bool {
type Error = ScalarConversionError;
fn try_from(scalar: Scalar) -> Result<Self, Self::Error> {
match scalar {
Scalar::Bool(v) => Ok(v),
_ => Err(ScalarConversionError),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct Metadata {
pub tensors: BTreeMap<String, TensorDescriptor>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub metadata: BTreeMap<String, String>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub scalars: BTreeMap<String, Scalar>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct TensorDescriptor {
pub dtype: DType,
pub shape: Vec<u64>,
pub data_offsets: (u64, u64),
#[serde(default, skip_serializing_if = "Option::is_none")]
pub param_id: Option<u64>,
}
#[derive(Debug)]
pub enum Error {
InvalidHeader,
InvalidMagicNumber,
InvalidVersion,
MetadataSerializationError(String),
MetadataDeserializationError(String),
IoError(String),
TensorNotFound(String),
TensorBytesSizeMismatch(String),
ValidationError(String),
}
impl core::fmt::Display for Error {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Error::InvalidHeader => write!(f, "Invalid header: insufficient bytes"),
Error::InvalidMagicNumber => write!(f, "Invalid magic number"),
Error::InvalidVersion => write!(f, "Unsupported version"),
Error::MetadataSerializationError(e) => {
write!(f, "Metadata serialization error: {}", e)
}
Error::MetadataDeserializationError(e) => {
write!(f, "Metadata deserialization error: {}", e)
}
Error::IoError(e) => write!(f, "I/O error: {}", e),
Error::TensorNotFound(name) => write!(f, "Tensor not found: {}", name),
Error::TensorBytesSizeMismatch(e) => {
write!(f, "Tensor bytes size mismatch: {}", e)
}
Error::ValidationError(e) => write!(f, "Validation error: {}", e),
}
}
}
impl core::error::Error for Error {}
#[cfg(test)]
mod scalar_tests {
use super::*;
#[test]
fn int_round_trips_through_checked_conversion() {
assert_eq!(i32::try_from(Scalar::from(-5i32)).unwrap(), -5);
assert_eq!(u8::try_from(Scalar::from(200u8)).unwrap(), 200);
assert_eq!(usize::try_from(Scalar::from(42usize)).unwrap(), 42);
}
#[test]
fn out_of_range_int_conversion_is_rejected() {
let big = Scalar::from(5_000_000_000u64);
assert!(i32::try_from(big).is_err());
assert!(u32::try_from(Scalar::from(-1i32)).is_err());
assert!(u8::try_from(Scalar::from(300u32)).is_err());
}
#[test]
fn float_and_bool_variant_mismatches_are_rejected() {
assert!(i64::try_from(Scalar::Float(1.5)).is_err());
assert!(bool::try_from(Scalar::Int(1)).is_err());
assert!(f64::try_from(Scalar::Bool(true)).is_err());
}
#[test]
fn float_accepts_int_variants_symmetrically() {
assert_eq!(f64::try_from(Scalar::Int(3)).unwrap(), 3.0);
assert_eq!(f32::try_from(Scalar::Int(3)).unwrap(), 3.0);
assert_eq!(f64::try_from(Scalar::Float(2.5)).unwrap(), 2.5);
assert_eq!(f32::try_from(Scalar::Float(2.5)).unwrap(), 2.5);
}
}