Skip to main content

tensor4all_tensorbackend/
lib.rs

1#![warn(missing_docs)]
2//! Tensor storage and linear algebra backend for tensor4all.
3//!
4//! [`CpuExecutionContext`] is the canonical CPU integration path. It requires a
5//! caller-supplied backend and owns plain, graph, and eager-AD runtime state.
6//!
7//! ## Feature flags
8//!
9//! - `explicit-context`: explicit CPU execution and logical tensor transfer.
10//! - `global-defaults`: legacy process-global tensor operations.
11//! - `backend-tenferro` (default): compatibility alias for `global-defaults`.
12
13#[cfg(feature = "global-defaults")]
14/// Dynamic scalar types supporting f32, f64, Complex32, and Complex64.
15mod any_scalar;
16#[cfg(feature = "global-defaults")]
17/// Backend dispatch for dense linear algebra operations.
18mod backend;
19#[cfg(feature = "explicit-context")]
20/// Explicit and optional process-global tenferro execution helpers.
21mod context;
22#[cfg(feature = "tenferro-cuda")]
23/// Explicit visible-ordinal-0 CUDA execution and transfer boundaries.
24mod cuda;
25#[cfg(feature = "explicit-context")]
26/// Backend-free tensor snapshots for execution-domain transfer.
27mod logical_tensor;
28#[cfg(feature = "global-defaults")]
29/// Dense column-major matrix type and backend-backed matrix utilities.
30mod matrix;
31#[cfg(feature = "global-defaults")]
32/// Process-level memory pressure helpers.
33mod memory;
34#[cfg(feature = "global-defaults")]
35/// Tensor snapshot storage types and low-level dense/diagonal kernels.
36mod storage;
37#[cfg(feature = "global-defaults")]
38pub(crate) mod tenferro_bridge;
39#[cfg(feature = "global-defaults")]
40/// Supported public tensor element types and native constructor hooks.
41mod 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/// Extract a result whose error branch means validated internal state is inconsistent.
93#[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}