tensor4all_tensorbackend/
lib.rs1#![warn(missing_docs)]
2#[cfg(feature = "global-defaults")]
14mod any_scalar;
16#[cfg(feature = "global-defaults")]
17mod backend;
19#[cfg(feature = "explicit-context")]
20mod context;
22#[cfg(feature = "tenferro-cuda")]
23mod cuda;
25#[cfg(feature = "explicit-context")]
26mod logical_tensor;
28#[cfg(feature = "global-defaults")]
29mod matrix;
31#[cfg(feature = "global-defaults")]
32mod memory;
34#[cfg(feature = "global-defaults")]
35mod storage;
37#[cfg(feature = "global-defaults")]
38pub(crate) mod tenferro_bridge;
39#[cfg(feature = "global-defaults")]
40mod tensor_element;
42
43#[cfg(feature = "global-defaults")]
44pub use any_scalar::BackendScalar;
45#[cfg(feature = "global-defaults")]
46pub use backend::{
47 full_piv_lu_backend, full_piv_lu_matrix, qr_backend, solve_backend, solve_matrix,
48 solve_matrix_owned, svd_backend, triangular_solve_backend, triangular_solve_matrix,
49 triangular_solve_matrix_owned, BackendLinalgError, BackendLinalgScalar, FullPivLuMatrixResult,
50 FullPivLuResult, FullPivLuScalar, MatrixSolveScalar, MatrixTriangularSolveScalar, SvdResult,
51};
52#[cfg(feature = "global-defaults")]
53pub use context::{default_eager_ctx, with_default_backend, EagerContextError};
54#[cfg(feature = "explicit-context")]
55pub use context::{CpuExecutionContext, CpuExecutionContextError};
56#[cfg(feature = "tenferro-cuda")]
57pub use cuda::{CudaExecutionContext, CudaExecutionContextError, CUDA_ORDINAL};
58#[cfg(feature = "explicit-context")]
59pub use logical_tensor::{LogicalTensor, LogicalTensorData, LogicalTensorError};
60#[cfg(feature = "global-defaults")]
61pub use matrix::{
62 batched_mat_mul_same_shape, batched_mat_mul_same_shape_owned, from_vec2d,
63 hermitian_eigendecomposition, hermitian_exponential_first_column, lowest_hermitian_eigenpair,
64 mat_mul, mat_mul_owned, submatrix, submatrix_argmax, swap_cols, swap_rows, transpose,
65 try_from_vec2d, BlasMul, HermitianEigenError, HermitianEigenScalar,
66 HermitianEigendecomposition, HermitianEigenpair, Matrix, MatrixScalar, MatrixShapeError,
67 MatrixTensorConversionError,
68};
69#[cfg(feature = "global-defaults")]
70pub use memory::{release_process_allocator_cached_memory, AllocatorPressureRelief};
71#[cfg(feature = "global-defaults")]
72pub use storage::{
73 contract_storage, make_mut_storage, min_dim, Storage, StorageError, StorageKind, StorageResult,
74 StorageScalar, StructuredStorage, SumFromStorage,
75};
76#[cfg(feature = "global-defaults")]
77pub use tenferro_bridge::{
78 axpby_native_tensor, axpby_storage_native, conj_native_tensor, contract_native_tensor,
79 contract_storage_native, dense_native_tensor_from_col_major, diag_native_tensor_from_col_major,
80 einsum_native_tensor_reads, einsum_native_tensors, einsum_native_tensors_owned,
81 native_tensor_primal_to_dense_col_major, native_tensor_primal_to_diag,
82 native_tensor_primal_to_storage, outer_product_native_tensor, outer_product_storage_native,
83 permute_native_tensor, permute_storage_native, print_and_reset_native_einsum_profile,
84 qr_native_tensor, reset_native_einsum_profile, reshape_col_major_native_tensor,
85 scale_native_tensor, scale_storage_native, storage_payload_native_read_input,
86 storage_to_native_tensor, sum_native_tensor, svd_native_tensor, tangent_native_tensor,
87 BridgeError, NativeTensorReadInput,
88};
89#[cfg(feature = "global-defaults")]
90pub use tensor_element::TensorElement;
91
92#[cfg(feature = "global-defaults")]
94pub(crate) fn require_invariant<T, E: std::fmt::Display>(
95 result: std::result::Result<T, E>,
96 context: &str,
97) -> T {
98 let valid = result.is_ok();
99 if let Err(error) = &result {
100 assert!(valid, "{context}: {error}");
101 }
102 match result {
103 Ok(value) => value,
104 Err(_) => loop {
105 std::hint::spin_loop();
106 },
107 }
108}
109
110#[cfg(all(test, feature = "global-defaults"))]
111mod invariant_tests {
112 use super::require_invariant;
113
114 #[test]
115 fn require_invariant_returns_success_and_reports_failure_context() {
116 assert_eq!(require_invariant::<_, &str>(Ok(7), "valid state"), 7);
117
118 let failure = std::panic::catch_unwind(|| {
119 require_invariant::<(), _>(Err("broken state"), "tensor invariant")
120 });
121 let message = failure
122 .unwrap_err()
123 .downcast::<String>()
124 .map(|message| *message)
125 .unwrap_or_default();
126 assert!(message.contains("tensor invariant: broken state"));
127 }
128}