use alloc::format;
use alloc::string::String;
#[cfg(feature = "std")]
use alloc::string::ToString;
use burn_pack::{Error as PackError, Tensor as PackTensor};
use burn_core::tensor::kind::Basic;
use burn_core::tensor::quantization::quantized_data_len;
use burn_core::tensor::{DType, Shape, Tensor, TensorData};
pub fn data_len(dtype: DType, shape: &Shape) -> usize {
match dtype {
DType::QFloat(scheme) => quantized_data_len(&scheme, shape),
_ => shape.iter().product::<usize>() * dtype.size(),
}
}
#[cfg(feature = "std")]
fn guarded(f: impl Fn() -> Result<TensorData, PackError>) -> Result<TensorData, PackError> {
std::panic::catch_unwind(std::panic::AssertUnwindSafe(f)).unwrap_or_else(|payload| {
let cause = payload
.downcast_ref::<&str>()
.map(|s| (*s).to_string())
.or_else(|| payload.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "unknown panic payload".to_string());
Err(PackError::ValidationError(format!(
"panic while producing tensor data: {cause}"
)))
})
}
#[cfg(not(feature = "std"))]
fn guarded(f: impl Fn() -> Result<TensorData, PackError>) -> Result<TensorData, PackError> {
f()
}
#[cfg(target_has_atomic = "ptr")]
pub trait MaybeSendSync: Send + Sync {}
#[cfg(target_has_atomic = "ptr")]
impl<T: Send + Sync> MaybeSendSync for T {}
#[cfg(not(target_has_atomic = "ptr"))]
pub trait MaybeSendSync {}
#[cfg(not(target_has_atomic = "ptr"))]
impl<T> MaybeSendSync for T {}
pub fn deferred(
name: String,
dtype: DType,
shape: Shape,
param_id: Option<u64>,
data_fn: impl Fn() -> Result<TensorData, PackError> + MaybeSendSync + 'static,
) -> PackTensor {
let byte_len = data_len(dtype, &shape);
let declared = shape.clone();
PackTensor::deferred(name, dtype, shape, param_id, byte_len, move || {
let data = guarded(&data_fn)?;
if data.dtype != dtype || data.shape != declared {
return Err(PackError::ValidationError(format!(
"provider produced {:?} {:?}, but the tensor declared {:?} {:?}",
data.dtype, data.shape, dtype, declared
)));
}
Ok(data.bytes)
})
}
pub fn from_tensor<const D: usize, K: Basic + 'static>(
tensor: &Tensor<D, K>,
name: String,
param_id: Option<u64>,
) -> PackTensor {
let (dtype, shape) = (tensor.dtype(), tensor.shape());
let tensor = tensor.clone();
deferred(name, dtype, shape, param_id, move || Ok(tensor.to_data()))
}
pub fn from_data(data: TensorData, name: String, param_id: Option<u64>) -> PackTensor {
PackTensor::new(name, data.dtype, data.shape.clone(), param_id, data.bytes)
}
pub fn to_data(tensor: &PackTensor) -> Result<TensorData, PackError> {
Ok(TensorData::from_bytes(
tensor.to_bytes()?,
tensor.shape.clone(),
tensor.dtype,
))
}
pub fn into_data(tensor: PackTensor) -> Result<TensorData, PackError> {
let (_, dtype, shape, _, bytes) = tensor.into_parts()?;
Ok(TensorData::from_bytes(bytes, shape, dtype))
}
pub fn map_data(
tensor: PackTensor,
name: String,
dtype: DType,
shape: Shape,
f: impl Fn(TensorData) -> TensorData + MaybeSendSync + 'static,
) -> PackTensor {
let param_id = tensor.param_id;
deferred(name, dtype, shape, param_id, move || {
to_data(&tensor).map(&f)
})
}
#[cfg(all(test, feature = "std"))]
mod tests {
use super::*;
use alloc::format;
use alloc::string::ToString;
use alloc::vec;
use burn_core::tensor::quantization::{QuantScheme, QuantStore, QuantValue, ScaleDtype};
use burn_core::tensor::{Bool, Device, Distribution, Int, shape};
fn assert_byte_len_matches(tensor: PackTensor) {
let (declared, name) = (tensor.byte_len(), tensor.name.clone());
match to_data(&tensor) {
Ok(data) => assert_eq!(
declared,
data.bytes.len(),
"byte_len disagrees with the materialized bytes for {name}"
),
Err(e) => panic!("byte_len disagrees with the materialized bytes for {name}: {e}"),
}
}
#[test]
fn data_len_covers_every_dtype_family() {
let device = Device::default();
let floats = Tensor::<2>::from_data([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]], &device);
let ints = Tensor::<2, Int>::from_data([[1, 2], [3, 4]], &device);
let bools = Tensor::<2, Bool>::from_data([[true, false], [false, true]], &device);
assert_byte_len_matches(from_tensor(&floats, "float".to_string(), None));
assert_byte_len_matches(from_tensor(&ints, "int".to_string(), None));
assert_byte_len_matches(from_tensor(&bools, "bool".to_string(), None));
}
#[test]
fn data_len_matches_materialized_bytes_when_quantized() {
let device = Device::default();
let schemes = [
("q8", QuantScheme::default()),
("q4", QuantScheme::default().with_value(QuantValue::Q4S)),
("q2", QuantScheme::default().with_value(QuantValue::Q2S)),
];
let shapes = [shape![32, 32], shape![3, 3], shape![5, 5], shape![2, 3]];
for (name, scheme) in schemes {
for shape in &shapes {
let tensor = Tensor::<2>::random(shape.clone(), Distribution::Default, &device)
.quantize_dynamic(&scheme);
assert_byte_len_matches(from_tensor(&tensor, format!("{name}-{shape:?}"), None));
}
}
}
#[test]
fn data_len_matches_quantized_bytes() {
let base = QuantScheme::default()
.with_value(QuantValue::Q8S)
.with_store(QuantStore::Native);
for (values, block) in [(8usize, 4usize), (6, 3), (10, 5)] {
let scales = vec![0.5f32; values / block];
let one_level = base.per_block([block as u8], ScaleDtype::UE4M3);
let two_level = base
.per_block([block as u8], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F32);
for (scheme, global) in [(one_level, None), (two_level, Some(3.0f32))] {
let data =
TensorData::quantized(vec![0i8; values], [values], scheme, &scales, global);
assert_eq!(
data_len(data.dtype, &data.shape),
data.bytes.len(),
"predicted size disagrees with the written bytes for {values} values \
in blocks of {block}, {scheme:?}"
);
}
}
}
#[test]
fn data_len_packs_per_line_for_packed_stores() {
let scheme = QuantScheme::default().with_value(QuantValue::Q4S);
let packed = |shape: Shape| data_len(DType::QFloat(scheme), &shape);
assert_eq!(packed(shape![3, 3]), 12 + 4);
assert_eq!(packed(shape![2, 5, 9]), 20 * 4 + 4);
assert_eq!(packed(shape![4, 8]), 4 * 4 + 4);
}
#[test]
fn a_panicking_provider_becomes_an_error() {
let tensor = deferred("weight".to_string(), DType::F32, shape![2, 2], None, || {
panic!("device readback panicked")
});
let err = to_data(&tensor).expect_err("a panicking provider must not unwind");
assert!(
matches!(&err, PackError::ValidationError(m) if m.contains("panic")),
"expected a validation error naming the panic, got {err:?}"
);
}
#[test]
fn a_write_time_failure_keeps_its_class_and_names_its_tensor() {
let tensor = deferred(
"encoder.weight".to_string(),
DType::F32,
shape![1],
None,
|| panic!("device readback panicked"),
);
let err = burn_pack::Writer::new(vec![tensor])
.into_bytes()
.expect_err("a panicking provider must fail the write");
assert!(
matches!(&err, PackError::ValidationError(m)
if m.contains("tensor 'encoder.weight'") && m.contains("device readback panicked")),
"expected a named ValidationError carrying the panic message, got {err:?}"
);
}
#[test]
fn a_provider_error_passes_through() {
let tensor = deferred("weight".to_string(), DType::F32, shape![2, 2], None, || {
Err(PackError::IoError("simulated IO error".to_string()))
});
let err = to_data(&tensor).unwrap_err();
assert!(
matches!(&err, PackError::IoError(m) if m == "simulated IO error"),
"expected the provider's own error, got {err:?}"
);
}
#[test]
fn deferred_runs_its_provider_only_when_asked() {
use core::sync::atomic::{AtomicUsize, Ordering};
use alloc::sync::Arc;
let calls = Arc::new(AtomicUsize::new(0));
let counter = calls.clone();
let data = TensorData::from([1.0f32, 2.0, 3.0, 4.0]);
let tensor = deferred(
"weight".to_string(),
DType::F32,
shape![4],
None,
move || {
counter.fetch_add(1, Ordering::Relaxed);
Ok(data.clone())
},
);
assert_eq!(tensor.byte_len(), 16);
assert_eq!(tensor.shape, shape![4]);
assert_eq!(calls.load(Ordering::Relaxed), 0);
to_data(&tensor).unwrap();
assert_eq!(calls.load(Ordering::Relaxed), 1);
}
#[test]
fn from_tensor_covers_every_kind() {
let device = Device::default();
let floats = from_tensor(
&Tensor::<2>::from_data([[1.0, 2.0], [3.0, 4.0]], &device),
"float".to_string(),
Some(7),
);
assert_eq!(floats.name, "float");
assert_eq!(floats.shape, shape![2, 2]);
assert_eq!(floats.param_id, Some(7));
assert_eq!(to_data(&floats).unwrap().shape, shape![2, 2]);
let ints = from_tensor(
&Tensor::<2, Int>::from_data([[1, 2], [3, 4]], &device),
"int".to_string(),
None,
);
assert_eq!(ints.dtype, device.settings().int_dtype.into());
let bools = from_tensor(
&Tensor::<2, Bool>::from_data([[true, false], [false, true]], &device),
"bool".to_string(),
None,
);
assert_eq!(to_data(&bools).unwrap().shape, shape![2, 2]);
}
#[test]
fn a_provider_contradicting_its_declared_dtype_is_refused() {
let tensor = deferred("weight".to_string(), DType::F32, shape![4], None, || {
Ok(TensorData::from([1i32, 2, 3, 4]))
});
let err = to_data(&tensor).unwrap_err();
assert!(
matches!(&err, PackError::ValidationError(m) if m.contains("I32") && m.contains("F32")),
"expected a validation error naming both dtypes, got {err:?}"
);
}
#[test]
fn a_provider_contradicting_its_declared_shape_is_refused() {
let tensor = deferred("weight".to_string(), DType::F32, shape![2, 3], None, || {
Ok(TensorData::new(vec![0.0f32; 6], shape![3, 2]))
});
assert!(matches!(
to_data(&tensor).unwrap_err(),
PackError::ValidationError(_)
));
}
#[test]
fn into_data_matches_to_data() {
let device = Device::default();
let tensor = from_tensor(
&Tensor::<2>::from_data([[1.0, 2.0], [3.0, 4.0]], &device),
"weight".to_string(),
None,
);
let borrowed = to_data(&tensor).unwrap();
let taken = into_data(tensor).unwrap();
assert_eq!(borrowed.shape, taken.shape);
assert_eq!(borrowed.dtype, taken.dtype);
assert_eq!(borrowed.bytes.to_vec(), taken.bytes.to_vec());
}
#[test]
fn map_data_declares_the_transformed_length() {
let device = Device::default();
let source = from_tensor(
&Tensor::<2>::from_data([[1.0, 2.0], [3.0, 4.0]], &device),
"weight".to_string(),
None,
);
assert_eq!(source.byte_len(), 16);
let shape = source.shape.clone();
let cast = map_data(source, "weight".to_string(), DType::F16, shape, |data| {
data.convert_dtype(DType::F16)
});
assert_eq!(cast.byte_len(), 8);
assert_byte_len_matches(cast);
}
}