use crate::{ModuleSnapshot, SafetensorsStore};
use burn_nn::LinearConfig;
#[test]
fn shape_mismatch_errors() {
let device = Default::default();
let module = LinearConfig::new(2, 2).with_bias(true).init(&device);
let mut save_store = SafetensorsStore::from_bytes(None);
module.save_into(&mut save_store).unwrap();
let mut incompatible_module = LinearConfig::new(3, 3).with_bias(true).init(&device);
let mut load_store = SafetensorsStore::from_bytes(None).validate(false); if let SafetensorsStore::Memory(ref mut p) = load_store
&& let SafetensorsStore::Memory(ref p_save) = save_store
{
let data_arc = p_save.data().unwrap();
p.set_data(data_arc.as_ref().clone());
}
let result = incompatible_module.load_from(&mut load_store).unwrap();
assert!(!result.errors.is_empty());
let mut load_store_with_validation = SafetensorsStore::from_bytes(None).validate(true);
if let SafetensorsStore::Memory(ref mut p) = load_store_with_validation
&& let SafetensorsStore::Memory(ref p_save) = save_store
{
let data_arc = p_save.data().unwrap();
p.set_data(data_arc.as_ref().clone());
}
let validation_result = incompatible_module.load_from(&mut load_store_with_validation);
assert!(validation_result.is_err());
}
#[test]
fn a_failing_tensor_does_not_unwind_out_of_collect_from() {
use crate::{ModuleAdapter, ModuleContext, bridge};
use alloc::boxed::Box;
use burn_pack::Tensor as PackTensor;
#[derive(Clone)]
struct PanickingAdapter;
impl ModuleAdapter for PanickingAdapter {
fn adapt(&self, tensor: PackTensor, _ctx: ModuleContext<'_>) -> PackTensor {
bridge::deferred(
tensor.name.clone(),
tensor.dtype,
tensor.shape.clone(),
None,
|| panic!("device readback panicked"),
)
}
fn clone_box(&self) -> Box<dyn ModuleAdapter> {
Box::new(self.clone())
}
}
let device = Default::default();
let module = LinearConfig::new(2, 2).init(&device);
let mut store = SafetensorsStore::from_bytes(None).with_to_adapter(PanickingAdapter);
let err = module
.save_into(&mut store)
.expect_err("a panicking provider must be returned, not unwound");
let message = alloc::format!("{err}");
assert!(
message.contains("tensor '") && message.contains("device readback panicked"),
"the error should name the tensor and carry the cause, got: {message}"
);
}