use std::ptr;
use singe_cuda::{
data_type::{DataType, DataTypeLike},
memory::DeviceMemory,
types::{Complex32, Complex64},
};
use crate::{
context::Context,
eigen::{
EigenSelection,
info::SyevjInfo,
validation::{
matrix_mut_parts, matrix_mut_ref_option, matrix_mut_ref_parts, matrix_ref_parts,
optional_xgeev_matrix_mut_ptr, optional_xgeev_matrix_ptr, require_host_workspace,
require_info_buffer, require_info_buffer_len, require_workspace,
require_workspace_bytes, selection_parts, validate_syev_buffers,
validate_syevj_batched_buffers, validate_sygvd_buffers, validate_sygvj_buffers,
validate_xgeev_inputs, validate_xsyev_batched_buffers, validate_xsyevd_buffers,
validate_xsyevdx_range, validate_xsyevdx_value_type,
},
},
error::Result,
layout::{ByteWorkspaceMut, MatrixMut, MatrixRef, SelectionWorkspaceSizes, WorkspaceSizes},
params::Params,
sys, try_ffi,
types::{EigenMode, EigenRange, EigenType, FillMode},
utility::{to_i32, to_i64, to_usize},
};
pub(crate) fn xsyevd_buffer_size<TA: DataTypeLike, TW: DataTypeLike>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<TA>,
lda: usize,
w: &DeviceMemory<TW>,
) -> Result<WorkspaceSizes> {
xsyevd_raw_buffer_size(
ctx,
params,
mode,
fill_mode,
n,
TA::data_type(),
a,
lda,
TW::data_type(),
w,
TA::data_type(),
)
}
pub(crate) fn xsyevd<TA: DataTypeLike, TW: DataTypeLike>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<TA>,
lda: usize,
w: &mut DeviceMemory<TW>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
xsyevd_raw(
ctx,
params,
mode,
fill_mode,
n,
TA::data_type(),
a,
lda,
TW::data_type(),
w,
TA::data_type(),
workspace,
dev_info,
)
}
pub(crate) fn xsyevdx_buffer_size<TA: DataTypeLike, TR: Copy + Default, TW: DataTypeLike>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<TR>,
n: usize,
a: &DeviceMemory<TA>,
lda: usize,
w: &DeviceMemory<TW>,
) -> Result<SelectionWorkspaceSizes> {
let (range, value_range, index_range) = selection_parts(selection);
xsyevdx_raw_buffer_size(
ctx,
params,
mode,
range,
fill_mode,
n,
TA::data_type(),
a,
lda,
value_range,
index_range,
TW::data_type(),
w,
TA::data_type(),
)
}
pub(crate) fn xsyevdx<TA: DataTypeLike, TR: Copy + Default, TW: DataTypeLike>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<TR>,
n: usize,
a: &mut DeviceMemory<TA>,
lda: usize,
w: &mut DeviceMemory<TW>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<usize> {
let (range, value_range, index_range) = selection_parts(selection);
xsyevdx_raw(
ctx,
params,
mode,
range,
fill_mode,
n,
TA::data_type(),
a,
lda,
value_range,
index_range,
TW::data_type(),
w,
TA::data_type(),
workspace,
dev_info,
)
}
pub(crate) fn xsyev_batched_buffer_size<TA: DataTypeLike, TW: DataTypeLike>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: MatrixRef<'_, TA>,
w: &DeviceMemory<TW>,
batch_count: usize,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
validate_xsyev_batched_buffers(
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
w.byte_len(),
TW::data_type(),
batch_count,
)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXsyevBatched_bufferSize(
ctx.as_raw(),
params.as_raw(),
mode.into(),
fill_mode.into(),
to_i64(n, "n")?,
TA::data_type().into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
TW::data_type().into(),
w.as_ptr().cast(),
TA::data_type().into(),
&raw mut device_bytes,
&raw mut host_bytes,
to_i64(batch_count, "batch_count")?,
))?;
}
Ok(WorkspaceSizes::new(
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
pub(crate) fn xsyev_batched<TA: DataTypeLike, TW: DataTypeLike>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: MatrixMut<'_, TA>,
w: &mut DeviceMemory<TW>,
batch_count: usize,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_xsyev_batched_buffers(
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
w.byte_len(),
TW::data_type(),
batch_count,
)?;
require_info_buffer_len(dev_info, batch_count)?;
let workspace_sizes =
xsyev_batched_buffer_size(ctx, params, mode, fill_mode, n, a.as_ref(), w, batch_count)?;
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::cusolverDnXsyevBatched(
ctx.as_raw(),
params.as_raw(),
mode.into(),
fill_mode.into(),
to_i64(n, "n")?,
TA::data_type().into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
TW::data_type().into(),
w.as_mut_ptr().cast(),
TA::data_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(),
to_i64(batch_count, "batch_count")?,
))?;
}
Ok(())
}
pub(crate) fn xgeev_buffer_size<TA: DataTypeLike, TW: DataTypeLike, TV: DataTypeLike>(
ctx: &Context,
params: &Params,
n: usize,
a: MatrixRef<'_, TA>,
eigenvalues: &DeviceMemory<TW>,
right_vectors: Option<MatrixRef<'_, TV>>,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
validate_xgeev_inputs(
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
eigenvalues.byte_len(),
TW::data_type(),
matrix_ref_parts(right_vectors),
TV::data_type(),
)?;
let (vr_ptr, ldvr) = optional_xgeev_matrix_ptr(matrix_ref_parts(right_vectors))?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXgeev_bufferSize(
ctx.as_raw(),
params.as_raw(),
EigenMode::NoVector.into(),
if right_vectors.is_some() {
EigenMode::Vector
} else {
EigenMode::NoVector
}
.into(),
to_i64(n, "n")?,
TA::data_type().into(),
a.data.as_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
TW::data_type().into(),
eigenvalues.as_ptr().cast(),
TA::data_type().into(),
ptr::null(),
1,
TV::data_type().into(),
vr_ptr.cast(),
ldvr,
TA::data_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(crate) fn xgeev<TA: DataTypeLike, TW: DataTypeLike, TV: DataTypeLike>(
ctx: &Context,
params: &Params,
n: usize,
a: MatrixMut<'_, TA>,
eigenvalues: &mut DeviceMemory<TW>,
right_vectors: Option<MatrixMut<'_, TV>>,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_xgeev_inputs(
n,
a.data.byte_len(),
a.leading_dimension,
TA::data_type(),
eigenvalues.byte_len(),
TW::data_type(),
matrix_mut_ref_parts(right_vectors.as_ref()),
TV::data_type(),
)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xgeev_buffer_size(
ctx,
params,
n,
a.as_ref(),
eigenvalues,
matrix_mut_ref_option(right_vectors.as_ref()),
)?;
require_workspace_bytes(workspace.device.byte_len(), workspace_sizes.device_bytes)?;
require_host_workspace(workspace.host.len(), workspace_sizes.host_bytes)?;
let (vr_ptr, ldvr) = optional_xgeev_matrix_mut_ptr(matrix_mut_parts(right_vectors))?;
unsafe {
try_ffi!(sys::cusolverDnXgeev(
ctx.as_raw(),
params.as_raw(),
EigenMode::NoVector.into(),
if vr_ptr.is_null() {
EigenMode::NoVector
} else {
EigenMode::Vector
}
.into(),
to_i64(n, "n")?,
TA::data_type().into(),
a.data.as_mut_ptr().cast(),
to_i64(a.leading_dimension, "lda")?,
TW::data_type().into(),
eigenvalues.as_mut_ptr().cast(),
TA::data_type().into(),
ptr::null_mut(),
1,
TV::data_type().into(),
vr_ptr.cast(),
ldvr,
TA::data_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(())
}
fn xsyevd_raw_buffer_size<TA, TW>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a_type: DataType,
a: &DeviceMemory<TA>,
lda: usize,
w_type: DataType,
w: &DeviceMemory<TW>,
compute_type: DataType,
) -> Result<WorkspaceSizes> {
ctx.bind()?;
validate_xsyevd_buffers(n, a.byte_len(), lda, a_type, w.byte_len(), w_type)?;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXsyevd_bufferSize(
ctx.as_raw(),
params.as_raw(),
mode.into(),
fill_mode.into(),
to_i64(n, "n")?,
a_type.into(),
a.as_ptr().cast(),
to_i64(lda, "lda")?,
w_type.into(),
w.as_ptr().cast(),
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")?,
))
}
fn xsyevd_raw<TA, TW>(
ctx: &Context,
params: &Params,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a_type: DataType,
a: &mut DeviceMemory<TA>,
lda: usize,
w_type: DataType,
w: &mut DeviceMemory<TW>,
compute_type: DataType,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_xsyevd_buffers(n, a.byte_len(), lda, a_type, w.byte_len(), w_type)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xsyevd_raw_buffer_size(
ctx,
params,
mode,
fill_mode,
n,
a_type,
a,
lda,
w_type,
w,
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::cusolverDnXsyevd(
ctx.as_raw(),
params.as_raw(),
mode.into(),
fill_mode.into(),
to_i64(n, "n")?,
a_type.into(),
a.as_mut_ptr().cast(),
to_i64(lda, "lda")?,
w_type.into(),
w.as_mut_ptr().cast(),
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(())
}
fn xsyevdx_raw_buffer_size<TA, TR, TW>(
ctx: &Context,
params: &Params,
mode: EigenMode,
range: EigenRange,
fill_mode: FillMode,
n: usize,
a_type: DataType,
a: &DeviceMemory<TA>,
lda: usize,
value_range: Option<(TR, TR)>,
index_range: Option<(usize, usize)>,
w_type: DataType,
w: &DeviceMemory<TW>,
compute_type: DataType,
) -> Result<SelectionWorkspaceSizes>
where
TR: Copy + Default,
{
ctx.bind()?;
validate_xsyevd_buffers(n, a.byte_len(), lda, a_type, w.byte_len(), w_type)?;
validate_xsyevdx_value_type::<TR>(w_type)?;
let (mut vl, mut vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig = 0;
let mut device_bytes = 0;
let mut host_bytes = 0;
unsafe {
try_ffi!(sys::cusolverDnXsyevdx_bufferSize(
ctx.as_raw(),
params.as_raw(),
mode.into(),
range.into(),
fill_mode.into(),
to_i64(n, "n")?,
a_type.into(),
a.as_ptr().cast(),
to_i64(lda, "lda")?,
(&raw mut vl).cast(),
(&raw mut vu).cast(),
to_i64(il, "il")?,
to_i64(iu, "iu")?,
&raw mut meig,
w_type.into(),
w.as_ptr().cast(),
compute_type.into(),
&raw mut device_bytes,
&raw mut host_bytes,
))?;
}
Ok(SelectionWorkspaceSizes::new(
to_usize(meig, "meig")?,
to_usize(device_bytes, "device workspace size")?,
to_usize(host_bytes, "host workspace size")?,
))
}
fn xsyevdx_raw<TA, TR, TW>(
ctx: &Context,
params: &Params,
mode: EigenMode,
range: EigenRange,
fill_mode: FillMode,
n: usize,
a_type: DataType,
a: &mut DeviceMemory<TA>,
lda: usize,
value_range: Option<(TR, TR)>,
index_range: Option<(usize, usize)>,
w_type: DataType,
w: &mut DeviceMemory<TW>,
compute_type: DataType,
workspace: ByteWorkspaceMut<'_>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<usize>
where
TR: Copy + Default,
{
ctx.bind()?;
validate_xsyevd_buffers(n, a.byte_len(), lda, a_type, w.byte_len(), w_type)?;
validate_xsyevdx_value_type::<TR>(w_type)?;
require_info_buffer(dev_info)?;
let workspace_sizes = xsyevdx_raw_buffer_size(
ctx,
params,
mode,
range,
fill_mode,
n,
a_type,
a,
lda,
value_range,
index_range,
w_type,
w,
compute_type,
)?;
require_workspace_bytes(
workspace.device.byte_len(),
workspace_sizes.workspace.device_bytes,
)?;
require_host_workspace(workspace.host.len(), workspace_sizes.workspace.host_bytes)?;
let (mut vl, mut vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig_raw = 0;
unsafe {
try_ffi!(sys::cusolverDnXsyevdx(
ctx.as_raw(),
params.as_raw(),
mode.into(),
range.into(),
fill_mode.into(),
to_i64(n, "n")?,
a_type.into(),
a.as_mut_ptr().cast(),
to_i64(lda, "lda")?,
(&raw mut vl).cast(),
(&raw mut vu).cast(),
to_i64(il, "il")?,
to_i64(iu, "iu")?,
&raw mut meig_raw,
w_type.into(),
w.as_mut_ptr().cast(),
compute_type.into(),
workspace.device.as_mut_ptr().cast(),
workspace_sizes.workspace.device_bytes as _,
workspace.host.as_mut_ptr().cast(),
workspace_sizes.workspace.host_bytes as _,
dev_info.as_mut_ptr().cast(),
))?;
}
debug_assert_eq!(workspace_sizes.selection_size, to_usize(meig_raw, "meig")?);
Ok(workspace_sizes.selection_size)
}
pub(crate) fn ssyevj_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f32>,
lda: usize,
w: &DeviceMemory<f32>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSsyevj_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn dsyevj_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f64>,
lda: usize,
w: &DeviceMemory<f64>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDsyevj_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn cheevj_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex32>,
lda: usize,
w: &DeviceMemory<f32>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnCheevj_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn zheevj_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex64>,
lda: usize,
w: &DeviceMemory<f64>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZheevj_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn ssyevj(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f32>,
lda: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
require_info_buffer(dev_info)?;
let lwork = ssyevj_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnSsyevj(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn dsyevj(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f64>,
lda: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
require_info_buffer(dev_info)?;
let lwork = dsyevj_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnDsyevj(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn cheevj(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex32>,
lda: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
require_info_buffer(dev_info)?;
let lwork = cheevj_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnCheevj(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn zheevj(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex64>,
lda: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_syev_buffers(n, a.len(), lda, w.len())?;
require_info_buffer(dev_info)?;
let lwork = zheevj_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnZheevj(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn ssyevj_batched_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f32>,
lda: usize,
w: &DeviceMemory<f32>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<usize> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSsyevjBatched_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn dsyevj_batched_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f64>,
lda: usize,
w: &DeviceMemory<f64>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<usize> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDsyevjBatched_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn cheevj_batched_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex32>,
lda: usize,
w: &DeviceMemory<f32>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<usize> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnCheevjBatched_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn zheevj_batched_buffer_size(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex64>,
lda: usize,
w: &DeviceMemory<f64>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<usize> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZheevjBatched_bufferSize(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn ssyevj_batched(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f32>,
lda: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<()> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
require_info_buffer_len(dev_info, batch_count)?;
let lwork =
ssyevj_batched_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params, batch_count)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnSsyevjBatched(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
Ok(())
}
pub(crate) fn dsyevj_batched(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f64>,
lda: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<()> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
require_info_buffer_len(dev_info, batch_count)?;
let lwork =
dsyevj_batched_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params, batch_count)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnDsyevjBatched(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
Ok(())
}
pub(crate) fn cheevj_batched(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex32>,
lda: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<()> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
require_info_buffer_len(dev_info, batch_count)?;
let lwork =
cheevj_batched_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params, batch_count)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnCheevjBatched(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
Ok(())
}
pub(crate) fn zheevj_batched(
ctx: &Context,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex64>,
lda: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
batch_count: usize,
) -> Result<()> {
ctx.bind()?;
validate_syevj_batched_buffers(n, a.len(), lda, w.len(), batch_count)?;
require_info_buffer_len(dev_info, batch_count)?;
let lwork =
zheevj_batched_buffer_size(ctx, mode, fill_mode, n, a, lda, w, params, batch_count)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnZheevjBatched(
ctx.as_raw(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
to_i32(batch_count, "batch_count")?,
))?;
}
Ok(())
}
pub(crate) fn ssygvj_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f32>,
lda: usize,
b: &DeviceMemory<f32>,
ldb: usize,
w: &DeviceMemory<f32>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSsygvj_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn dsygvj_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f64>,
lda: usize,
b: &DeviceMemory<f64>,
ldb: usize,
w: &DeviceMemory<f64>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDsygvj_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn chegvj_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex32>,
lda: usize,
b: &DeviceMemory<Complex32>,
ldb: usize,
w: &DeviceMemory<f32>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnChegvj_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn zhegvj_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex64>,
lda: usize,
b: &DeviceMemory<Complex64>,
ldb: usize,
w: &DeviceMemory<f64>,
params: &SyevjInfo,
) -> Result<usize> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZhegvj_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
params.as_raw(),
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn ssygvj(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f32>,
lda: usize,
b: &mut DeviceMemory<f32>,
ldb: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = ssygvj_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnSsygvj(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn dsygvj(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f64>,
lda: usize,
b: &mut DeviceMemory<f64>,
ldb: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = dsygvj_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnDsygvj(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn chegvj(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex32>,
lda: usize,
b: &mut DeviceMemory<Complex32>,
ldb: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = chegvj_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnChegvj(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn zhegvj(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex64>,
lda: usize,
b: &mut DeviceMemory<Complex64>,
ldb: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
params: &SyevjInfo,
) -> Result<()> {
ctx.bind()?;
validate_sygvj_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = zhegvj_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w, params)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnZhegvj(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
params.as_raw(),
))?;
}
Ok(())
}
pub(crate) fn ssygvd_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f32>,
lda: usize,
b: &DeviceMemory<f32>,
ldb: usize,
w: &DeviceMemory<f32>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSsygvd_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn dsygvd_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<f64>,
lda: usize,
b: &DeviceMemory<f64>,
ldb: usize,
w: &DeviceMemory<f64>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDsygvd_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn chegvd_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex32>,
lda: usize,
b: &DeviceMemory<Complex32>,
ldb: usize,
w: &DeviceMemory<f32>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnChegvd_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn zhegvd_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &DeviceMemory<Complex64>,
lda: usize,
b: &DeviceMemory<Complex64>,
ldb: usize,
w: &DeviceMemory<f64>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZhegvd_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
to_usize(lwork, "lwork")
}
pub(crate) fn ssygvd(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f32>,
lda: usize,
b: &mut DeviceMemory<f32>,
ldb: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = ssygvd_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnSsygvd(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub(crate) fn dsygvd(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<f64>,
lda: usize,
b: &mut DeviceMemory<f64>,
ldb: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = dsygvd_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnDsygvd(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub(crate) fn chegvd(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex32>,
lda: usize,
b: &mut DeviceMemory<Complex32>,
ldb: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = chegvd_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnChegvd(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub(crate) fn zhegvd(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
n: usize,
a: &mut DeviceMemory<Complex64>,
lda: usize,
b: &mut DeviceMemory<Complex64>,
ldb: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<()> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let lwork = zhegvd_buffer_size(ctx, eig_type, mode, fill_mode, n, a, lda, b, ldb, w)?;
require_workspace(workspace.len(), lwork)?;
unsafe {
try_ffi!(sys::cusolverDnZhegvd(
ctx.as_raw(),
eig_type.into(),
mode.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
Ok(())
}
pub(crate) fn ssygvdx_selected_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f32>,
n: usize,
a: &DeviceMemory<f32>,
lda: usize,
b: &DeviceMemory<f32>,
ldb: usize,
w: &DeviceMemory<f32>,
) -> Result<(usize, usize)> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let (range, value_range, index_range) = selection_parts(selection);
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig = 0;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnSsygvdx_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
Ok((to_usize(meig, "meig")?, to_usize(lwork, "lwork")?))
}
pub(crate) fn dsygvdx_selected_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f64>,
n: usize,
a: &DeviceMemory<f64>,
lda: usize,
b: &DeviceMemory<f64>,
ldb: usize,
w: &DeviceMemory<f64>,
) -> Result<(usize, usize)> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let (range, value_range, index_range) = selection_parts(selection);
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig = 0;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnDsygvdx_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
Ok((to_usize(meig, "meig")?, to_usize(lwork, "lwork")?))
}
pub(crate) fn chegvdx_selected_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f32>,
n: usize,
a: &DeviceMemory<Complex32>,
lda: usize,
b: &DeviceMemory<Complex32>,
ldb: usize,
w: &DeviceMemory<f32>,
) -> Result<(usize, usize)> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let (range, value_range, index_range) = selection_parts(selection);
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig = 0;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnChegvdx_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
Ok((to_usize(meig, "meig")?, to_usize(lwork, "lwork")?))
}
pub(crate) fn zhegvdx_selected_buffer_size(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f64>,
n: usize,
a: &DeviceMemory<Complex64>,
lda: usize,
b: &DeviceMemory<Complex64>,
ldb: usize,
w: &DeviceMemory<f64>,
) -> Result<(usize, usize)> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
let (range, value_range, index_range) = selection_parts(selection);
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig = 0;
let mut lwork = 0;
unsafe {
try_ffi!(sys::cusolverDnZhegvdx_bufferSize(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_ptr().cast(),
to_i32(lda, "lda")?,
b.as_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig,
w.as_ptr().cast(),
&raw mut lwork,
))?;
}
Ok((to_usize(meig, "meig")?, to_usize(lwork, "lwork")?))
}
pub(crate) fn ssygvdx_selected(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f32>,
n: usize,
a: &mut DeviceMemory<f32>,
lda: usize,
b: &mut DeviceMemory<f32>,
ldb: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<f32>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let (range, value_range, index_range) = selection_parts(selection);
let (meig, lwork) = ssygvdx_selected_buffer_size(
ctx, eig_type, mode, fill_mode, selection, n, a, lda, b, ldb, w,
)?;
require_workspace(workspace.len(), lwork)?;
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig_raw = 0;
unsafe {
try_ffi!(sys::cusolverDnSsygvdx(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig_raw,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
debug_assert_eq!(meig, to_usize(meig_raw, "meig")?);
Ok(meig)
}
pub(crate) fn dsygvdx_selected(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f64>,
n: usize,
a: &mut DeviceMemory<f64>,
lda: usize,
b: &mut DeviceMemory<f64>,
ldb: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<f64>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let (range, value_range, index_range) = selection_parts(selection);
let (meig, lwork) = dsygvdx_selected_buffer_size(
ctx, eig_type, mode, fill_mode, selection, n, a, lda, b, ldb, w,
)?;
require_workspace(workspace.len(), lwork)?;
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig_raw = 0;
unsafe {
try_ffi!(sys::cusolverDnDsygvdx(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig_raw,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
debug_assert_eq!(meig, to_usize(meig_raw, "meig")?);
Ok(meig)
}
pub(crate) fn chegvdx_selected(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f32>,
n: usize,
a: &mut DeviceMemory<Complex32>,
lda: usize,
b: &mut DeviceMemory<Complex32>,
ldb: usize,
w: &mut DeviceMemory<f32>,
workspace: &mut DeviceMemory<Complex32>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let (range, value_range, index_range) = selection_parts(selection);
let (meig, lwork) = chegvdx_selected_buffer_size(
ctx, eig_type, mode, fill_mode, selection, n, a, lda, b, ldb, w,
)?;
require_workspace(workspace.len(), lwork)?;
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig_raw = 0;
unsafe {
try_ffi!(sys::cusolverDnChegvdx(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig_raw,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
debug_assert_eq!(meig, to_usize(meig_raw, "meig")?);
Ok(meig)
}
pub(crate) fn zhegvdx_selected(
ctx: &Context,
eig_type: EigenType,
mode: EigenMode,
fill_mode: FillMode,
selection: EigenSelection<f64>,
n: usize,
a: &mut DeviceMemory<Complex64>,
lda: usize,
b: &mut DeviceMemory<Complex64>,
ldb: usize,
w: &mut DeviceMemory<f64>,
workspace: &mut DeviceMemory<Complex64>,
dev_info: &mut DeviceMemory<i32>,
) -> Result<usize> {
ctx.bind()?;
validate_sygvd_buffers(n, a.len(), lda, b.len(), ldb, w.len())?;
require_info_buffer(dev_info)?;
let (range, value_range, index_range) = selection_parts(selection);
let (meig, lwork) = zhegvdx_selected_buffer_size(
ctx, eig_type, mode, fill_mode, selection, n, a, lda, b, ldb, w,
)?;
require_workspace(workspace.len(), lwork)?;
let (vl, vu, il, iu) = validate_xsyevdx_range(range, n, value_range, index_range)?;
let mut meig_raw = 0;
unsafe {
try_ffi!(sys::cusolverDnZhegvdx(
ctx.as_raw(),
eig_type.into(),
mode.into(),
range.into(),
fill_mode.into(),
to_i32(n, "n")?,
a.as_mut_ptr().cast(),
to_i32(lda, "lda")?,
b.as_mut_ptr().cast(),
to_i32(ldb, "ldb")?,
vl,
vu,
to_i32(il, "il")?,
to_i32(iu, "iu")?,
&raw mut meig_raw,
w.as_mut_ptr().cast(),
workspace.as_mut_ptr().cast(),
to_i32(lwork, "lwork")?,
dev_info.as_mut_ptr().cast(),
))?;
}
debug_assert_eq!(meig, to_usize(meig_raw, "meig")?);
Ok(meig)
}