lbfgsbrs 0.1.2

Rust port of L-BFGS-B-C
Documentation
use super::miniblas::{ddot, daxpy};

/// Error type for LINPACK operations
#[derive(Debug)]
pub enum LinpackError {
    NonPositiveDefinite(usize),
    ZeroDiagonal(usize),
    InvalidJob(i32)
}

/// Performs Cholesky decomposition of a positive-definite matrix.
/// The matrix is stored in column-major format.
/// 
/// # Arguments
/// * `a` - Input/output matrix stored as a flat array in column-major format
/// * `lda` - Leading dimension of the matrix
/// * `n` - Order of the matrix
/// 
/// # Returns
/// * `Ok(0)` on success
/// * `Ok(j)` if stopped at column j (1-based) due to non-positive definite matrix
pub fn dpofa(a: &mut [f64], lda: usize, n: usize) -> Result<usize, LinpackError> {
    // Validate input parameters
    if n > lda {
        return Err(LinpackError::NonPositiveDefinite(0));
    }

    // Main loop for processing the matrix
    for j in 0..n {
        let mut s = 0.0;
        
        // Update elements below diagonal in column j
        for k in 0..j {
            let mut t = a[k + j * lda];
            
            if k > 0 {
                // Compute dot product
                t -= ddot(
                    k as i32,
                    &a[k * lda..k * lda + k],
                    1,
                    &a[j * lda..j * lda + k],
                    1
                );
            }
            
            t /= a[k + k * lda];
            a[k + j * lda] = t;
            s += t * t;
        }
        
        // Update diagonal element
        s = a[j + j * lda] - s;
        
        // Check for non-positive definite matrix
        if s <= 0.0 {
            return Ok(j + 1); // Return 1-based index where we failed
        }
        
        a[j + j * lda] = s.sqrt();
    }
    
    Ok(0)
}

/// Solves triangular systems of linear equations
/// 
/// # Arguments
/// * `t` - Input matrix stored as a flat array in column-major format
/// * `ldt` - Leading dimension of matrix T
/// * `n` - Order of matrix T
/// * `b` - Right-hand side vector on input, solution vector on output
/// * `job` - Specifies the system to solve:
///   * job % 10 = 0: solve T*x=b
///   * job % 10 = 1: solve T'*x=b
///   * job / 10 = 0: T is lower triangular
///   * job / 10 = 1: T is upper triangular
pub fn dtrsl(
    t: &mut [f64],
    ldt: usize,
    n: usize,
    b: &mut [f64],
    job: i32
) -> Result<usize, LinpackError> {
    // Check for zero diagonal elements
    for i in 0..n {
        if t[i + i * ldt] == 0.0 {
            return Err(LinpackError::ZeroDiagonal(i + 1));
        }
    }
    
    // Determine case based on job parameter
    let case = match (job % 10, job % 100 / 10) {
        (0, 0) => 1, // Forward solve with transpose
        (1, 0) => 2, // Backward solve with transpose
        (0, 1) => 3, // Forward solve
        (1, 1) => 4, // Backward solve
        _ => return Err(LinpackError::InvalidJob(job)),
    };
    
    match case {
        1 => { // Solve T*x=b, T lower triangular
            b[0] /= t[0];
            for j in 1..n {
                let temp = -b[j - 1];
                daxpy(
                    (n - j) as i32,
                    temp,
                    &t[j + (j - 1) * ldt..],
                    1,
                    &mut b[j..],
                    1
                );
                b[j] /= t[j + j * ldt];
            }
        },
        2 => { // Solve T*x=b, T upper triangular
            b[n - 1] /= t[(n - 1) + (n - 1) * ldt];
            for jj in 1..n {
                let j = n - jj - 1;
                let temp = -b[j + 1];
                daxpy(
                    (j + 1) as i32,
                    temp,
                    &t[(j + 1) * ldt..],
                    1,
                    &mut b[..j + 1],
                    1
                );
                b[j] /= t[j + j * ldt];
            }
        },
        3 => { // Solve T'*x=b, T lower triangular
            b[n - 1] /= t[(n - 1) + (n - 1) * ldt];
            for jj in 1..n {
                let j = n - jj - 1;
                let mut sum = 0.0;
                for k in (j + 1)..n {
                    sum += t[k + j * ldt] * b[k];
                }
                b[j] = (b[j] - sum) / t[j + j * ldt];
            }
        },
        4 => { // Solve T'*x=b, T upper triangular
            b[0] /= t[0];
            for j in 1..n {
                let mut sum = 0.0;
                for k in 0..j {
                    sum += t[k + j * ldt] * b[k];
                }
                b[j] = (b[j] - sum) / t[j + j * ldt];
            }
        },
        _ => unreachable!(),
    }
    
    Ok(0)
}