use singe_cuda::{
data_type::{DataType, DataTypeLike},
memory::DeviceMemory,
};
use crate::{
context::Context,
dense::validation::{
require_host_workspace, require_info_buffer, require_pivot64_buffer,
require_workspace_bytes, validate_x_matrix, validate_xlarft_inputs,
},
error::Result,
layout::{ByteWorkspaceMut, MatrixMut, MatrixRef, VectorRef, WorkspaceSizes},
params::Params,
sys, try_ffi,
types::{DiagonalType, DirectMode, FillMode, Operation, StorevMode},
utility::{to_i64, to_usize},
};
pub fn xpotrf_buffer_size<TA: DataTypeLike>(
ctx: &Context,
params: &Params,
fill_mode: FillMode,
n: usize,
a: MatrixRef<'_, TA>,
compute_type: DataType,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
let a_type = TA::data_type();
validate_x_matrix(n, n, a.data.byte_len(), a.leading_dimension, a_type)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXpotrf_bufferSize(
ctx.as_raw(),
params.as_raw(),
fill_mode.into(),
to_i64(n, "n")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
compute_type.into(),
&raw mut device_bytes,
&raw mut host_bytes,
))?;
}
Ok(WorkspaceSizes::new(
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
pub fn xpotrf<TA: DataTypeLike>(
ctx: &Context,
params: &Params,
fill_mode: FillMode,
n: usize,
a: MatrixMut<'_, TA>,
compute_type: DataType,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
let a_type = TA::data_type();
validate_x_matrix(n, n, a.data.byte_len(), a.leading_dimension, a_type)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xpotrf_buffer_size(ctx, params, fill_mode, n, a.as_ref(), compute_type)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
unsafe {
try_ffi!(sys::cusolverDnXpotrf(
ctx.as_raw(),
params.as_raw(),
fill_mode.into(),
to_i64(n, "n")?,
a_type.into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
compute_type.into(),
workspace.device.as_mut_ptr().cast(),
workspace_sizes.device_bytes as _,
workspace.host.as_mut_ptr().cast(),
workspace_sizes.host_bytes as _,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub fn xpotrs<TA: DataTypeLike, TB: DataTypeLike>(
ctx: &Context,
params: &Params,
fill_mode: FillMode,
n: usize,
nrhs: usize,
a: MatrixRef<'_, TA>,
b: MatrixMut<'_, TB>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
let a_type = TA::data_type();
let b_type = TB::data_type();
validate_x_matrix(n, n, a.data.byte_len(), a.leading_dimension, a_type)?;
validate_x_matrix(n, nrhs, b.data.byte_len(), b.leading_dimension, b_type)?;
require_info_buffer(dev_info)?;
unsafe {
try_ffi!(sys::cusolverDnXpotrs(
ctx.as_raw(),
params.as_raw(),
fill_mode.into(),
to_i64(n, "n")?,
to_i64(nrhs, "nrhs")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
b_type.into(),
b.data.as_mut_ptr().cast(),
to_i64(b.leading_dimension, "ldb")?,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub fn xtrtri_buffer_size<TA: DataTypeLike>(
ctx: &Context,
fill_mode: FillMode,
diagonal_type: DiagonalType,
n: usize,
a: MatrixRef<'_, TA>,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
validate_x_matrix(
n,
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXtrtri_bufferSize(
ctx.as_raw(),
fill_mode.into(),
diagonal_type.into(),
to_i64(n, "n")?,
TA::data_type().into(),
a.data.as_ptr().cast_mut().cast(),
to_i64(a.leading_dimension, "lda")?,
&raw mut device_bytes,
&raw mut host_bytes,
))?;
}
Ok(WorkspaceSizes::new(
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
pub fn xtrtri<TA: DataTypeLike>(
ctx: &Context,
fill_mode: FillMode,
diagonal_type: DiagonalType,
n: usize,
a: MatrixMut<'_, TA>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_x_matrix(
n,
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xtrtri_buffer_size(ctx, fill_mode, diagonal_type, n, a.as_ref())?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
unsafe {
try_ffi!(sys::cusolverDnXtrtri(
ctx.as_raw(),
fill_mode.into(),
diagonal_type.into(),
to_i64(n, "n")?,
TA::data_type().into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
workspace.device.as_mut_ptr().cast(),
workspace_sizes.device_bytes as _,
workspace.host.as_mut_ptr().cast(),
workspace_sizes.host_bytes as _,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub fn xgetrf_buffer_size<TA: DataTypeLike>(
ctx: &Context,
params: &Params,
m: usize,
n: usize,
a: MatrixRef<'_, TA>,
compute_type: DataType,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
let a_type = TA::data_type();
validate_x_matrix(m, n, a.data.byte_len(), a.leading_dimension, a_type)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXgetrf_bufferSize(
ctx.as_raw(),
params.as_raw(),
to_i64(m, "m")?,
to_i64(n, "n")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
compute_type.into(),
&raw mut device_bytes,
&raw mut host_bytes,
))?;
}
Ok(WorkspaceSizes::new(
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
pub fn xgetrf<TA: DataTypeLike>(
ctx: &Context,
params: &Params,
m: usize,
n: usize,
a: MatrixMut<'_, TA>,
pivots: Option<&mut DeviceMemory<i64>>,
compute_type: DataType,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
let a_type = TA::data_type();
validate_x_matrix(m, n, a.data.byte_len(), a.leading_dimension, a_type)?;
if let Some(pivots) = pivots.as_ref() {
require_pivot64_buffer(pivots, m.min(n))?;
}
require_info_buffer(dev_info)?;
let workspace_sizes = xgetrf_buffer_size(ctx, params, m, n, a.as_ref(), compute_type)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
unsafe {
try_ffi!(sys::cusolverDnXgetrf(
ctx.as_raw(),
params.as_raw(),
to_i64(m, "m")?,
to_i64(n, "n")?,
a_type.into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
pivots.map_or(std::ptr::null_mut(), |p| p.as_mut_ptr()),
compute_type.into(),
workspace.device.as_mut_ptr().cast(),
workspace_sizes.device_bytes as _,
workspace.host.as_mut_ptr().cast(),
workspace_sizes.host_bytes as _,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub fn xgetrs<TA: DataTypeLike, TB: DataTypeLike>(
ctx: &Context,
params: &Params,
operation: Operation,
n: usize,
nrhs: usize,
a: MatrixRef<'_, TA>,
pivots: &DeviceMemory<i64>,
b: MatrixMut<'_, TB>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
let a_type = TA::data_type();
let b_type = TB::data_type();
validate_x_matrix(n, n, a.data.byte_len(), a.leading_dimension, a_type)?;
require_pivot64_buffer(pivots, n)?;
validate_x_matrix(n, nrhs, b.data.byte_len(), b.leading_dimension, b_type)?;
require_info_buffer(dev_info)?;
unsafe {
try_ffi!(sys::cusolverDnXgetrs(
ctx.as_raw(),
params.as_raw(),
operation.into(),
to_i64(n, "n")?,
to_i64(nrhs, "nrhs")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
pivots.as_ptr().cast(),
b_type.into(),
b.data.as_mut_ptr().cast(),
to_i64(b.leading_dimension, "ldb")?,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub fn xsytrs_buffer_size<TA: DataTypeLike, TB: DataTypeLike>(
ctx: &Context,
fill_mode: FillMode,
n: usize,
nrhs: usize,
a: MatrixRef<'_, TA>,
pivots: Option<&DeviceMemory<i64>>,
b: MatrixRef<'_, TB>,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
validate_x_matrix(
n,
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
)?;
validate_x_matrix(
n,
nrhs,
b.data.byte_len(),
b.leading_dimension,
TB::data_type(),
)?;
if let Some(pivots) = pivots {
require_pivot64_buffer(pivots, n)?;
}
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXsytrs_bufferSize(
ctx.as_raw(),
fill_mode.into(),
to_i64(n, "n")?,
to_i64(nrhs, "nrhs")?,
TA::data_type().into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
pivots.map_or(std::ptr::null(), DeviceMemory::as_ptr),
TB::data_type().into(),
b.data.as_ptr().cast_mut().cast(),
to_i64(b.leading_dimension, "ldb")?,
&raw mut device_bytes,
&raw mut host_bytes,
))?;
}
Ok(WorkspaceSizes::new(
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
pub fn xsytrs<TA: DataTypeLike, TB: DataTypeLike>(
ctx: &Context,
fill_mode: FillMode,
n: usize,
nrhs: usize,
a: MatrixRef<'_, TA>,
pivots: Option<&DeviceMemory<i64>>,
b: MatrixMut<'_, TB>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_x_matrix(
n,
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
)?;
validate_x_matrix(
n,
nrhs,
b.data.byte_len(),
b.leading_dimension,
TB::data_type(),
)?;
if let Some(pivots) = pivots {
require_pivot64_buffer(pivots, n)?;
}
require_info_buffer(dev_info)?;
let workspace_sizes = xsytrs_buffer_size(ctx, fill_mode, n, nrhs, a, pivots, b.as_ref())?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
unsafe {
try_ffi!(sys::cusolverDnXsytrs(
ctx.as_raw(),
fill_mode.into(),
to_i64(n, "n")?,
to_i64(nrhs, "nrhs")?,
TA::data_type().into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
pivots.map_or(std::ptr::null(), DeviceMemory::as_ptr),
TB::data_type().into(),
b.data.as_mut_ptr().cast(),
to_i64(b.leading_dimension, "ldb")?,
workspace.device.as_mut_ptr().cast(),
workspace_sizes.device_bytes as _,
workspace.host.as_mut_ptr().cast(),
workspace_sizes.host_bytes as _,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub fn xlarft_buffer_size<TV: DataTypeLike, TTau: DataTypeLike, TT: DataTypeLike>(
ctx: &Context,
params: &Params,
direct: DirectMode,
storev: StorevMode,
n: usize,
k: usize,
v: MatrixRef<'_, TV>,
tau: VectorRef<'_, TTau>,
t: MatrixRef<'_, TT>,
compute_type: DataType,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
let v_type = TV::data_type();
let tau_type = TTau::data_type();
let t_type = TT::data_type();
validate_xlarft_inputs(
n,
k,
storev,
v.data.byte_len(),
v.leading_dimension,
v_type,
tau.data.byte_len(),
tau_type,
t.data.byte_len(),
t.leading_dimension,
t_type,
)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXlarft_bufferSize(
ctx.as_raw(),
params.as_raw(),
direct.into(),
storev.into(),
to_i64(n, "n")?,
to_i64(k, "k")?,
v_type.into(),
v.data.as_ptr().cast(),
to_i64(v.leading_dimension, "ldv")?,
tau_type.into(),
tau.data.as_ptr().cast(),
t_type.into(),
t.data.as_ptr().cast_mut().cast(),
to_i64(t.leading_dimension, "ldt")?,
compute_type.into(),
&raw mut device_bytes,
&raw mut host_bytes,
))?;
}
Ok(WorkspaceSizes::new(
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
pub fn xlarft<TV: DataTypeLike, TTau: DataTypeLike, TT: DataTypeLike>(
ctx: &Context,
params: &Params,
direct: DirectMode,
storev: StorevMode,
n: usize,
k: usize,
v: MatrixRef<'_, TV>,
tau: VectorRef<'_, TTau>,
t: MatrixMut<'_, TT>,
compute_type: DataType,
workspace: ByteWorkspaceMut<'_>,
) -> Result<()> {
ctx.bind()?;
let v_type = TV::data_type();
let tau_type = TTau::data_type();
let t_type = TT::data_type();
validate_xlarft_inputs(
n,
k,
storev,
v.data.byte_len(),
v.leading_dimension,
v_type,
tau.data.byte_len(),
tau_type,
t.data.byte_len(),
t.leading_dimension,
t_type,
)?;
let workspace_sizes = xlarft_buffer_size(
ctx,
params,
direct,
storev,
n,
k,
v,
tau,
t.as_ref(),
compute_type,
)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
unsafe {
try_ffi!(sys::cusolverDnXlarft(
ctx.as_raw(),
params.as_raw(),
direct.into(),
storev.into(),
to_i64(n, "n")?,
to_i64(k, "k")?,
v_type.into(),
v.data.as_ptr().cast(),
to_i64(v.leading_dimension, "ldv")?,
tau_type.into(),
tau.data.as_ptr().cast(),
t_type.into(),
t.data.as_mut_ptr().cast(),
to_i64(t.leading_dimension, "ldt")?,
compute_type.into(),
workspace.device.as_mut_ptr().cast(),
workspace_sizes.device_bytes as _,
workspace.host.as_mut_ptr().cast(),
workspace_sizes.host_bytes as _,
))?;
}
Ok(())
}