tenferro_tensor/validate/mod.rs
1//! Validation helpers shared across backends and exec layers.
2//!
3//! # Examples
4//!
5//! ```rust
6//! use tenferro_tensor::validate::validate_nonsingular_u;
7//! use tenferro_tensor::{Tensor, TypedTensor};
8//!
9//! let t = Tensor::from_typed::<f64>(TypedTensor::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 1.0]).unwrap());
10//! assert!(validate_nonsingular_u(&t).is_ok());
11//! ```
12
13use num_complex::{Complex32, Complex64};
14
15use crate::{
16 DType, DotGeneralConfig, Error, ErrorKind, Result, ShapeMismatch, Tensor, TensorScalar,
17 TypedTensor, ValidationError,
18};
19
20/// Domain-specific reasons reported by triangular-factor validation.
21///
22/// The outer tensor error classifies these failures as numerical or unsupported
23/// while retaining this value as its typed source.
24///
25/// # Examples
26///
27/// ```rust
28/// use tenferro_tensor::validate::DiagonalError;
29///
30/// let error = DiagonalError::SingularOrNonFinite {
31/// index: 1,
32/// };
33/// assert!(error.to_string().contains("position [1,1]"));
34/// ```
35#[derive(Debug, thiserror::Error)]
36pub enum DiagonalError {
37 #[error("singular or non-finite diagonal at position [{index},{index}]")]
38 SingularOrNonFinite { index: usize },
39 #[error("singular or non-finite diagonal at batch {batch}, position [{index},{index}]")]
40 BatchedSingularOrNonFinite { batch: usize, index: usize },
41 #[error("triangular solve does not support dtype {dtype:?}")]
42 UnsupportedDType { dtype: DType },
43}
44
45/// Promote two dtypes according to tenferro's public dtype-promotion lattice.
46///
47/// # Examples
48///
49/// ```rust
50/// use tenferro_tensor::validate::promote_dtype;
51/// use tenferro_tensor::DType;
52///
53/// assert_eq!(promote_dtype(DType::I32, DType::F32), DType::F64);
54/// ```
55pub fn promote_dtype(lhs: DType, rhs: DType) -> DType {
56 // The lattice belongs to the scalar set that declares the members, and is
57 // derived from each member's declared kind, rank, and width. A set that
58 // declares a different set of scalars promotes within that set instead of
59 // using this one.
60 <crate::DefaultScalars as crate::ScalarSet>::promote(lhs, rhs)
61}
62
63/// Return whether public `convert` may change `from` into `to`.
64///
65/// Checked conversion follows the same dtype lattice as implicit promotion.
66/// Use explicit `cast` for value-changing projections outside this lattice.
67///
68/// # Examples
69///
70/// ```rust
71/// use tenferro_tensor::validate::can_convert_dtype;
72/// use tenferro_tensor::DType;
73///
74/// assert!(can_convert_dtype(DType::F32, DType::F64));
75/// assert!(!can_convert_dtype(DType::F64, DType::I32));
76/// ```
77pub fn can_convert_dtype(from: DType, to: DType) -> bool {
78 promote_dtype(from, to) == to
79}
80
81/// Validate a public checked dtype conversion.
82///
83/// # Examples
84///
85/// ```rust
86/// use tenferro_tensor::validate::validate_convert_dtype;
87/// use tenferro_tensor::DType;
88///
89/// assert!(validate_convert_dtype("convert", DType::F32, DType::F64).is_ok());
90/// assert!(validate_convert_dtype("convert", DType::C64, DType::F64).is_err());
91/// ```
92/// # Errors
93///
94/// Returns [`crate::Error::UnsupportedDTypeConversion`] when the requested
95/// conversion is outside the checked promotion lattice. Use an explicit cast
96/// for lossy projections.
97pub fn validate_convert_dtype(op: &'static str, from: DType, to: DType) -> Result<()> {
98 if can_convert_dtype(from, to) {
99 return Ok(());
100 }
101
102 Err(Error::unsupported_dtype_conversion(
103 op,
104 from,
105 to,
106 "checked convert only accepts conversions allowed by dtype promotion; use explicit cast for lossy dtype projection",
107 ))
108}
109
110/// Compute a shape product with overflow reported as a typed tensor error.
111///
112/// # Examples
113///
114/// ```rust
115/// use tenferro_tensor::validate::checked_shape_product;
116///
117/// assert_eq!(checked_shape_product("zeros", "shape", &[2, 3])?, 6);
118/// # Ok::<(), tenferro_tensor::Error>(())
119/// ```
120/// # Errors
121///
122/// Returns [`crate::Error::Validation`] containing
123/// [`tenferro_tensor_core::ValidationError::InvalidArgument`] when the
124/// product of `shape` exceeds `usize::MAX`; `role` identifies the shape-like
125/// argument in the diagnostic.
126pub fn checked_shape_product(
127 op: &'static str,
128 role: &'static str,
129 shape: &[usize],
130) -> Result<usize> {
131 if shape.contains(&0) {
132 return Ok(0);
133 }
134 shape
135 .iter()
136 .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
137 .ok_or_else(|| {
138 Error::invalid_argument(op, role, format!("product overflows for shape {shape:?}"))
139 })
140}
141
142/// Validate a full permutation for a tensor rank.
143///
144/// # Examples
145///
146/// ```rust
147/// use tenferro_tensor::validate::validate_permutation_axes;
148///
149/// validate_permutation_axes("transpose", 2, &[1, 0])?;
150/// # Ok::<(), tenferro_tensor::Error>(())
151/// ```
152/// # Errors
153///
154/// Returns [`crate::Error::Validation`] with `RankMismatch`,
155/// `AxisOutOfBounds`, or `DuplicateAxis` as appropriate.
156pub fn validate_permutation_axes(op: &'static str, rank: usize, perm: &[usize]) -> Result<()> {
157 if perm.len() != rank {
158 return Err(Error::validation(
159 op,
160 ValidationError::RankMismatch {
161 expected: rank,
162 actual: perm.len(),
163 },
164 ));
165 }
166
167 let mut seen = vec![false; rank];
168 for &axis in perm {
169 if axis >= rank {
170 return Err(Error::validation(
171 op,
172 ValidationError::AxisOutOfBounds { axis, rank },
173 ));
174 }
175 if seen[axis] {
176 return Err(Error::validation(
177 op,
178 ValidationError::DuplicateAxis {
179 axis,
180 role: "permutation",
181 },
182 ));
183 }
184 seen[axis] = true;
185 }
186 Ok(())
187}
188
189/// Validate a subset of axes for a tensor rank.
190///
191/// # Examples
192///
193/// ```rust
194/// use tenferro_tensor::validate::validate_unique_axes;
195///
196/// validate_unique_axes("reduce_sum", "axis", 3, &[0, 2])?;
197/// assert!(validate_unique_axes("reduce_sum", "axis", 2, &[2]).is_err());
198/// assert!(validate_unique_axes("reduce_sum", "axis", 2, &[0, 0]).is_err());
199/// # Ok::<(), tenferro_tensor::Error>(())
200/// ```
201/// # Errors
202///
203/// Returns [`crate::Error::Validation`] with `AxisOutOfBounds` for an invalid
204/// axis or `DuplicateAxis` when an axis occurs more than once.
205pub fn validate_unique_axes(
206 op: &'static str,
207 role: &'static str,
208 rank: usize,
209 axes: &[usize],
210) -> Result<()> {
211 let mut seen = vec![false; rank];
212 for &axis in axes {
213 if axis >= rank {
214 return Err(Error::validation(
215 op,
216 ValidationError::AxisOutOfBounds { axis, rank },
217 ));
218 }
219 if seen[axis] {
220 return Err(Error::validation(
221 op,
222 ValidationError::DuplicateAxis { axis, role },
223 ));
224 }
225 seen[axis] = true;
226 }
227 Ok(())
228}
229
230/// Validate rank-2 matrix multiplication shapes and return its dot-general config.
231///
232/// # Examples
233///
234/// ```rust
235/// use tenferro_tensor::validate::matmul_config_for_shapes;
236///
237/// let config = matmul_config_for_shapes("matmul", &[2, 3], &[3, 4])?;
238/// assert_eq!(config.lhs_contracting_dims.as_slice(), &[1]);
239/// # Ok::<(), tenferro_tensor::Error>(())
240/// ```
241/// # Errors
242///
243/// Returns [`crate::Error::Validation`] with `RankMismatch` for a non-matrix
244/// input or `ShapeMismatch` when the contracting dimensions differ.
245pub fn matmul_config_for_shapes(
246 op: &'static str,
247 lhs_shape: &[usize],
248 rhs_shape: &[usize],
249) -> Result<DotGeneralConfig> {
250 if lhs_shape.len() != 2 {
251 return Err(Error::validation(
252 op,
253 ValidationError::RankMismatch {
254 expected: 2,
255 actual: lhs_shape.len(),
256 },
257 ));
258 }
259 if rhs_shape.len() != 2 {
260 return Err(Error::validation(
261 op,
262 ValidationError::RankMismatch {
263 expected: 2,
264 actual: rhs_shape.len(),
265 },
266 ));
267 }
268 if lhs_shape[1] != rhs_shape[0] {
269 return Err(Error::validation(
270 op,
271 ShapeMismatch::IncompatibleShapes {
272 lhs: lhs_shape.to_vec().into(),
273 rhs: rhs_shape.to_vec().into(),
274 }
275 .into(),
276 ));
277 }
278
279 Ok(DotGeneralConfig {
280 lhs_contracting_dims: [1].as_slice().into(),
281 rhs_contracting_dims: [0].as_slice().into(),
282 lhs_batch_dims: [].as_slice().into(),
283 rhs_batch_dims: [].as_slice().into(),
284 })
285}
286
287/// Trait for detecting singular or non-finite diagonal entries.
288///
289/// Implemented for `f32`, `f64`, `Complex32`, and `Complex64`.
290/// A value is considered singular if it is zero, NaN, infinite,
291/// or (for complex types) if either component is non-finite.
292pub trait DiagSingularity {
293 /// Returns `true` if the value is singular or non-finite.
294 fn is_singular_or_nonfinite(&self) -> bool;
295}
296
297macro_rules! impl_diag_singularity_float {
298 ($($t:ty),* $(,)?) => {
299 $(
300 impl DiagSingularity for $t {
301 fn is_singular_or_nonfinite(&self) -> bool {
302 !self.is_finite() || *self == 0.0
303 }
304 }
305 )*
306 };
307}
308
309impl_diag_singularity_float!(f64, f32);
310
311macro_rules! impl_diag_singularity_complex {
312 ($($t:ty),* $(,)?) => {
313 $(
314 impl DiagSingularity for $t {
315 fn is_singular_or_nonfinite(&self) -> bool {
316 // Why not `norm_sqr() == 0`: squaring a representable tiny
317 // component can underflow and relabel a nonzero pivot as zero.
318 !self.re.is_finite()
319 || !self.im.is_finite()
320 || (self.re == 0.0 && self.im == 0.0)
321 }
322 }
323 )*
324 };
325}
326
327impl_diag_singularity_complex!(Complex64, Complex32);
328
329/// Checks that every diagonal element of a (possibly batched) upper-triangular
330/// factor is non-singular and finite.
331///
332/// Iterates over all batch slices and inspects the diagonal entries
333/// `data[i + i * rows]` for `i` in `0..min(rows, cols)`. Returns
334/// a numerical [`crate::Error::Extension`] carrying a typed
335/// [`DiagonalError::SingularOrNonFinite`] source for the first offending entry,
336/// or [`ValidationError::RankMismatch`] wrapped in [`Error::Validation`] when
337/// `t` has rank less than two.
338///
339/// # Examples
340///
341/// ```rust
342/// use tenferro_tensor::validate::check_singular_diagonal;
343/// use tenferro_tensor::TypedTensor;
344///
345/// let t = TypedTensor::from_vec_col_major(vec![2, 2], vec![1.0f32, 0.0, 0.0, 2.0]).unwrap();
346/// assert!(check_singular_diagonal(&t).is_ok());
347/// ```
348/// # Errors
349///
350/// Returns [`crate::Error::Validation`] with `RankMismatch` for a non-matrix
351/// tensor, a numerical [`crate::Error::Extension`] with a typed
352/// [`DiagonalError::SingularOrNonFinite`] source for a singular or non-finite
353/// diagonal, or an unsupported [`crate::Error::Extension`] with a typed
354/// [`DiagonalError::UnsupportedDType`] source for integer and boolean inputs.
355pub fn check_singular_diagonal<T: DiagSingularity + TensorScalar + std::fmt::Debug>(
356 t: &TypedTensor<T>,
357) -> Result<()> {
358 if t.shape().len() < 2 {
359 return Err(Error::validation(
360 "solve",
361 ValidationError::RankMismatch {
362 expected: 2,
363 actual: t.shape().len(),
364 },
365 ));
366 }
367 let rows = t.shape()[0];
368 let cols = t.shape()[1];
369 let n = rows.min(cols);
370 let batch_total = checked_shape_product("solve", "batch shape", &t.shape()[2..])?;
371 let slice_size = checked_shape_product("solve", "matrix shape", &t.shape()[..2])?;
372 let data = t.host_data()?;
373 for batch_idx in 0..batch_total {
374 let batch = &data[batch_idx * slice_size..(batch_idx + 1) * slice_size];
375 for i in 0..n {
376 let diag = batch[i + i * rows];
377 if diag.is_singular_or_nonfinite() {
378 return Err(Error::extension(
379 "solve",
380 "tensor-validation",
381 ErrorKind::NumericalFailure,
382 if batch_total > 1 {
383 DiagonalError::BatchedSingularOrNonFinite {
384 batch: batch_idx,
385 index: i,
386 }
387 } else {
388 DiagonalError::SingularOrNonFinite { index: i }
389 },
390 ));
391 }
392 }
393 }
394 Ok(())
395}
396
397/// Validates that the upper-triangular factor `u` of a matrix decomposition
398/// has no singular (zero) or non-finite diagonal entries.
399///
400/// Dispatches to [`check_singular_diagonal`] after unpacking the concrete
401/// tensor variant. Returns `Ok(())` when all diagonal entries are valid.
402///
403/// # Examples
404///
405/// ```rust
406/// use tenferro_tensor::validate::validate_nonsingular_u;
407/// use tenferro_tensor::{Tensor, TypedTensor};
408///
409/// let t = Tensor::from_typed::<f64>(TypedTensor::from_vec_col_major(vec![2, 2], vec![1.0, 0.0, 0.0, 1.0]).unwrap());
410/// assert!(validate_nonsingular_u(&t).is_ok());
411/// ```
412/// # Errors
413///
414/// Returns [`crate::Error::Validation`] with the applicable typed shape, rank,
415/// axis, dtype, or argument source when validation fails. Singular or
416/// non-finite diagonal checks return [`crate::Error::BackendFailure`].
417pub fn validate_nonsingular_u(u: &Tensor) -> Result<()> {
418 match u.dtype() {
419 DType::F64 => check_singular_diagonal(
420 u.as_typed::<f64>()
421 .ok_or_else(|| unsupported_diagonal_dtype(u))?,
422 ),
423 DType::F32 => check_singular_diagonal(
424 u.as_typed::<f32>()
425 .ok_or_else(|| unsupported_diagonal_dtype(u))?,
426 ),
427 DType::C64 => check_singular_diagonal(
428 u.as_typed::<Complex64>()
429 .ok_or_else(|| unsupported_diagonal_dtype(u))?,
430 ),
431 DType::C32 => check_singular_diagonal(
432 u.as_typed::<Complex32>()
433 .ok_or_else(|| unsupported_diagonal_dtype(u))?,
434 ),
435 DType::I32 | DType::I64 | DType::Bool | DType::External(_) => {
436 Err(unsupported_diagonal_dtype(u))
437 }
438 }
439}
440
441/// The refusal this module produces for a dtype it cannot validate a diagonal in.
442///
443/// Returning it from an accessor is the same refusal the wildcard arm produced; a
444/// caller reaches that accessor from a match on `u.dtype()`, so it is unreachable in
445/// practice rather than a caller mistake.
446fn unsupported_diagonal_dtype(u: &Tensor) -> Error {
447 Error::extension(
448 "solve",
449 "tensor-validation",
450 ErrorKind::Unsupported,
451 DiagonalError::UnsupportedDType { dtype: u.dtype() },
452 )
453}
454
455#[cfg(test)]
456mod tests;