#![cfg(feature = "cuda")]
use crate::error::{LinalgError, LinalgResult};
use oxicuda_blas::types::{FillMode, Layout, MatrixDesc, MatrixDescMut, Transpose};
use scirs2_core::ndarray::{Array1, Array2, ArrayView1, ArrayView2};
pub fn cuda_is_available() -> bool {
oxicuda_driver::init().is_ok()
&& oxicuda_driver::device::Device::count()
.map(|c| c > 0)
.unwrap_or(false)
}
fn build_context() -> LinalgResult<std::sync::Arc<oxicuda_driver::Context>> {
oxicuda_driver::init()
.map_err(|e| LinalgError::ComputationError(format!("CUDA unavailable: {e}")))?;
let count = oxicuda_driver::device::Device::count()
.map_err(|e| LinalgError::ComputationError(format!("device count: {e}")))?;
if count <= 0 {
return Err(LinalgError::ComputationError(
"no NVIDIA CUDA device available".into(),
));
}
let dev = oxicuda_driver::device::Device::get(0).map_err(cuda_err)?;
Ok(std::sync::Arc::new(
oxicuda_driver::Context::new(&dev).map_err(cuda_err)?,
))
}
fn solver_err(e: oxicuda_solver::error::SolverError) -> LinalgError {
LinalgError::ComputationError(format!("oxicuda-solver: {e}"))
}
fn blas_err(e: oxicuda_blas::error::BlasError) -> LinalgError {
LinalgError::ComputationError(format!("oxicuda-blas: {e}"))
}
fn cuda_err(e: oxicuda_driver::CudaError) -> LinalgError {
LinalgError::ComputationError(format!("oxicuda CUDA driver: {e}"))
}
pub fn cuda_gemm(a: &ArrayView2<f64>, b: &ArrayView2<f64>) -> LinalgResult<Array2<f64>> {
let (m, k) = (a.nrows(), a.ncols());
let (k_b, n) = (b.nrows(), b.ncols());
if k != k_b {
return Err(LinalgError::ShapeError(format!(
"cuda_gemm: inner dimensions disagree: A is {m}x{k}, B is {k_b}x{n}"
)));
}
if m == 0 || k == 0 || n == 0 {
return Ok(Array2::zeros((m, n)));
}
let a_std = a.as_standard_layout();
let b_std = b.as_standard_layout();
let a_slice = a_std
.as_slice()
.ok_or_else(|| LinalgError::ComputationError("cuda_gemm: A not contiguous".into()))?;
let b_slice = b_std
.as_slice()
.ok_or_else(|| LinalgError::ComputationError("cuda_gemm: B not contiguous".into()))?;
let ctx = build_context()?;
let handle = oxicuda_blas::BlasHandle::new(&ctx).map_err(blas_err)?;
let d_a = oxicuda_memory::DeviceBuffer::from_host(a_slice).map_err(cuda_err)?;
let d_b = oxicuda_memory::DeviceBuffer::from_host(b_slice).map_err(cuda_err)?;
let mut d_c = oxicuda_memory::DeviceBuffer::<f64>::alloc(m * n).map_err(cuda_err)?;
let a_desc =
MatrixDesc::from_buffer(&d_a, m as u32, k as u32, Layout::RowMajor).map_err(blas_err)?;
let b_desc =
MatrixDesc::from_buffer(&d_b, k as u32, n as u32, Layout::RowMajor).map_err(blas_err)?;
let mut c_desc = MatrixDescMut::from_buffer(&mut d_c, m as u32, n as u32, Layout::RowMajor)
.map_err(blas_err)?;
oxicuda_blas::level3::gemm_api::gemm::<f64>(
&handle,
Transpose::NoTrans,
Transpose::NoTrans,
1.0,
&a_desc,
&b_desc,
0.0,
&mut c_desc,
)
.map_err(blas_err)?;
let mut c_host = vec![0.0f64; m * n];
d_c.copy_to_host(&mut c_host).map_err(cuda_err)?;
Array2::from_shape_vec((m, n), c_host)
.map_err(|e| LinalgError::ComputationError(format!("cuda_gemm: reshape failed: {e}")))
}
pub fn cuda_solve_spd(a: &ArrayView2<f64>, b: &ArrayView1<f64>) -> LinalgResult<Array1<f64>> {
let n = a.nrows();
if a.ncols() != n {
return Err(LinalgError::ShapeError(format!(
"cuda_solve_spd: A must be square, got {n}x{}",
a.ncols()
)));
}
if b.len() != n {
return Err(LinalgError::DimensionError(format!(
"cuda_solve_spd: b length {} does not match A dimension {n}",
b.len()
)));
}
if n == 0 {
return Ok(Array1::zeros(0));
}
let a_std = a.as_standard_layout();
let a_slice = a_std
.as_slice()
.ok_or_else(|| LinalgError::ComputationError("cuda_solve_spd: A not contiguous".into()))?;
let b_vec: Vec<f64> = b.iter().copied().collect();
let ctx = build_context()?;
let mut handle = oxicuda_solver::SolverHandle::new(&ctx).map_err(solver_err)?;
let mut d_a = oxicuda_memory::DeviceBuffer::from_host(a_slice).map_err(cuda_err)?;
let mut d_b = oxicuda_memory::DeviceBuffer::from_host(&b_vec).map_err(cuda_err)?;
oxicuda_solver::dense::cholesky::<f64>(
&mut handle,
FillMode::Lower,
&mut d_a,
n as u32,
n as u32,
)
.map_err(solver_err)?;
oxicuda_solver::dense::cholesky_solve::<f64>(
&handle,
FillMode::Lower,
&d_a,
&mut d_b,
n as u32,
1,
)
.map_err(solver_err)?;
let mut x = vec![0.0f64; n];
d_b.copy_to_host(&mut x).map_err(cuda_err)?;
Ok(Array1::from_vec(x))
}
#[cfg(test)]
mod tests {
use super::*;
fn mat(rows: usize, cols: usize, data: Vec<f64>) -> Array2<f64> {
Array2::from_shape_vec((rows, cols), data).expect("valid test matrix shape")
}
#[test]
fn cuda_gemm_or_skip() {
if !cuda_is_available() {
eprintln!("skipping: no NVIDIA CUDA device");
assert!(!cuda_is_available());
return;
}
let a = mat(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
let b = mat(3, 2, vec![7.0, 8.0, 9.0, 10.0, 11.0, 12.0]);
let c = cuda_gemm(&a.view(), &b.view()).expect("cuda_gemm failed");
let expected = mat(2, 2, vec![58.0, 64.0, 139.0, 154.0]);
let max_diff = c
.iter()
.zip(expected.iter())
.map(|(g, e)| (g - e).abs())
.fold(0.0f64, f64::max);
assert!(max_diff < 1e-9, "max abs diff {max_diff} exceeds 1e-9");
}
#[test]
fn cuda_gemm_nonsquare_stress_or_skip() {
if !cuda_is_available() {
eprintln!("skipping: no NVIDIA CUDA device");
assert!(!cuda_is_available());
return;
}
let m = 17usize;
let k = 5usize;
let n = 31usize;
let a_data: Vec<f64> = (0..m * k).map(|idx| (idx as f64 + 1.0) * 0.1).collect();
let b_data: Vec<f64> = (0..k * n)
.map(|idx| (idx as f64 * 1.3 + 0.5) * 0.07)
.collect();
let a = mat(m, k, a_data);
let b = mat(k, n, b_data);
let c_gpu = cuda_gemm(&a.view(), &b.view()).expect("cuda_gemm failed on 17x5 . 5x31");
let c_cpu = a.dot(&b);
assert_eq!(c_gpu.shape(), &[m, n], "output shape must be {m}x{n}");
let max_diff = c_gpu
.iter()
.zip(c_cpu.iter())
.map(|(g, e)| (g - e).abs())
.fold(0.0f64, f64::max);
assert!(
max_diff < 1e-9,
"17x5 . 5x31 max abs diff {max_diff} exceeds 1e-9"
);
}
#[test]
fn cuda_solve_spd_or_skip() {
if !cuda_is_available() {
eprintln!("skipping: no NVIDIA CUDA device");
assert!(!cuda_is_available());
return;
}
let a = mat(3, 3, vec![4.0, 1.0, 0.0, 1.0, 3.0, 1.0, 0.0, 1.0, 2.0]);
let b = Array1::from_vec(vec![1.0, 2.0, 3.0]);
let x = cuda_solve_spd(&a.view(), &b.view()).expect("cuda_solve_spd failed");
let ax = a.dot(&x);
let max_diff = ax
.iter()
.zip(b.iter())
.map(|(g, e)| (g - e).abs())
.fold(0.0f64, f64::max);
assert!(max_diff < 1e-9, "max abs diff {max_diff} exceeds 1e-9");
}
#[test]
fn cuda_gemm_shape_mismatch_errors() {
let a = mat(1, 3, vec![1.0, 2.0, 3.0]);
let b = mat(1, 2, vec![1.0, 2.0]);
assert!(cuda_gemm(&a.view(), &b.view()).is_err());
}
#[test]
fn cuda_solve_spd_shape_mismatch_errors() {
let a = mat(2, 2, vec![1.0, 0.0, 0.0, 1.0]);
let b = Array1::from_vec(vec![1.0, 2.0, 3.0]);
assert!(cuda_solve_spd(&a.view(), &b.view()).is_err());
let nonsquare = mat(1, 3, vec![1.0, 2.0, 3.0]);
let b2 = Array1::from_vec(vec![1.0]);
assert!(cuda_solve_spd(&nonsquare.view(), &b2.view()).is_err());
}
#[test]
fn cuda_solve_spd_large_n_blocked_or_skip() {
if !cuda_is_available() {
eprintln!("skipping: no NVIDIA CUDA device");
assert!(!cuda_is_available());
return;
}
for &n in &[80_usize, 128_usize] {
let m_flat: Vec<f64> = (0..n * n)
.map(|idx| {
let i = idx / n;
let j = idx % n;
(i + j + 1) as f64 / (2 * n) as f64
})
.collect();
let m = Array2::from_shape_vec((n, n), m_flat).expect("deterministic M must be valid");
let mut a = m.t().dot(&m);
for i in 0..n {
a[[i, i]] += n as f64;
}
let ones = Array1::from_elem(n, 1.0_f64);
let b = a.dot(&ones);
let x = cuda_solve_spd(&a.view(), &b.view())
.unwrap_or_else(|e| panic!("cuda_solve_spd n={n}: {e}"));
let ax = a.dot(&x);
let tol = (n as f64) * 1e-10;
let max_diff = ax
.iter()
.zip(b.iter())
.map(|(gpu_val, expected)| (gpu_val - expected).abs())
.fold(0.0_f64, f64::max);
eprintln!("n={n}: blocked-Cholesky residual max_diff={max_diff:.3e} tol={tol:.3e}");
assert!(
max_diff < tol,
"n={n}: blocked Cholesky residual {max_diff:.3e} exceeds tight tol {tol:.3e}"
);
}
}
}