use std::error::Error as _;
use std::panic::{catch_unwind, AssertUnwindSafe};
use num_complex::{Complex32, Complex64};
use tenferro_cpu::CpuBackend;
use tenferro_linalg::{LinalgBackend, TensorLinalgExt};
use tenferro_tensor::{
BackendCachedDot, BackendRuntimeCache, BackendSession, BackendSessionHost,
BackendStorageHandle, CompareDir, DType, DotGeneralConfig, Error, ErrorKind, GatherConfig,
MemoryKind, PadConfig, Placement, ScatterConfig, SliceConfig, StorageBuffer, Tensor,
TensorAnalytic, TensorBackend, TensorBuffer, TensorDeviceTransfer, TensorDot,
TensorElementwise, TensorFusion, TensorIndexing, TensorRead, TensorReduction, TensorStructural,
TensorView, TensorWrite, TypedTensor, TypedTensorView, ValidationError,
};
use super::support;
fn f64_tensor(shape: Vec<usize>, data: Vec<f64>) -> Tensor {
Tensor::F64(TypedTensor::from_vec_col_major(shape, data).unwrap())
}
fn f32_tensor(shape: Vec<usize>, data: Vec<f32>) -> Tensor {
Tensor::F32(TypedTensor::from_vec_col_major(shape, data).unwrap())
}
fn c64_tensor(shape: Vec<usize>, data: Vec<Complex64>) -> Tensor {
Tensor::C64(TypedTensor::from_vec_col_major(shape, data).unwrap())
}
fn c32_tensor(shape: Vec<usize>, data: Vec<Complex32>) -> Tensor {
Tensor::C32(TypedTensor::from_vec_col_major(shape, data).unwrap())
}
fn i32_tensor(shape: Vec<usize>, data: Vec<i32>) -> Tensor {
Tensor::I32(TypedTensor::from_vec_col_major(shape, data).unwrap())
}
fn f64_values(tensor: &Tensor) -> Vec<f64> {
match tensor {
Tensor::F64(tensor) => tensor.host_data().unwrap().to_vec(),
other => panic!("expected F64 tensor, got {:?}", other.dtype()),
}
}
fn c64_values(tensor: &Tensor) -> Vec<Complex64> {
match tensor {
Tensor::C64(tensor) => tensor.host_data().unwrap().to_vec(),
other => panic!("expected C64 tensor, got {:?}", other.dtype()),
}
}
fn opaque_backend_placement() -> Placement {
Placement {
memory_kind: MemoryKind::Device,
device: None,
cpu_affinity: None,
}
}
fn backend_f64_tensor(shape: Vec<usize>, handle_id: u64) -> Tensor {
let len = shape.iter().product();
Tensor::F64(
TypedTensor::<f64>::from_buffer_col_major(
shape,
StorageBuffer::Backend(Box::new(BackendStorageHandle::<f64>::new_with_len(
handle_id, len,
))),
opaque_backend_placement(),
)
.unwrap(),
)
}
fn assert_backend_download_error<T>(result: tenferro_tensor::Result<T>, expected_op: &'static str) {
let err = match result {
Ok(_) => panic!("expected {expected_op} to reject backend buffer"),
Err(err) => err,
};
assert!(matches!(
err,
Error::RuntimeState {
op,
ref message,
} if op == expected_op && message.contains("download")
));
}
fn assert_no_panic_backend_download_error<T>(
expected_op: &'static str,
f: impl FnOnce() -> tenferro_tensor::Result<T>,
) {
let result = catch_unwind(AssertUnwindSafe(f));
assert!(
result.is_ok(),
"{expected_op} should return Err for backend buffers, not panic"
);
assert_backend_download_error(result.unwrap(), expected_op);
}
#[test]
fn default_svd_read_returns_explicit_backend_boundary_error() {
struct DefaultOnlyLinalgBackend {
eig_values_result: Option<Tensor>,
eig_values_calls: usize,
}
macro_rules! panic_backend_methods {
($($name:ident($($arg:ident : $argty:ty),*) -> $ret:ty;)+) => {
$(
fn $name(&mut self, $($arg: $argty),*) -> $ret {
$(let _ = &$arg;)*
panic!(concat!(stringify!($name), " should not be called by this test"))
}
)+
};
}
impl BackendRuntimeCache for DefaultOnlyLinalgBackend {
type RuntimeCache = ();
}
impl TensorElementwise for DefaultOnlyLinalgBackend {
panic_backend_methods! {
add(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
sub(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
mul(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
neg(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
conj(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
div(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
abs(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
sign(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
maximum(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
minimum(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
compare(lhs: &Tensor, rhs: &Tensor, dir: &CompareDir) -> tenferro_tensor::Result<Tensor>;
select(pred: &Tensor, on_true: &Tensor, on_false: &Tensor) -> tenferro_tensor::Result<Tensor>;
clamp(input: &Tensor, lower: &Tensor, upper: &Tensor) -> tenferro_tensor::Result<Tensor>;
}
}
impl TensorAnalytic for DefaultOnlyLinalgBackend {
panic_backend_methods! {
exp(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
log(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
sin(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
cos(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
tanh(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
sqrt(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
rsqrt(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
pow(lhs: &Tensor, rhs: &Tensor) -> tenferro_tensor::Result<Tensor>;
expm1(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
log1p(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
}
}
impl TensorStructural for DefaultOnlyLinalgBackend {
panic_backend_methods! {
transpose(input: &Tensor, perm: &[usize]) -> tenferro_tensor::Result<Tensor>;
reshape(input: &Tensor, shape: &[usize]) -> tenferro_tensor::Result<Tensor>;
broadcast_in_dim(input: &Tensor, shape: &[usize], dims: &[usize]) -> tenferro_tensor::Result<Tensor>;
cast(input: &Tensor, to: DType) -> tenferro_tensor::Result<Tensor>;
extract_diagonal(input: &Tensor, axis_a: usize, axis_b: usize) -> tenferro_tensor::Result<Tensor>;
embed_diagonal(input: &Tensor, axis_a: usize, axis_b: usize) -> tenferro_tensor::Result<Tensor>;
tril(input: &Tensor, k: i64) -> tenferro_tensor::Result<Tensor>;
triu(input: &Tensor, k: i64) -> tenferro_tensor::Result<Tensor>;
}
fn copy_read_into(
&mut self,
src: TensorRead<'_>,
dst: TensorWrite<'_>,
) -> tenferro_tensor::Result<()> {
let Some(Tensor::F64(src)) = src.as_tensor() else {
panic!("the default solve_read_into test uses an owned f64 source")
};
let TensorWrite::Tensor(Tensor::F64(dst)) = dst else {
panic!("the default solve_read_into test uses an owned f64 destination")
};
dst.host_data_mut()?.copy_from_slice(src.host_data()?);
Ok(())
}
}
impl TensorReduction for DefaultOnlyLinalgBackend {
panic_backend_methods! {
reduce_sum(input: &Tensor, axes: &[usize]) -> tenferro_tensor::Result<Tensor>;
reduce_prod(input: &Tensor, axes: &[usize]) -> tenferro_tensor::Result<Tensor>;
reduce_max(input: &Tensor, axes: &[usize]) -> tenferro_tensor::Result<Tensor>;
reduce_min(input: &Tensor, axes: &[usize]) -> tenferro_tensor::Result<Tensor>;
}
}
impl TensorIndexing for DefaultOnlyLinalgBackend {
panic_backend_methods! {
gather(operand: &Tensor, start_indices: &Tensor, config: &GatherConfig) -> tenferro_tensor::Result<Tensor>;
scatter(operand: &Tensor, scatter_indices: &Tensor, updates: &Tensor, config: &ScatterConfig) -> tenferro_tensor::Result<Tensor>;
slice(input: &Tensor, config: &SliceConfig) -> tenferro_tensor::Result<Tensor>;
dynamic_slice(input: &Tensor, starts: &Tensor, slice_sizes: &[usize]) -> tenferro_tensor::Result<Tensor>;
dynamic_update_slice(operand: &Tensor, update: &Tensor, starts: &Tensor) -> tenferro_tensor::Result<Tensor>;
pad(input: &Tensor, config: &PadConfig) -> tenferro_tensor::Result<Tensor>;
concatenate(inputs: &[&Tensor], axis: usize) -> tenferro_tensor::Result<Tensor>;
reverse(input: &Tensor, axes: &[usize]) -> tenferro_tensor::Result<Tensor>;
}
}
impl TensorDot for DefaultOnlyLinalgBackend {
panic_backend_methods! {
dot_general(lhs: &Tensor, rhs: &Tensor, config: &DotGeneralConfig) -> tenferro_tensor::Result<Tensor>;
}
}
impl TensorFusion for DefaultOnlyLinalgBackend {}
impl TensorBuffer for DefaultOnlyLinalgBackend {}
impl TensorDeviceTransfer for DefaultOnlyLinalgBackend {
fn download_to_host(&mut self, _tensor: TensorRead<'_>) -> tenferro_tensor::Result<Tensor> {
Err(Error::unsupported(
"DefaultOnlyLinalgBackend::download_to_host",
"test backend does not transfer tensors",
))
}
fn upload_host_tensor(
&mut self,
_tensor: TensorRead<'_>,
) -> tenferro_tensor::Result<Tensor> {
Err(Error::unsupported(
"DefaultOnlyLinalgBackend::upload_host_tensor",
"test backend does not transfer tensors",
))
}
}
impl BackendCachedDot for DefaultOnlyLinalgBackend {}
#[doc(hidden)]
struct DefaultOnlyLinalgBackendSessionMarker;
impl BackendSession for DefaultOnlyLinalgBackend {
fn session_type_id(&self) -> std::any::TypeId {
std::any::TypeId::of::<DefaultOnlyLinalgBackendSessionMarker>()
}
unsafe fn session_data_mut(&mut self) -> *mut () {
self as *mut Self as *mut ()
}
}
impl BackendSessionHost for DefaultOnlyLinalgBackend {}
impl TensorBackend for DefaultOnlyLinalgBackend {}
impl LinalgBackend for DefaultOnlyLinalgBackend {
panic_backend_methods! {
cholesky(input: &Tensor) -> tenferro_tensor::Result<Tensor>;
triangular_solve(a: &Tensor, b: &Tensor, left_side: bool, lower: bool, transpose_a: bool, unit_diagonal: bool) -> tenferro_tensor::Result<Tensor>;
lu(input: &Tensor) -> tenferro_tensor::Result<Vec<Tensor>>;
full_piv_lu(input: &Tensor) -> tenferro_tensor::Result<Vec<Tensor>>;
full_piv_lu_solve(a: &Tensor, b: &Tensor, transpose_a: bool) -> tenferro_tensor::Result<Tensor>;
svd(input: &Tensor) -> tenferro_tensor::Result<Vec<Tensor>>;
qr(input: &Tensor) -> tenferro_tensor::Result<Vec<Tensor>>;
eigh(input: &Tensor) -> tenferro_tensor::Result<Vec<Tensor>>;
eig(input: &Tensor) -> tenferro_tensor::Result<Vec<Tensor>>;
solve(a: &Tensor, b: &Tensor) -> tenferro_tensor::Result<Tensor>;
}
fn solve_read(
&mut self,
_a: TensorRead<'_>,
b: TensorRead<'_>,
) -> tenferro_tensor::Result<Tensor> {
Ok(Tensor::F64(
TypedTensor::from_vec_col_major(b.shape().to_vec(), vec![2.0, 3.0]).unwrap(),
))
}
fn eig_values(&mut self, _input: &Tensor) -> tenferro_tensor::Result<Tensor> {
self.eig_values_calls += 1;
Ok(self
.eig_values_result
.take()
.expect("eig_values result configured by the test"))
}
}
let input =
TypedTensor::<f64>::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 2.0]).unwrap();
let expected_eig_values = Tensor::C64(
TypedTensor::from_vec_col_major(
vec![2],
vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)],
)
.unwrap(),
);
let mut backend = DefaultOnlyLinalgBackend {
eig_values_result: Some(expected_eig_values),
eig_values_calls: 0,
};
let eig_values = Tensor::F64(input.duplicate().unwrap())
.eigvals(&mut backend)
.unwrap();
assert_eq!(backend.eig_values_calls, 1);
assert_eq!(
c64_values(&eig_values),
vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)]
);
let err = backend
.lu_factor(&Tensor::F64(input.duplicate().unwrap()))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "lu_factor",
ref message,
} if message.contains("does not implement")
));
let err = backend
.svd_values(&Tensor::F64(input.duplicate().unwrap()))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "svd_values",
ref message,
} if message.contains("does not implement")
));
let owned_input = Tensor::F64(input.duplicate().unwrap());
let rhs = Tensor::from_vec_col_major(vec![2, 1], vec![7.0_f64, 11.0]).unwrap();
let mut output = Tensor::from_vec_col_major(vec![2, 1], vec![-1.0_f64; 2]).unwrap();
backend
.solve_read_into(
TensorRead::from_tensor(&owned_input),
TensorRead::from_tensor(&rhs),
TensorWrite::from_tensor(&mut output),
)
.unwrap();
assert_eq!(output.as_slice::<f64>().unwrap(), &[2.0, 3.0]);
let mut bad_dtype = Tensor::from_vec_col_major(vec![2, 1], vec![0_i32; 2]).unwrap();
let err = backend
.solve_read_into(
TensorRead::from_tensor(&owned_input),
TensorRead::from_tensor(&rhs),
TensorWrite::from_tensor(&mut bad_dtype),
)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "solve_read_into",
source: ValidationError::DTypeMismatch { .. },
}
));
let mut wrong_placement = backend_f64_tensor(vec![2, 1], 180);
let err = backend
.solve_read_into(
TensorRead::from_tensor(&owned_input),
TensorRead::from_tensor(&rhs),
TensorWrite::from_tensor(&mut wrong_placement),
)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "solve_read_into",
source: ValidationError::InvalidArgument {
argument: "out",
..
},
}
));
let err = backend
.svd_read(tenferro_tensor::TensorRead::from_tensor(&owned_input))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "svd",
ref message,
} if message.contains("tensor reads")
));
let err = backend
.qr_read(TensorRead::from_view(TensorView::F64(input.as_view())))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "qr",
ref message,
} if message.contains("tensor reads")
));
let err = backend
.eigh_read(TensorRead::from_view(TensorView::F64(input.as_view())))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "eigh",
ref message,
} if message.contains("tensor reads")
));
let err = backend
.eigh_values(&Tensor::F64(input.duplicate().unwrap()))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "eigh_values",
ref message,
} if message.contains("does not implement")
));
let pivots = Tensor::I32(TypedTensor::from_vec_col_major(vec![2], vec![1, 2]).unwrap());
let err = backend
.lu_solve_prepared(
&Tensor::F64(input.duplicate().unwrap()),
&Tensor::F64(input.duplicate().unwrap()),
&pivots,
&Tensor::F64(input.duplicate().unwrap()),
false,
false,
)
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "lu_solve_prepared",
ref message,
} if message.contains("does not implement")
));
let err = backend
.cholesky_read(TensorRead::from_view(TensorView::F64(input.as_view())))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "cholesky",
ref message,
} if message.contains("tensor reads")
));
let err = backend
.lu_read(TensorRead::from_view(TensorView::F64(input.as_view())))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "lu",
ref message,
} if message.contains("tensor reads")
));
let err = backend
.full_piv_lu_read(TensorRead::from_view(TensorView::F64(input.as_view())))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "full_piv_lu",
ref message,
} if message.contains("tensor reads")
));
let err = backend
.eig_read(TensorRead::from_view(TensorView::F64(input.as_view())))
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "eig",
ref message,
} if message.contains("tensor reads")
));
let rhs =
Tensor::F64(TypedTensor::<f64>::from_vec_col_major(vec![2, 1], vec![1.0, 2.0]).unwrap());
let solved = backend
.solve_read(
TensorRead::from_tensor(&owned_input),
TensorRead::from_tensor(&rhs),
)
.unwrap();
assert_eq!(solved.as_slice::<f64>().unwrap(), &[2.0, 3.0]);
let err = backend
.triangular_solve_read(
TensorRead::from_tensor(&owned_input),
TensorRead::from_tensor(&rhs),
true,
true,
false,
false,
)
.unwrap_err();
assert!(matches!(
err,
Error::Unsupported {
op: "triangular_solve",
ref message,
} if message.contains("tensor reads")
));
}
#[test]
fn cpu_lu_solve_prepared_consumes_packed_factor_outputs() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let a = f64_tensor(vec![2, 2], vec![2.0, 0.0, 0.0, 3.0]);
let b = f64_tensor(vec![2, 1], vec![4.0, 9.0]);
let factors = backend.lu_factor(&a).unwrap();
let x = backend
.lu_solve_prepared(&a, &factors[0], &factors[1], &b, false, false)
.unwrap();
assert_eq!(f64_values(&x), vec![2.0, 3.0]);
let a = f64_tensor(vec![2, 2], vec![1.0, 0.0, 2.0, 3.0]);
let b = f64_tensor(vec![2, 1], vec![5.0, 31.0]);
let factors = backend.lu_factor(&a).unwrap();
let x = backend
.lu_solve_prepared(&a, &factors[0], &factors[1], &b, true, false)
.unwrap();
let values = f64_values(&x);
assert!((values[0] - 5.0).abs() < 1.0e-12);
assert!((values[1] - 7.0).abs() < 1.0e-12);
let a = c64_tensor(
vec![2, 2],
vec![
Complex64::new(1.0, 1.0),
Complex64::new(0.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(2.0, -1.0),
],
);
let b = c64_tensor(
vec![2, 1],
vec![Complex64::new(2.0, -2.0), Complex64::new(6.0, 3.0)],
);
let factors = backend.lu_factor(&a).unwrap();
let x = backend
.lu_solve_prepared(&a, &factors[0], &factors[1], &b, true, true)
.unwrap();
let values = c64_values(&x);
assert!((values[0] - Complex64::new(2.0, 0.0)).norm() < 1.0e-12);
assert!((values[1] - Complex64::new(3.0, 0.0)).norm() < 1.0e-12);
});
}
#[test]
fn cpu_lu_factor_covers_pivoted_real_and_complex_dtypes() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let a = f32_tensor(vec![2, 2], vec![0.0, 1.0, 1.0, 0.0]);
let factors = backend.lu_factor(&a).unwrap();
assert!(matches!(&factors[0], Tensor::F32(t) if t.shape() == [2, 2]));
assert!(matches!(&factors[1], Tensor::I32(t) if t.host_data().unwrap() == [2, 2]));
assert!(matches!(&factors[2], Tensor::F32(t) if t.host_data().unwrap() == [-1.0]));
let a = c32_tensor(
vec![2, 2],
vec![
Complex32::new(2.0, 0.0),
Complex32::new(0.0, 0.0),
Complex32::new(0.0, 0.0),
Complex32::new(3.0, 1.0),
],
);
let factors = backend.lu_factor(&a).unwrap();
assert!(matches!(&factors[0], Tensor::C32(t) if t.shape() == [2, 2]));
assert!(matches!(&factors[1], Tensor::I32(t) if t.host_data().unwrap() == [1, 2]));
assert!(
matches!(&factors[2], Tensor::C32(t) if t.host_data().unwrap() == [Complex32::new(1.0, 0.0)])
);
});
}
#[test]
fn cpu_values_only_decompositions_cover_real_complex_and_batched_inputs() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let s = backend
.svd_values(&f32_tensor(vec![2, 2], vec![3.0, 0.0, 0.0, 4.0]))
.unwrap();
assert!(matches!(s, Tensor::F32(ref t) if t.shape() == [2]));
let s = backend
.svd_values(&f64_tensor(
vec![2, 2, 2],
vec![3.0, 0.0, 0.0, 4.0, 5.0, 0.0, 0.0, 6.0],
))
.unwrap();
assert!(matches!(s, Tensor::F64(ref t) if t.shape() == [2, 2]));
let s = backend
.svd_values(&c32_tensor(
vec![2, 2],
vec![
Complex32::new(3.0, 0.0),
Complex32::new(0.0, 0.0),
Complex32::new(0.0, 0.0),
Complex32::new(4.0, 0.0),
],
))
.unwrap();
assert!(matches!(s, Tensor::F32(ref t) if t.shape() == [2]));
let s = backend
.svd_values(&c64_tensor(
vec![2, 2],
vec![
Complex64::new(3.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(4.0, 0.0),
],
))
.unwrap();
assert!(matches!(s, Tensor::F64(ref t) if t.shape() == [2]));
let values = backend
.eigh_values(&f32_tensor(vec![2, 2], vec![3.0, 0.0, 0.0, 4.0]))
.unwrap();
assert!(matches!(values, Tensor::F32(ref t) if t.shape() == [2]));
let values = backend
.eigh_values(&f64_tensor(
vec![2, 2, 2],
vec![3.0, 0.0, 0.0, 4.0, 5.0, 0.0, 0.0, 6.0],
))
.unwrap();
assert!(matches!(values, Tensor::F64(ref t) if t.shape() == [2, 2]));
let values = backend
.eigh_values(&c32_tensor(
vec![2, 2],
vec![
Complex32::new(3.0, 0.0),
Complex32::new(0.0, 0.0),
Complex32::new(0.0, 0.0),
Complex32::new(4.0, 0.0),
],
))
.unwrap();
assert!(matches!(values, Tensor::F32(ref t) if t.shape() == [2]));
let values = backend
.eigh_values(&c64_tensor(
vec![2, 2],
vec![
Complex64::new(3.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(4.0, 0.0),
],
))
.unwrap();
assert!(matches!(values, Tensor::F64(ref t) if t.shape() == [2]));
});
}
#[test]
fn cpu_lu_solve_prepared_restores_vector_rhs_and_validates_inputs() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let a = f64_tensor(vec![2, 2], vec![2.0, 0.0, 0.0, 4.0]);
let factors = backend.lu_factor(&a).unwrap();
let b = f64_tensor(vec![2], vec![6.0, 20.0]);
let x = backend
.lu_solve_prepared(&a, &factors[0], &factors[1], &b, false, false)
.unwrap();
assert_eq!(x.shape(), &[2]);
assert_eq!(f64_values(&x), vec![3.0, 5.0]);
let empty_a = f64_tensor(vec![0, 0], Vec::new());
let empty_b = f64_tensor(vec![0, 1], Vec::new());
let empty_pivots = i32_tensor(vec![0], Vec::new());
let x = backend
.lu_solve_prepared(&empty_a, &empty_a, &empty_pivots, &empty_b, false, false)
.unwrap();
assert_eq!(x.shape(), &[0, 1]);
let bad_pivots = f64_tensor(vec![2], vec![1.0, 2.0]);
let err = backend
.lu_solve_prepared(&a, &factors[0], &bad_pivots, &b, false, false)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "lu_solve_prepared",
source: ValidationError::DTypeMismatch { .. },
}
));
let bad_b = c64_tensor(
vec![2, 1],
vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)],
);
let err = backend
.lu_solve_prepared(&a, &factors[0], &factors[1], &bad_b, false, false)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "lu_solve_prepared",
source: ValidationError::DTypeMismatch { .. },
}
));
let bad_pivots = i32_tensor(vec![2], vec![0, 2]);
let err = backend
.lu_solve_prepared(&a, &a, &bad_pivots, &b, false, false)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "lu_solve_prepared",
source: ValidationError::InvalidArgument {
argument: "pivot",
..
},
}
));
});
}
#[test]
fn cpu_lu_solve_prepared_rejects_rank_less_than_two() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let a = f64_tensor(vec![2], vec![1.0, 2.0]);
let b = f64_tensor(vec![2], vec![1.0, 2.0]);
let pivots = i32_tensor(vec![2], vec![1, 2]);
let err = backend
.lu_solve_prepared(&a, &a, &pivots, &b, true, false)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "lu_solve_prepared",
source: ValidationError::RankMismatch { .. },
}
));
});
}
#[test]
fn cpu_linalg_rejects_backend_buffers_without_panicking_or_downloading() {
let a = backend_f64_tensor(vec![2, 2], 101);
let b = backend_f64_tensor(vec![2, 1], 102);
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
assert_no_panic_backend_download_error("cholesky", || backend.cholesky(&a));
assert_no_panic_backend_download_error("triangular_solve", || {
backend.triangular_solve(&a, &b, true, true, false, false)
});
assert_no_panic_backend_download_error("lu", || backend.lu(&a));
assert_no_panic_backend_download_error("full_piv_lu", || backend.full_piv_lu(&a));
assert_no_panic_backend_download_error("full_piv_lu_solve", || {
backend.full_piv_lu_solve(&a, &b, false)
});
assert_no_panic_backend_download_error("svd", || backend.svd(&a));
assert_no_panic_backend_download_error("qr", || backend.qr(&a));
assert_no_panic_backend_download_error("eigh", || backend.eigh(&a));
assert_no_panic_backend_download_error("eig", || backend.eig(&a));
assert_no_panic_backend_download_error("solve", || backend.solve(&a, &b));
assert_no_panic_backend_download_error("solve", || {
backend.solve_read(TensorRead::from_tensor(&a), TensorRead::from_tensor(&b))
});
assert_no_panic_backend_download_error("triangular_solve", || {
backend.triangular_solve_read(
TensorRead::from_tensor(&a),
TensorRead::from_tensor(&b),
true,
true,
false,
false,
)
});
let host_c64 = c64_tensor(
vec![2, 1],
vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)],
);
assert_no_panic_backend_download_error("solve", || {
backend.solve_read(
TensorRead::from_tensor(&a),
TensorRead::from_tensor(&host_c64),
)
});
});
}
#[test]
fn cpu_linalg_rejects_backend_rhs_before_zero_dim_fast_paths() {
let a = f64_tensor(vec![0, 0], Vec::new());
let b = backend_f64_tensor(vec![0, 1], 103);
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
assert_no_panic_backend_download_error("solve", || backend.solve(&a, &b));
assert_no_panic_backend_download_error("full_piv_lu_solve", || {
backend.full_piv_lu_solve(&a, &b, false)
});
});
}
#[test]
fn solve_rejects_invalid_dtype_pairs_before_zero_dim_fast_path() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let f64_a = f64_tensor(vec![0, 0], Vec::new());
let c64_b = c64_tensor(vec![0, 1], Vec::new());
let i32_a = i32_tensor(vec![0, 0], Vec::new());
let i32_b = i32_tensor(vec![0, 1], Vec::new());
let err = backend.solve(&f64_a, &c64_b).unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "solve",
source: ValidationError::DTypeMismatch { .. },
}
));
let err = backend
.solve_read(
TensorRead::from_tensor(&f64_a),
TensorRead::from_tensor(&c64_b),
)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "solve",
source: ValidationError::DTypeMismatch { .. },
}
));
let err = backend.solve(&i32_a, &i32_b).unwrap_err();
assert!(matches!(
err,
Error::Extension {
op: "solve",
family: tenferro_linalg::LINALG_EXTENSION_FAMILY_ID,
kind: tenferro_tensor::ErrorKind::Unsupported,
..
}
));
});
}
#[test]
fn full_piv_lu_solve_rejects_invalid_dtype_pairs_before_zero_dim_fast_path() {
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let f64_a = f64_tensor(vec![0, 0], Vec::new());
let c64_b = c64_tensor(vec![0, 1], Vec::new());
let i32_a = i32_tensor(vec![0, 0], Vec::new());
let i32_b = i32_tensor(vec![0, 1], Vec::new());
let err = backend
.full_piv_lu_solve(&f64_a, &c64_b, false)
.unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "full_piv_lu_solve",
source: ValidationError::DTypeMismatch { .. },
}
));
let err = backend
.full_piv_lu_solve(&i32_a, &i32_b, false)
.unwrap_err();
assert!(matches!(
err,
Error::Extension {
op: "full_piv_lu_solve",
family: tenferro_linalg::LINALG_EXTENSION_FAMILY_ID,
kind: tenferro_tensor::ErrorKind::Unsupported,
..
}
));
});
}
#[test]
fn cholesky_rejects_rank_less_than_two_even_when_zero_dim() {
let input = f64_tensor(vec![0], Vec::new());
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let result = catch_unwind(AssertUnwindSafe(|| backend.cholesky(&input)));
assert!(result.is_ok(), "cholesky should return Err, not panic");
let err = result.unwrap().unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "cholesky",
source: ValidationError::RankMismatch {
expected: 2,
actual: 1,
},
}
));
});
}
#[test]
fn solve_rejects_singular_matrix() {
let a = f64_tensor(vec![2, 2], vec![1.0, 2.0, 2.0, 4.0]);
let b = f64_tensor(vec![2, 1], vec![1.0, 2.0]);
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let err = backend.solve(&a, &b).unwrap_err();
assert_eq!(err.kind(), ErrorKind::NumericalFailure);
assert!(matches!(
err.source()
.and_then(|source| source.downcast_ref::<tenferro_linalg::Error>()),
Some(tenferro_linalg::Error::Singular { op: "solve" })
));
let err = backend
.solve_read(TensorRead::from_tensor(&a), TensorRead::from_tensor(&b))
.unwrap_err();
assert_eq!(err.kind(), ErrorKind::NumericalFailure);
assert!(matches!(
err.source()
.and_then(|source| source.downcast_ref::<tenferro_linalg::Error>()),
Some(tenferro_linalg::Error::Singular { op: "solve" })
));
});
}
#[test]
fn triangular_solve_rejects_batch_mismatch_without_backend_panic() {
let mut backend = CpuBackend::new();
let a = f64_tensor(vec![2, 2, 2], vec![1.0, 0.0, 0.0, 1.0, 2.0, 0.0, 0.0, 2.0]);
let b = f64_tensor(vec![2, 1, 3], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
support::with_cpu_linalg(&mut backend, |backend| {
let result = catch_unwind(AssertUnwindSafe(|| {
backend.triangular_solve(&a, &b, true, true, false, false)
}));
assert!(
result.is_ok(),
"triangular_solve should return Err on batch mismatch, not panic"
);
let err = result.unwrap().unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "triangular_solve",
source: ValidationError::ShapeMismatch(_),
}
));
});
}
#[test]
fn full_piv_lu_solve_rejects_batch_mismatch_without_backend_panic() {
let mut backend = CpuBackend::new();
let a = f64_tensor(vec![2, 2, 2], vec![1.0, 0.0, 0.0, 1.0, 2.0, 0.0, 0.0, 2.0]);
let b = f64_tensor(vec![2, 1, 3], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
support::with_cpu_linalg(&mut backend, |backend| {
let result = catch_unwind(AssertUnwindSafe(|| {
backend.full_piv_lu_solve(&a, &b, false)
}));
assert!(
result.is_ok(),
"full_piv_lu_solve should return Err on batch mismatch, not panic"
);
let err = result.unwrap().unwrap_err();
assert!(matches!(
err,
Error::Validation {
op: "full_piv_lu_solve",
source: ValidationError::ShapeMismatch(_),
}
));
});
}
#[test]
fn cpu_solve_read_covers_direct_vector_and_matrix_rhs_views() {
let a = f64_tensor(vec![2, 2], vec![2.0, 0.0, 0.0, 3.0]);
let vector = f64_tensor(vec![2], vec![4.0, 6.0]);
let mut backend = CpuBackend::new();
support::with_cpu_linalg(&mut backend, |backend| {
let vector_output = backend
.solve_read(
TensorRead::from_tensor(&a),
TensorRead::from_tensor(&vector),
)
.unwrap();
assert_eq!(f64_values(&vector_output), vec![2.0, 2.0]);
let mut storage = vec![-1.0_f64; 8];
storage[1] = 4.0;
storage[2] = 6.0;
storage[5] = 8.0;
storage[6] = 9.0;
let matrix_view = TypedTensorView::from_slice(vec![2, 2], vec![1, 4], 1, &storage).unwrap();
let matrix_output = backend
.solve_read(
TensorRead::from_tensor(&a),
TensorRead::from_view(TensorView::F64(matrix_view)),
)
.unwrap();
assert_eq!(f64_values(&matrix_output), vec![2.0, 2.0, 4.0, 3.0]);
});
}