#[cfg(not(feature = "std"))]
use alloc::{vec, vec::Vec};
use crate::packed::{PackedMat, PackedMut, PackedRef, TriangularKind};
use crate::{Mat, MatMut, MatRef};
use oxiblas_core::scalar::Scalar;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DiagonalKind {
NonUnit,
Unit,
}
#[derive(Clone, Copy)]
pub struct TriangularView<'a, T: Scalar> {
inner: MatRef<'a, T>,
uplo: TriangularKind,
diag: DiagonalKind,
}
impl<'a, T: Scalar> TriangularView<'a, T> {
#[inline]
pub fn new(mat: MatRef<'a, T>, uplo: TriangularKind, diag: DiagonalKind) -> Self {
assert!(mat.is_square(), "Triangular matrix must be square");
TriangularView {
inner: mat,
uplo,
diag,
}
}
#[inline]
pub fn dim(&self) -> usize {
self.inner.nrows()
}
#[inline]
pub fn shape(&self) -> (usize, usize) {
self.inner.shape()
}
#[inline]
pub fn uplo(&self) -> TriangularKind {
self.uplo
}
#[inline]
pub fn diag(&self) -> DiagonalKind {
self.diag
}
#[inline]
pub fn in_triangle(&self, row: usize, col: usize) -> bool {
match self.uplo {
TriangularKind::Upper => row <= col,
TriangularKind::Lower => row >= col,
}
}
#[inline]
pub fn get(&self, row: usize, col: usize) -> Option<&T> {
if !self.in_triangle(row, col) {
return None;
}
if self.diag == DiagonalKind::Unit && row == col {
return None; }
self.inner.get(row, col)
}
#[inline]
pub fn as_inner(&self) -> MatRef<'a, T> {
self.inner
}
#[inline]
pub fn as_ptr(&self) -> *const T {
self.inner.as_ptr()
}
#[inline]
pub fn row_stride(&self) -> usize {
self.inner.row_stride()
}
pub fn to_dense(&self) -> Mat<T>
where
T: bytemuck::Zeroable,
{
let n = self.dim();
let mut mat = Mat::zeros(n, n);
for j in 0..n {
for i in 0..n {
if self.in_triangle(i, j) {
if self.diag == DiagonalKind::Unit && i == j {
mat[(i, j)] = T::one();
} else {
mat[(i, j)] = self.inner[(i, j)];
}
}
}
}
mat
}
pub fn to_packed(&self) -> PackedMat<T>
where
T: bytemuck::Zeroable,
{
PackedMat::from_dense(&self.inner, self.uplo)
}
}
pub struct TriangularViewMut<'a, T: Scalar> {
inner: MatMut<'a, T>,
uplo: TriangularKind,
diag: DiagonalKind,
}
impl<'a, T: Scalar> TriangularViewMut<'a, T> {
#[inline]
pub fn new(mat: MatMut<'a, T>, uplo: TriangularKind, diag: DiagonalKind) -> Self {
assert!(mat.is_square(), "Triangular matrix must be square");
TriangularViewMut {
inner: mat,
uplo,
diag,
}
}
#[inline]
pub fn dim(&self) -> usize {
self.inner.nrows()
}
#[inline]
pub fn shape(&self) -> (usize, usize) {
self.inner.shape()
}
#[inline]
pub fn uplo(&self) -> TriangularKind {
self.uplo
}
#[inline]
pub fn diag(&self) -> DiagonalKind {
self.diag
}
#[inline]
pub fn in_triangle(&self, row: usize, col: usize) -> bool {
match self.uplo {
TriangularKind::Upper => row <= col,
TriangularKind::Lower => row >= col,
}
}
#[inline]
pub fn get(&self, row: usize, col: usize) -> Option<&T> {
if !self.in_triangle(row, col) {
return None;
}
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.inner.get(row, col)
}
#[inline]
pub fn get_mut(&mut self, row: usize, col: usize) -> Option<&mut T> {
if !self.in_triangle(row, col) {
return None;
}
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.inner.get_mut(row, col)
}
#[inline]
pub fn set(&mut self, row: usize, col: usize, value: T) {
assert!(
self.in_triangle(row, col),
"Element outside stored triangle"
);
assert!(
!(self.diag == DiagonalKind::Unit && row == col),
"Cannot set diagonal of unit triangular matrix"
);
self.inner.set(row, col, value);
}
#[inline]
pub fn rb(&self) -> TriangularView<'_, T> {
TriangularView {
inner: self.inner.rb(),
uplo: self.uplo,
diag: self.diag,
}
}
#[inline]
pub fn rb_mut(&mut self) -> TriangularViewMut<'_, T> {
TriangularViewMut {
inner: self.inner.rb_mut(),
uplo: self.uplo,
diag: self.diag,
}
}
#[inline]
pub fn as_mut_ptr(&mut self) -> *mut T {
self.inner.as_mut_ptr()
}
pub fn fill(&mut self, value: T) {
let n = self.dim();
for j in 0..n {
for i in 0..n {
if self.in_triangle(i, j) {
if self.diag == DiagonalKind::Unit && i == j {
continue;
}
self.inner.set(i, j, value);
}
}
}
}
pub fn scale(&mut self, alpha: T) {
let n = self.dim();
for j in 0..n {
for i in 0..n {
if self.in_triangle(i, j) {
if self.diag == DiagonalKind::Unit && i == j {
continue;
}
if let Some(val) = self.inner.get(i, j) {
self.inner.set(i, j, *val * alpha);
}
}
}
}
}
pub fn zero_non_triangle(&mut self)
where
T: num_traits::Zero,
{
let n = self.dim();
for j in 0..n {
for i in 0..n {
if !self.in_triangle(i, j) {
self.inner.set(i, j, T::zero());
}
}
}
}
}
#[derive(Clone)]
pub struct TriangularMat<T: Scalar> {
packed: PackedMat<T>,
diag: DiagonalKind,
}
impl<T: Scalar> TriangularMat<T> {
pub fn zeros(n: usize, uplo: TriangularKind, diag: DiagonalKind) -> Self
where
T: bytemuck::Zeroable,
{
TriangularMat {
packed: PackedMat::zeros(n, uplo),
diag,
}
}
pub fn unit_zeros(n: usize, uplo: TriangularKind) -> Self
where
T: bytemuck::Zeroable,
{
Self::zeros(n, uplo, DiagonalKind::Unit)
}
#[inline]
pub fn from_packed(packed: PackedMat<T>, diag: DiagonalKind) -> Self {
TriangularMat { packed, diag }
}
pub fn from_dense(mat: &MatRef<'_, T>, uplo: TriangularKind, diag: DiagonalKind) -> Self
where
T: bytemuck::Zeroable,
{
TriangularMat {
packed: PackedMat::from_dense(mat, uplo),
diag,
}
}
#[inline]
pub fn dim(&self) -> usize {
self.packed.dim()
}
#[inline]
pub fn shape(&self) -> (usize, usize) {
let n = self.dim();
(n, n)
}
#[inline]
pub fn uplo(&self) -> TriangularKind {
self.packed.kind()
}
#[inline]
pub fn diag(&self) -> DiagonalKind {
self.diag
}
#[inline]
pub fn len(&self) -> usize {
self.packed.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.packed.is_empty()
}
#[inline]
pub fn in_triangle(&self, row: usize, col: usize) -> bool {
self.packed.packed_index(row, col).is_some()
}
#[inline]
pub fn get(&self, row: usize, col: usize) -> Option<&T> {
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.packed.get(row, col)
}
#[inline]
pub fn get_mut(&mut self, row: usize, col: usize) -> Option<&mut T> {
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.packed.get_mut(row, col)
}
#[inline]
pub fn set(&mut self, row: usize, col: usize, value: T) {
assert!(
!(self.diag == DiagonalKind::Unit && row == col),
"Cannot set diagonal of unit triangular matrix"
);
assert!(
self.in_triangle(row, col),
"Element outside stored triangle"
);
self.packed
.set(row, col, value)
.expect("in_triangle check above guarantees this index is valid");
}
#[inline]
pub fn as_ptr(&self) -> *const T {
self.packed.as_ptr()
}
#[inline]
pub fn as_mut_ptr(&mut self) -> *mut T {
self.packed.as_mut_ptr()
}
#[inline]
pub fn as_slice(&self) -> &[T] {
self.packed.as_slice()
}
#[inline]
pub fn as_slice_mut(&mut self) -> &mut [T] {
self.packed.as_slice_mut()
}
#[inline]
pub fn as_packed(&self) -> &PackedMat<T> {
&self.packed
}
#[inline]
pub fn as_packed_mut(&mut self) -> &mut PackedMat<T> {
&mut self.packed
}
pub fn to_dense(&self) -> Mat<T>
where
T: bytemuck::Zeroable,
{
let n = self.dim();
let mut mat = Mat::zeros(n, n);
for j in 0..n {
for i in 0..n {
if let Some(&val) = self.packed.get(i, j) {
if self.diag == DiagonalKind::Unit && i == j {
mat[(i, j)] = T::one();
} else {
mat[(i, j)] = val;
}
} else if self.diag == DiagonalKind::Unit && i == j {
mat[(i, j)] = T::one();
}
}
}
mat
}
pub fn diagonal(&self) -> Vec<T> {
let n = self.dim();
if self.diag == DiagonalKind::Unit {
vec![T::one(); n]
} else {
self.packed.diagonal()
}
}
pub fn set_diagonal(&mut self, diag: &[T]) {
assert!(
self.diag != DiagonalKind::Unit,
"Cannot set diagonal of unit triangular matrix"
);
self.packed.set_diagonal(diag);
}
pub fn fill(&mut self, value: T) {
if self.diag == DiagonalKind::Unit {
let n = self.dim();
for j in 0..n {
for i in 0..n {
if self.in_triangle(i, j) && i != j {
self.packed
.set(i, j, value)
.expect("in_triangle check above guarantees this index is valid");
}
}
}
} else {
self.packed.fill(value);
}
}
pub fn scale(&mut self, alpha: T) {
if self.diag == DiagonalKind::Unit {
let n = self.dim();
for j in 0..n {
for i in 0..n {
if self.in_triangle(i, j) && i != j {
if let Some(val) = self.packed.get_mut(i, j) {
*val *= alpha;
}
}
}
}
} else {
self.packed.scale(alpha);
}
}
pub fn transpose(&self) -> Self
where
T: bytemuck::Zeroable,
{
TriangularMat {
packed: self.packed.transpose(),
diag: self.diag,
}
}
}
impl<T: Scalar + core::fmt::Debug> core::fmt::Debug for TriangularMat<T> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
let n = self.dim();
writeln!(
f,
"TriangularMat {}×{} ({:?}, {:?}) {{",
n,
n,
self.uplo(),
self.diag
)?;
let max_dim = 8.min(n);
for i in 0..max_dim {
write!(f, " [")?;
for j in 0..max_dim {
if j > 0 {
write!(f, ", ")?;
}
if self.in_triangle(i, j) {
if self.diag == DiagonalKind::Unit && i == j {
write!(f, "{:8.4?}", T::one())?;
} else if let Some(v) = self.packed.get(i, j) {
write!(f, "{:8.4?}", v)?;
} else {
write!(f, " * ")?;
}
} else {
write!(f, " 0 ")?;
}
}
if n > max_dim {
write!(f, ", ...")?;
}
writeln!(f, "]")?;
}
if n > max_dim {
writeln!(f, " ...")?;
}
write!(f, "}}")
}
}
#[derive(Clone, Copy)]
pub struct TriangularRef<'a, T: Scalar> {
packed: PackedRef<'a, T>,
diag: DiagonalKind,
}
impl<'a, T: Scalar> TriangularRef<'a, T> {
#[inline]
pub fn new(packed: PackedRef<'a, T>, diag: DiagonalKind) -> Self {
TriangularRef { packed, diag }
}
#[inline]
pub fn from_slice(data: &'a [T], n: usize, uplo: TriangularKind, diag: DiagonalKind) -> Self {
TriangularRef {
packed: PackedRef::from_slice(data, n, uplo),
diag,
}
}
#[inline]
pub fn dim(&self) -> usize {
self.packed.dim()
}
#[inline]
pub fn uplo(&self) -> TriangularKind {
self.packed.kind()
}
#[inline]
pub fn diag(&self) -> DiagonalKind {
self.diag
}
#[inline]
pub fn get(&self, row: usize, col: usize) -> Option<&T> {
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.packed.get(row, col)
}
#[inline]
pub fn as_packed(&self) -> PackedRef<'a, T> {
self.packed
}
}
pub struct TriangularMut<'a, T: Scalar> {
packed: PackedMut<'a, T>,
diag: DiagonalKind,
}
impl<'a, T: Scalar> TriangularMut<'a, T> {
#[inline]
pub fn new(packed: PackedMut<'a, T>, diag: DiagonalKind) -> Self {
TriangularMut { packed, diag }
}
#[inline]
pub fn from_slice(
data: &'a mut [T],
n: usize,
uplo: TriangularKind,
diag: DiagonalKind,
) -> Self {
TriangularMut {
packed: PackedMut::from_slice(data, n, uplo),
diag,
}
}
#[inline]
pub fn dim(&self) -> usize {
self.packed.dim()
}
#[inline]
pub fn uplo(&self) -> TriangularKind {
self.packed.kind()
}
#[inline]
pub fn diag(&self) -> DiagonalKind {
self.diag
}
#[inline]
pub fn get(&self, row: usize, col: usize) -> Option<&T> {
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.packed.get(row, col)
}
#[inline]
pub fn get_mut(&mut self, row: usize, col: usize) -> Option<&mut T> {
if self.diag == DiagonalKind::Unit && row == col {
return None;
}
self.packed.get_mut(row, col)
}
#[inline]
pub fn set(&mut self, row: usize, col: usize, value: T) {
assert!(
!(self.diag == DiagonalKind::Unit && row == col),
"Cannot set diagonal of unit triangular matrix"
);
assert!(
self.packed.packed_index(row, col).is_some(),
"Element outside stored triangle"
);
self.packed
.set(row, col, value)
.expect("packed_index check above guarantees this index is valid");
}
#[inline]
pub fn rb(&self) -> TriangularRef<'_, T> {
TriangularRef {
packed: self.packed.rb(),
diag: self.diag,
}
}
#[inline]
pub fn rb_mut(&mut self) -> TriangularMut<'_, T> {
TriangularMut {
packed: self.packed.rb_mut(),
diag: self.diag,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_triangular_view_upper() {
let mut m: Mat<f64> = Mat::zeros(3, 3);
m[(0, 0)] = 1.0;
m[(0, 1)] = 2.0;
m[(0, 2)] = 3.0;
m[(1, 1)] = 4.0;
m[(1, 2)] = 5.0;
m[(2, 2)] = 6.0;
m[(1, 0)] = 99.0;
m[(2, 0)] = 99.0;
m[(2, 1)] = 99.0;
let tri = TriangularView::new(m.as_ref(), TriangularKind::Upper, DiagonalKind::NonUnit);
assert_eq!(tri.get(0, 0), Some(&1.0));
assert_eq!(tri.get(0, 1), Some(&2.0));
assert_eq!(tri.get(0, 2), Some(&3.0));
assert_eq!(tri.get(1, 1), Some(&4.0));
assert_eq!(tri.get(1, 2), Some(&5.0));
assert_eq!(tri.get(2, 2), Some(&6.0));
assert_eq!(tri.get(1, 0), None);
assert_eq!(tri.get(2, 0), None);
assert_eq!(tri.get(2, 1), None);
let dense = tri.to_dense();
assert_eq!(dense[(0, 0)], 1.0);
assert_eq!(dense[(0, 1)], 2.0);
assert_eq!(dense[(1, 0)], 0.0); assert_eq!(dense[(2, 1)], 0.0);
}
#[test]
fn test_triangular_view_unit() {
let mut m: Mat<f64> = Mat::zeros(3, 3);
m[(0, 0)] = 99.0; m[(0, 1)] = 2.0;
m[(0, 2)] = 3.0;
m[(1, 1)] = 99.0; m[(1, 2)] = 4.0;
m[(2, 2)] = 99.0;
let tri = TriangularView::new(m.as_ref(), TriangularKind::Upper, DiagonalKind::Unit);
assert_eq!(tri.get(0, 0), None);
assert_eq!(tri.get(1, 1), None);
assert_eq!(tri.get(2, 2), None);
assert_eq!(tri.get(0, 1), Some(&2.0));
assert_eq!(tri.get(0, 2), Some(&3.0));
assert_eq!(tri.get(1, 2), Some(&4.0));
let dense = tri.to_dense();
assert_eq!(dense[(0, 0)], 1.0);
assert_eq!(dense[(1, 1)], 1.0);
assert_eq!(dense[(2, 2)], 1.0);
assert_eq!(dense[(0, 1)], 2.0);
}
#[test]
fn test_triangular_view_lower() {
let mut m: Mat<f64> = Mat::zeros(3, 3);
m[(0, 0)] = 1.0;
m[(1, 0)] = 2.0;
m[(1, 1)] = 3.0;
m[(2, 0)] = 4.0;
m[(2, 1)] = 5.0;
m[(2, 2)] = 6.0;
let tri = TriangularView::new(m.as_ref(), TriangularKind::Lower, DiagonalKind::NonUnit);
assert_eq!(tri.get(0, 0), Some(&1.0));
assert_eq!(tri.get(1, 0), Some(&2.0));
assert_eq!(tri.get(1, 1), Some(&3.0));
assert_eq!(tri.get(2, 0), Some(&4.0));
assert_eq!(tri.get(2, 1), Some(&5.0));
assert_eq!(tri.get(2, 2), Some(&6.0));
assert_eq!(tri.get(0, 1), None);
assert_eq!(tri.get(0, 2), None);
assert_eq!(tri.get(1, 2), None);
}
#[test]
fn test_triangular_mat_packed() {
let mut tri: TriangularMat<f64> =
TriangularMat::zeros(3, TriangularKind::Upper, DiagonalKind::NonUnit);
tri.set(0, 0, 1.0);
tri.set(0, 1, 2.0);
tri.set(0, 2, 3.0);
tri.set(1, 1, 4.0);
tri.set(1, 2, 5.0);
tri.set(2, 2, 6.0);
assert_eq!(tri.get(0, 0), Some(&1.0));
assert_eq!(tri.get(0, 1), Some(&2.0));
assert_eq!(tri.get(1, 2), Some(&5.0));
assert_eq!(tri.get(1, 0), None);
let diag = tri.diagonal();
assert_eq!(diag, vec![1.0, 4.0, 6.0]);
}
#[test]
fn test_triangular_mat_unit() {
let mut tri: TriangularMat<f64> = TriangularMat::unit_zeros(3, TriangularKind::Lower);
tri.set(1, 0, 2.0);
tri.set(2, 0, 3.0);
tri.set(2, 1, 4.0);
assert_eq!(tri.get(0, 0), None);
assert_eq!(tri.get(1, 1), None);
assert_eq!(tri.get(2, 2), None);
assert_eq!(tri.get(1, 0), Some(&2.0));
assert_eq!(tri.get(2, 0), Some(&3.0));
assert_eq!(tri.get(2, 1), Some(&4.0));
let diag = tri.diagonal();
assert_eq!(diag, vec![1.0, 1.0, 1.0]);
let dense = tri.to_dense();
assert_eq!(dense[(0, 0)], 1.0);
assert_eq!(dense[(1, 1)], 1.0);
assert_eq!(dense[(2, 2)], 1.0);
assert_eq!(dense[(1, 0)], 2.0);
assert_eq!(dense[(0, 1)], 0.0);
}
#[test]
fn test_triangular_mat_transpose() {
let mut tri: TriangularMat<f64> =
TriangularMat::zeros(3, TriangularKind::Upper, DiagonalKind::NonUnit);
tri.set(0, 0, 1.0);
tri.set(0, 1, 2.0);
tri.set(0, 2, 3.0);
tri.set(1, 1, 4.0);
tri.set(1, 2, 5.0);
tri.set(2, 2, 6.0);
let tri_t = tri.transpose();
assert_eq!(tri_t.uplo(), TriangularKind::Lower);
assert_eq!(tri_t.get(0, 0), Some(&1.0));
assert_eq!(tri_t.get(1, 0), Some(&2.0)); assert_eq!(tri_t.get(2, 0), Some(&3.0)); assert_eq!(tri_t.get(1, 1), Some(&4.0));
assert_eq!(tri_t.get(2, 1), Some(&5.0)); assert_eq!(tri_t.get(2, 2), Some(&6.0));
}
#[test]
fn test_triangular_view_mut() {
let mut m: Mat<f64> = Mat::zeros(3, 3);
let mut view =
TriangularViewMut::new(m.as_mut(), TriangularKind::Upper, DiagonalKind::NonUnit);
view.set(0, 0, 1.0);
view.set(0, 1, 2.0);
view.set(1, 1, 3.0);
assert_eq!(view.get(0, 0), Some(&1.0));
assert_eq!(view.get(0, 1), Some(&2.0));
assert_eq!(view.get(1, 1), Some(&3.0));
view.zero_non_triangle();
view.fill(5.0);
assert_eq!(view.get(0, 0), Some(&5.0));
assert_eq!(view.get(0, 1), Some(&5.0));
}
#[test]
fn test_triangular_from_dense() {
let m = Mat::from_rows(&[&[1.0, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let tri =
TriangularMat::from_dense(&m.as_ref(), TriangularKind::Upper, DiagonalKind::NonUnit);
assert_eq!(tri.get(0, 0), Some(&1.0));
assert_eq!(tri.get(0, 1), Some(&2.0));
assert_eq!(tri.get(0, 2), Some(&3.0));
assert_eq!(tri.get(1, 1), Some(&5.0));
assert_eq!(tri.get(1, 2), Some(&6.0));
assert_eq!(tri.get(2, 2), Some(&9.0));
assert_eq!(tri.get(1, 0), None);
assert_eq!(tri.get(2, 0), None);
assert_eq!(tri.get(2, 1), None);
}
#[test]
fn test_triangular_ref_mut() {
let mut data = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut tmut =
TriangularMut::from_slice(&mut data, 3, TriangularKind::Upper, DiagonalKind::NonUnit);
assert_eq!(tmut.get(0, 0), Some(&1.0));
assert_eq!(tmut.get(0, 1), Some(&2.0));
tmut.set(0, 0, 10.0);
assert_eq!(tmut.get(0, 0), Some(&10.0));
assert_eq!(data[0], 10.0);
}
}