use std::ptr;
use singe_cuda::{
data_type::DataTypeLike,
memory::DeviceMemory,
types::{Complex32, Complex64},
};
use crate::{
context::Context,
error::{Error, Result},
layout::{
ByteWorkspaceMut, MatrixMut, MatrixRef, StridedBatchedMatrixMut, StridedBatchedMatrixRef,
StridedBatchedVectorMut, StridedBatchedVectorRef, WorkspaceSizes,
},
params::Params,
svd::{
info::GesvdjInfo,
validation::{
matrix_mut_parts, matrix_mut_ref_option, matrix_mut_ref_parts, matrix_ref_parts,
optional_gesvda_output_mut_ptr, optional_gesvda_output_ptr,
optional_gesvdj_matrix_mut_ptr, optional_gesvdj_matrix_ptr,
optional_x_eig_matrix_mut_ptr, optional_x_eig_matrix_ptr, optional_x_matrix_mut_ptr,
optional_x_matrix_ptr, optional_x_truncated_u_mut_ptr, optional_x_truncated_u_ptr,
optional_x_truncated_v_mut_ptr, optional_x_truncated_v_ptr, require_host_workspace,
require_info_buffer, require_info_buffer_len, require_workspace,
require_workspace_bytes, strided_batched_matrix_mut_parts,
strided_batched_matrix_mut_ref_option, strided_batched_matrix_ref_parts,
validate_gesvd_dims, validate_gesvda_strided_batched_inputs,
validate_gesvdj_batched_inputs, validate_gesvdj_inputs, validate_x_matrix,
validate_x_svd_output, validate_x_vector, validate_xgesvdp_inputs,
validate_xgesvdr_inputs,
},
},
sys, try_ffi,
types::{EigenMode, SvdMode, TruncatedSvdMode},
utility::{to_i32, to_i64, to_usize},
};
pub fn xgesvd_buffer_size<
TA: DataTypeLike,
TS: DataTypeLike,
TU: DataTypeLike,
TVT: DataTypeLike,
>(
ctx: &Context,
params: &Params,
job_u: SvdMode,
job_vt: SvdMode,
m: usize,
n: usize,
a: MatrixRef<'_, TA>,
s: &DeviceMemory<TS>,
u: Option<MatrixRef<'_, TU>>,
vt: Option<MatrixRef<'_, TVT>>,
) -> Result<WorkspaceSizes> {
let a_type = TA::data_type();
let s_type = TS::data_type();
let u_type = TU::data_type();
let vt_type = TVT::data_type();
ctx.bind()?;
validate_gesvd_dims(m, n)?;
validate_x_matrix(m, n, a.data.byte_len(), a.leading_dimension, a_type)?;
validate_x_vector(m.min(n), s.byte_len(), s_type)?;
validate_x_svd_output(m, m, matrix_ref_parts(u), job_u, u_type)?;
validate_x_svd_output(n, n, matrix_ref_parts(vt), job_vt, vt_type)?;
if matches!(job_u, SvdMode::Overwrite) && matches!(job_vt, SvdMode::Overwrite) {
return Err(Error::InvalidSvdMode);
}
let (u_ptr, ldu) = optional_x_matrix_ptr(matrix_ref_parts(u), m, m, job_u, u_type)?;
let (vt_ptr, ldvt) = optional_x_matrix_ptr(matrix_ref_parts(vt), n, n, job_vt, vt_type)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXgesvd_bufferSize(
ctx.as_raw(),
params.as_raw(),
job_u.as_raw(),
job_vt.as_raw(),
to_i64(m, "m")?,
to_i64(n, "n")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
s_type.into(),
s.as_ptr().cast(),
u_type.into(),
u_ptr.cast(),
ldu,
vt_type.into(),
vt_ptr.cast(),
ldvt,
a_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 xgesvd<TA: DataTypeLike, TS: DataTypeLike, TU: DataTypeLike, TVT: DataTypeLike>(
ctx: &Context,
params: &Params,
job_u: SvdMode,
job_vt: SvdMode,
m: usize,
n: usize,
a: MatrixMut<'_, TA>,
s: &mut DeviceMemory<TS>,
u: Option<MatrixMut<'_, TU>>,
vt: Option<MatrixMut<'_, TVT>>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
let a_type = TA::data_type();
let s_type = TS::data_type();
let u_type = TU::data_type();
let vt_type = TVT::data_type();
ctx.bind()?;
validate_gesvd_dims(m, n)?;
validate_x_matrix(m, n, a.data.byte_len(), a.leading_dimension, a_type)?;
validate_x_vector(m.min(n), s.byte_len(), s_type)?;
validate_x_svd_output(m, m, matrix_mut_ref_parts(u.as_ref()), job_u, u_type)?;
validate_x_svd_output(n, n, matrix_mut_ref_parts(vt.as_ref()), job_vt, vt_type)?;
if matches!(job_u, SvdMode::Overwrite) && matches!(job_vt, SvdMode::Overwrite) {
return Err(Error::InvalidSvdMode);
}
require_info_buffer(dev_info)?;
let workspace_sizes = xgesvd_buffer_size(
ctx,
params,
job_u,
job_vt,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(vt.as_ref()),
)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
let (u_ptr, ldu) = optional_x_matrix_mut_ptr(matrix_mut_parts(u), m, m, job_u, u_type)?;
let (vt_ptr, ldvt) = optional_x_matrix_mut_ptr(matrix_mut_parts(vt), n, n, job_vt, vt_type)?;
unsafe {
try_ffi!(sys::cusolverDnXgesvd(
ctx.as_raw(),
params.as_raw(),
job_u.as_raw(),
job_vt.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")?,
s_type.into(),
s.as_mut_ptr().cast(),
u_type.into(),
u_ptr.cast(),
ldu,
vt_type.into(),
vt_ptr.cast(),
ldvt,
a_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 xgesvdp_buffer_size<
TA: DataTypeLike,
TS: DataTypeLike,
TU: DataTypeLike,
TV: DataTypeLike,
>(
ctx: &Context,
params: &Params,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixRef<'_, TA>,
s: &DeviceMemory<TS>,
u: Option<MatrixRef<'_, TU>>,
v: Option<MatrixRef<'_, TV>>,
) -> Result<WorkspaceSizes> {
let a_type = TA::data_type();
let s_type = TS::data_type();
let u_type = TU::data_type();
let v_type = TV::data_type();
ctx.bind()?;
validate_xgesvdp_inputs(
m,
n,
a.data.byte_len(),
a.leading_dimension,
a_type,
s.byte_len(),
s_type,
jobz,
econ,
matrix_ref_parts(u).as_ref(),
u_type,
matrix_ref_parts(v).as_ref(),
v_type,
)?;
let (u_ptr, ldu) = optional_x_eig_matrix_ptr(matrix_ref_parts(u), m, n, jobz, econ, u_type)?;
let (v_ptr, ldv) = optional_x_eig_matrix_ptr(matrix_ref_parts(v), n, n, jobz, econ, v_type)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXgesvdp_bufferSize(
ctx.as_raw(),
params.as_raw(),
jobz.into(),
i32::from(econ),
to_i64(m, "m")?,
to_i64(n, "n")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
s_type.into(),
s.as_ptr().cast(),
u_type.into(),
u_ptr.cast(),
ldu,
v_type.into(),
v_ptr.cast(),
ldv,
a_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 xgesvdp<TA: DataTypeLike, TS: DataTypeLike, TU: DataTypeLike, TV: DataTypeLike>(
ctx: &Context,
params: &Params,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixMut<'_, TA>,
s: &mut DeviceMemory<TS>,
u: Option<MatrixMut<'_, TU>>,
v: Option<MatrixMut<'_, TV>>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
err_sigma: Option<&mut f64>,
) -> Result<()> {
let a_type = TA::data_type();
let s_type = TS::data_type();
let u_type = TU::data_type();
let v_type = TV::data_type();
ctx.bind()?;
validate_xgesvdp_inputs(
m,
n,
a.data.byte_len(),
a.leading_dimension,
a_type,
s.byte_len(),
s_type,
jobz,
econ,
matrix_mut_ref_parts(u.as_ref()).as_ref(),
u_type,
matrix_mut_ref_parts(v.as_ref()).as_ref(),
v_type,
)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xgesvdp_buffer_size(
ctx,
params,
jobz,
econ,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
let (u_ptr, ldu) =
optional_x_eig_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, econ, u_type)?;
let (v_ptr, ldv) =
optional_x_eig_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, econ, v_type)?;
unsafe {
try_ffi!(sys::cusolverDnXgesvdp(
ctx.as_raw(),
params.as_raw(),
jobz.into(),
i32::from(econ),
to_i64(m, "m")?,
to_i64(n, "n")?,
a_type.into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
s_type.into(),
s.as_mut_ptr().cast(),
u_type.into(),
u_ptr.cast(),
ldu,
v_type.into(),
v_ptr.cast(),
ldv,
a_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(),
err_sigma.map_or(ptr::null_mut(), |value| value as *mut f64),
))?;
}
Ok(())
}
pub fn xgesvdr_buffer_size<
TA: DataTypeLike,
TS: DataTypeLike,
TU: DataTypeLike,
TV: DataTypeLike,
>(
ctx: &Context,
params: &Params,
job_u: TruncatedSvdMode,
job_v: TruncatedSvdMode,
m: usize,
n: usize,
k: usize,
p: usize,
niters: usize,
a: MatrixRef<'_, TA>,
s: &DeviceMemory<TS>,
u: Option<MatrixRef<'_, TU>>,
v: Option<MatrixRef<'_, TV>>,
) -> Result<WorkspaceSizes> {
let a_type = TA::data_type();
let s_type = TS::data_type();
let u_type = TU::data_type();
let v_type = TV::data_type();
ctx.bind()?;
validate_xgesvdr_inputs(
m,
n,
k,
p,
niters,
a.data.byte_len(),
a.leading_dimension,
a_type,
s.byte_len(),
s_type,
job_u,
matrix_ref_parts(u).as_ref(),
u_type,
job_v,
matrix_ref_parts(v).as_ref(),
v_type,
)?;
let (u_ptr, ldu) = optional_x_truncated_u_ptr(matrix_ref_parts(u), m, k, job_u, u_type)?;
let (v_ptr, ldv) = optional_x_truncated_v_ptr(matrix_ref_parts(v), n, k, job_v, v_type)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXgesvdr_bufferSize(
ctx.as_raw(),
params.as_raw(),
job_u.as_raw(),
job_v.as_raw(),
to_i64(m, "m")?,
to_i64(n, "n")?,
to_i64(k, "k")?,
to_i64(p, "p")?,
to_i64(niters, "niters")?,
a_type.into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
s_type.into(),
s.as_ptr().cast(),
u_type.into(),
u_ptr.cast(),
ldu,
v_type.into(),
v_ptr.cast(),
ldv,
a_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 xgesvdr<TA: DataTypeLike, TS: DataTypeLike, TU: DataTypeLike, TV: DataTypeLike>(
ctx: &Context,
params: &Params,
job_u: TruncatedSvdMode,
job_v: TruncatedSvdMode,
m: usize,
n: usize,
k: usize,
p: usize,
niters: usize,
a: MatrixMut<'_, TA>,
s: &mut DeviceMemory<TS>,
u: Option<MatrixMut<'_, TU>>,
v: Option<MatrixMut<'_, TV>>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
let a_type = TA::data_type();
let s_type = TS::data_type();
let u_type = TU::data_type();
let v_type = TV::data_type();
ctx.bind()?;
validate_xgesvdr_inputs(
m,
n,
k,
p,
niters,
a.data.byte_len(),
a.leading_dimension,
a_type,
s.byte_len(),
s_type,
job_u,
matrix_mut_ref_parts(u.as_ref()).as_ref(),
u_type,
job_v,
matrix_mut_ref_parts(v.as_ref()).as_ref(),
v_type,
)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xgesvdr_buffer_size(
ctx,
params,
job_u,
job_v,
m,
n,
k,
p,
niters,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
let (u_ptr, ldu) = optional_x_truncated_u_mut_ptr(matrix_mut_parts(u), m, k, job_u, u_type)?;
let (v_ptr, ldv) = optional_x_truncated_v_mut_ptr(matrix_mut_parts(v), n, k, job_v, v_type)?;
unsafe {
try_ffi!(sys::cusolverDnXgesvdr(
ctx.as_raw(),
params.as_raw(),
job_u.as_raw(),
job_v.as_raw(),
to_i64(m, "m")?,
to_i64(n, "n")?,
to_i64(k, "k")?,
to_i64(p, "p")?,
to_i64(niters, "niters")?,
a_type.into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
s_type.into(),
s.as_mut_ptr().cast(),
u_type.into(),
u_ptr.cast(),
ldu,
v_type.into(),
v_ptr.cast(),
ldv,
a_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 sgesvdj_buffer_size(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixRef<'_, f32>,
s: &DeviceMemory<f32>,
u: Option<MatrixRef<'_, f32>>,
v: Option<MatrixRef<'_, f32>>,
params: &GesvdjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_ref_parts(u),
matrix_ref_parts(v),
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, econ)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSgesvdj_bufferSize(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub fn dgesvdj_buffer_size(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixRef<'_, f64>,
s: &DeviceMemory<f64>,
u: Option<MatrixRef<'_, f64>>,
v: Option<MatrixRef<'_, f64>>,
params: &GesvdjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_ref_parts(u),
matrix_ref_parts(v),
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, econ)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDgesvdj_bufferSize(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub fn cgesvdj_buffer_size(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixRef<'_, Complex32>,
s: &DeviceMemory<f32>,
u: Option<MatrixRef<'_, Complex32>>,
v: Option<MatrixRef<'_, Complex32>>,
params: &GesvdjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_ref_parts(u),
matrix_ref_parts(v),
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, econ)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnCgesvdj_bufferSize(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub fn zgesvdj_buffer_size(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixRef<'_, Complex64>,
s: &DeviceMemory<f64>,
u: Option<MatrixRef<'_, Complex64>>,
v: Option<MatrixRef<'_, Complex64>>,
params: &GesvdjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_ref_parts(u),
matrix_ref_parts(v),
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, econ)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZgesvdj_bufferSize(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub fn sgesvdj(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixMut<'_, f32>,
s: &mut DeviceMemory<f32>,
u: Option<MatrixMut<'_, f32>>,
v: Option<MatrixMut<'_, f32>>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
)?;
require_info_buffer(dev_info)?;
let lwork = sgesvdj_buffer_size(
ctx,
jobz,
econ,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, econ)?;
unsafe {
try_ffi!(sys::cusolverDnSgesvdj(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub fn dgesvdj(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixMut<'_, f64>,
s: &mut DeviceMemory<f64>,
u: Option<MatrixMut<'_, f64>>,
v: Option<MatrixMut<'_, f64>>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
)?;
require_info_buffer(dev_info)?;
let lwork = dgesvdj_buffer_size(
ctx,
jobz,
econ,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, econ)?;
unsafe {
try_ffi!(sys::cusolverDnDgesvdj(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub fn cgesvdj(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixMut<'_, Complex32>,
s: &mut DeviceMemory<f32>,
u: Option<MatrixMut<'_, Complex32>>,
v: Option<MatrixMut<'_, Complex32>>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
)?;
require_info_buffer(dev_info)?;
let lwork = cgesvdj_buffer_size(
ctx,
jobz,
econ,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, econ)?;
unsafe {
try_ffi!(sys::cusolverDnCgesvdj(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub fn zgesvdj(
ctx: &Context,
jobz: EigenMode,
econ: bool,
m: usize,
n: usize,
a: MatrixMut<'_, Complex64>,
s: &mut DeviceMemory<f64>,
u: Option<MatrixMut<'_, Complex64>>,
v: Option<MatrixMut<'_, Complex64>>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
econ,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
)?;
require_info_buffer(dev_info)?;
let lwork = zgesvdj_buffer_size(
ctx,
jobz,
econ,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, econ)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, econ)?;
unsafe {
try_ffi!(sys::cusolverDnZgesvdj(
ctx.as_raw(),
jobz.into(),
i32::from(econ),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub fn sgesvdj_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixRef<'_, f32>,
s: &DeviceMemory<f32>,
u: Option<MatrixRef<'_, f32>>,
v: Option<MatrixRef<'_, f32>>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_ref_parts(u),
matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, true)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSgesvdjBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn dgesvdj_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixRef<'_, f64>,
s: &DeviceMemory<f64>,
u: Option<MatrixRef<'_, f64>>,
v: Option<MatrixRef<'_, f64>>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_ref_parts(u),
matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, true)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDgesvdjBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn cgesvdj_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixRef<'_, Complex32>,
s: &DeviceMemory<f32>,
u: Option<MatrixRef<'_, Complex32>>,
v: Option<MatrixRef<'_, Complex32>>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_ref_parts(u),
matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, true)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnCgesvdjBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn zgesvdj_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixRef<'_, Complex64>,
s: &DeviceMemory<f64>,
u: Option<MatrixRef<'_, Complex64>>,
v: Option<MatrixRef<'_, Complex64>>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_ref_parts(u),
matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_ptr(matrix_ref_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_ptr(matrix_ref_parts(v), n, n, jobz, true)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZgesvdjBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
&raw mut lwork,
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn sgesvdj_batched(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixMut<'_, f32>,
s: &mut DeviceMemory<f32>,
u: Option<MatrixMut<'_, f32>>,
v: Option<MatrixMut<'_, f32>>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = sgesvdj_batched_buffer_size(
ctx,
jobz,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, true)?;
unsafe {
try_ffi!(sys::cusolverDnSgesvdjBatched(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn dgesvdj_batched(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixMut<'_, f64>,
s: &mut DeviceMemory<f64>,
u: Option<MatrixMut<'_, f64>>,
v: Option<MatrixMut<'_, f64>>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = dgesvdj_batched_buffer_size(
ctx,
jobz,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, true)?;
unsafe {
try_ffi!(sys::cusolverDnDgesvdjBatched(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn cgesvdj_batched(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixMut<'_, Complex32>,
s: &mut DeviceMemory<f32>,
u: Option<MatrixMut<'_, Complex32>>,
v: Option<MatrixMut<'_, Complex32>>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = cgesvdj_batched_buffer_size(
ctx,
jobz,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, true)?;
unsafe {
try_ffi!(sys::cusolverDnCgesvdjBatched(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn zgesvdj_batched(
ctx: &Context,
jobz: EigenMode,
m: usize,
n: usize,
a: MatrixMut<'_, Complex64>,
s: &mut DeviceMemory<f64>,
u: Option<MatrixMut<'_, Complex64>>,
v: Option<MatrixMut<'_, Complex64>>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
params: &GesvdjInfo,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvdj_batched_inputs(
m,
n,
a.data.len(),
a.leading_dimension,
s.len(),
jobz,
matrix_mut_ref_parts(u.as_ref()),
matrix_mut_ref_parts(v.as_ref()),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = zgesvdj_batched_buffer_size(
ctx,
jobz,
m,
n,
a.as_ref(),
s,
matrix_mut_ref_option(u.as_ref()),
matrix_mut_ref_option(v.as_ref()),
params,
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(u), m, n, jobz, true)?;
let (v_ptr, ldv) = optional_gesvdj_matrix_mut_ptr(matrix_mut_parts(v), n, n, jobz, true)?;
unsafe {
try_ffi!(sys::cusolverDnZgesvdjBatched(
ctx.as_raw(),
jobz.into(),
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_mut_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
s.as_mut_ptr().cast(),
u_ptr.cast(),
ldu,
v_ptr.cast(),
ldv,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn sgesvda_strided_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, f32>,
s: StridedBatchedVectorRef<'_, f32>,
u: Option<StridedBatchedMatrixRef<'_, f32>>,
v: Option<StridedBatchedMatrixRef<'_, f32>>,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_ref_parts(u),
strided_batched_matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(v), n, rank, jobz)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSgesvdaStridedBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
&raw mut lwork,
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn dgesvda_strided_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, f64>,
s: StridedBatchedVectorRef<'_, f64>,
u: Option<StridedBatchedMatrixRef<'_, f64>>,
v: Option<StridedBatchedMatrixRef<'_, f64>>,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_ref_parts(u),
strided_batched_matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(v), n, rank, jobz)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDgesvdaStridedBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
&raw mut lwork,
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn cgesvda_strided_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, Complex32>,
s: StridedBatchedVectorRef<'_, f32>,
u: Option<StridedBatchedMatrixRef<'_, Complex32>>,
v: Option<StridedBatchedMatrixRef<'_, Complex32>>,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_ref_parts(u),
strided_batched_matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(v), n, rank, jobz)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnCgesvdaStridedBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
&raw mut lwork,
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn zgesvda_strided_batched_buffer_size(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, Complex64>,
s: StridedBatchedVectorRef<'_, f64>,
u: Option<StridedBatchedMatrixRef<'_, Complex64>>,
v: Option<StridedBatchedMatrixRef<'_, Complex64>>,
batch_size: usize,
) -> Result<usize> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_ref_parts(u),
strided_batched_matrix_ref_parts(v),
batch_size,
)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_ptr(strided_batched_matrix_ref_parts(v), n, rank, jobz)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZgesvdaStridedBatched_bufferSize(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
&raw mut lwork,
to_i32(batch_size, "batch_size")?,
))?;
}
to_usize(lwork, "lwork")
}
pub fn sgesvda_strided_batched(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, f32>,
s: StridedBatchedVectorMut<'_, f32>,
u: Option<StridedBatchedMatrixMut<'_, f32>>,
v: Option<StridedBatchedMatrixMut<'_, f32>>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
residual: Option<&mut f64>,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_mut_ref_option(u.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
strided_batched_matrix_mut_ref_option(v.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = sgesvda_strided_batched_buffer_size(
ctx,
jobz,
rank,
m,
n,
a,
s.as_ref(),
strided_batched_matrix_mut_ref_option(u.as_ref()),
strided_batched_matrix_mut_ref_option(v.as_ref()),
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(v), n, rank, jobz)?;
unsafe {
try_ffi!(sys::cusolverDnSgesvdaStridedBatched(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_mut_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
residual.map_or(ptr::null_mut(), |value| value as *mut f64),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn dgesvda_strided_batched(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, f64>,
s: StridedBatchedVectorMut<'_, f64>,
u: Option<StridedBatchedMatrixMut<'_, f64>>,
v: Option<StridedBatchedMatrixMut<'_, f64>>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
residual: Option<&mut f64>,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_mut_ref_option(u.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
strided_batched_matrix_mut_ref_option(v.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = dgesvda_strided_batched_buffer_size(
ctx,
jobz,
rank,
m,
n,
a,
s.as_ref(),
strided_batched_matrix_mut_ref_option(u.as_ref()),
strided_batched_matrix_mut_ref_option(v.as_ref()),
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(v), n, rank, jobz)?;
unsafe {
try_ffi!(sys::cusolverDnDgesvdaStridedBatched(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_mut_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
residual.map_or(ptr::null_mut(), |value| value as *mut f64),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn cgesvda_strided_batched(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, Complex32>,
s: StridedBatchedVectorMut<'_, f32>,
u: Option<StridedBatchedMatrixMut<'_, Complex32>>,
v: Option<StridedBatchedMatrixMut<'_, Complex32>>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
residual: Option<&mut f64>,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_mut_ref_option(u.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
strided_batched_matrix_mut_ref_option(v.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = cgesvda_strided_batched_buffer_size(
ctx,
jobz,
rank,
m,
n,
a,
s.as_ref(),
strided_batched_matrix_mut_ref_option(u.as_ref()),
strided_batched_matrix_mut_ref_option(v.as_ref()),
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(v), n, rank, jobz)?;
unsafe {
try_ffi!(sys::cusolverDnCgesvdaStridedBatched(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_mut_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
residual.map_or(ptr::null_mut(), |value| value as *mut f64),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}
pub fn zgesvda_strided_batched(
ctx: &Context,
jobz: EigenMode,
rank: usize,
m: usize,
n: usize,
a: StridedBatchedMatrixRef<'_, Complex64>,
s: StridedBatchedVectorMut<'_, f64>,
u: Option<StridedBatchedMatrixMut<'_, Complex64>>,
v: Option<StridedBatchedMatrixMut<'_, Complex64>>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
residual: Option<&mut f64>,
batch_size: usize,
) -> Result<()> {
ctx.bind()?;
validate_gesvda_strided_batched_inputs(
rank,
m,
n,
a.data.len(),
a.leading_dimension,
a.stride,
s.data.len(),
s.stride,
jobz,
strided_batched_matrix_mut_ref_option(u.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
strided_batched_matrix_mut_ref_option(v.as_ref())
.map(|m| (m.data, m.leading_dimension, m.stride)),
batch_size,
)?;
require_info_buffer_len(dev_info, batch_size)?;
let lwork = zgesvda_strided_batched_buffer_size(
ctx,
jobz,
rank,
m,
n,
a,
s.as_ref(),
strided_batched_matrix_mut_ref_option(u.as_ref()),
strided_batched_matrix_mut_ref_option(v.as_ref()),
batch_size,
)?;
require_workspace(workspace.len(), lwork)?;
let (u_ptr, ldu, stride_u) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(u), m, rank, jobz)?;
let (v_ptr, ldv, stride_v) =
optional_gesvda_output_mut_ptr(strided_batched_matrix_mut_parts(v), n, rank, jobz)?;
unsafe {
try_ffi!(sys::cusolverDnZgesvdaStridedBatched(
ctx.as_raw(),
jobz.into(),
to_i32(rank, "rank")?,
to_i32(m, "m")?,
to_i32(n, "n")?,
a.data.as_ptr().cast(),
to_i32(a.leading_dimension, "lda")?,
to_i64(a.stride, "stride_a")?,
s.data.as_mut_ptr().cast(),
to_i64(s.stride, "stride_s")?,
u_ptr.cast(),
ldu,
stride_u,
v_ptr.cast(),
ldv,
stride_v,
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
residual.map_or(ptr::null_mut(), |value| value as *mut f64),
to_i32(batch_size, "batch_size")?,
))?;
}
Ok(())
}