use std::fmt::{self, Display, Formatter};
use singe_core::{impl_enum_conversion, impl_enum_display};
use num_enum::{IntoPrimitive, TryFromPrimitive};
use singe_cutensor_sys as sys;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(i32)]
#[non_exhaustive]
pub enum Algorithm {
DefaultPatient = sys::cutensorAlgo_t::CUTENSOR_ALGO_DEFAULT_PATIENT as _,
Gett = sys::cutensorAlgo_t::CUTENSOR_ALGO_GETT as _,
Tgett = sys::cutensorAlgo_t::CUTENSOR_ALGO_TGETT as _,
Ttgt = sys::cutensorAlgo_t::CUTENSOR_ALGO_TTGT as _,
Default = sys::cutensorAlgo_t::CUTENSOR_ALGO_DEFAULT as _,
}
impl_enum_conversion!(i32, sys::cutensorAlgo_t, Algorithm);
impl_enum_display!(Algorithm, {
Algorithm::DefaultPatient => "CUTENSOR_ALGO_DEFAULT_PATIENT",
Algorithm::Gett => "CUTENSOR_ALGO_GETT",
Algorithm::Tgett => "CUTENSOR_ALGO_TGETT",
Algorithm::Ttgt => "CUTENSOR_ALGO_TTGT",
Algorithm::Default => "CUTENSOR_ALGO_DEFAULT",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum ComputeType {
F16 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_16F as _,
Bf16 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_16BF as _,
Tf32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_TF32 as _,
Tf32x3 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_3XTF32 as _,
F32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_32F as _,
F64 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_64F as _,
U8 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_8U as _,
I8 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_8I as _,
U32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_32U as _,
I32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_32I as _,
}
impl_enum_conversion!(sys::cutensorComputeType_t, ComputeType);
impl_enum_display!(ComputeType, {
ComputeType::F16 => "CUTENSOR_COMPUTE_16F",
ComputeType::Bf16 => "CUTENSOR_COMPUTE_16BF",
ComputeType::Tf32 => "CUTENSOR_COMPUTE_TF32",
ComputeType::Tf32x3 => "CUTENSOR_COMPUTE_3XTF32",
ComputeType::F32 => "CUTENSOR_COMPUTE_32F",
ComputeType::F64 => "CUTENSOR_COMPUTE_64F",
ComputeType::U8 => "CUTENSOR_COMPUTE_8U",
ComputeType::I8 => "CUTENSOR_COMPUTE_8I",
ComputeType::U32 => "CUTENSOR_COMPUTE_32U",
ComputeType::I32 => "CUTENSOR_COMPUTE_32I",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum Operator {
Identity = sys::cutensorOperator_t::CUTENSOR_OP_IDENTITY as _,
Sqrt = sys::cutensorOperator_t::CUTENSOR_OP_SQRT as _,
Relu = sys::cutensorOperator_t::CUTENSOR_OP_RELU as _,
Conj = sys::cutensorOperator_t::CUTENSOR_OP_CONJ as _,
Rcp = sys::cutensorOperator_t::CUTENSOR_OP_RCP as _,
Sigmoid = sys::cutensorOperator_t::CUTENSOR_OP_SIGMOID as _,
Tanh = sys::cutensorOperator_t::CUTENSOR_OP_TANH as _,
Exp = sys::cutensorOperator_t::CUTENSOR_OP_EXP as _,
Log = sys::cutensorOperator_t::CUTENSOR_OP_LOG as _,
Abs = sys::cutensorOperator_t::CUTENSOR_OP_ABS as _,
Neg = sys::cutensorOperator_t::CUTENSOR_OP_NEG as _,
Sin = sys::cutensorOperator_t::CUTENSOR_OP_SIN as _,
Cos = sys::cutensorOperator_t::CUTENSOR_OP_COS as _,
Tan = sys::cutensorOperator_t::CUTENSOR_OP_TAN as _,
Sinh = sys::cutensorOperator_t::CUTENSOR_OP_SINH as _,
Cosh = sys::cutensorOperator_t::CUTENSOR_OP_COSH as _,
Asin = sys::cutensorOperator_t::CUTENSOR_OP_ASIN as _,
Acos = sys::cutensorOperator_t::CUTENSOR_OP_ACOS as _,
Atan = sys::cutensorOperator_t::CUTENSOR_OP_ATAN as _,
Asinh = sys::cutensorOperator_t::CUTENSOR_OP_ASINH as _,
Acosh = sys::cutensorOperator_t::CUTENSOR_OP_ACOSH as _,
Atanh = sys::cutensorOperator_t::CUTENSOR_OP_ATANH as _,
Ceil = sys::cutensorOperator_t::CUTENSOR_OP_CEIL as _,
Floor = sys::cutensorOperator_t::CUTENSOR_OP_FLOOR as _,
Mish = sys::cutensorOperator_t::CUTENSOR_OP_MISH as _,
Swish = sys::cutensorOperator_t::CUTENSOR_OP_SWISH as _,
SoftPlus = sys::cutensorOperator_t::CUTENSOR_OP_SOFT_PLUS as _,
SoftSign = sys::cutensorOperator_t::CUTENSOR_OP_SOFT_SIGN as _,
Add = sys::cutensorOperator_t::CUTENSOR_OP_ADD as _,
Mul = sys::cutensorOperator_t::CUTENSOR_OP_MUL as _,
Max = sys::cutensorOperator_t::CUTENSOR_OP_MAX as _,
Min = sys::cutensorOperator_t::CUTENSOR_OP_MIN as _,
Unknown = sys::cutensorOperator_t::CUTENSOR_OP_UNKNOWN as _,
}
impl_enum_conversion!(sys::cutensorOperator_t, Operator);
impl_enum_display!(Operator, {
Operator::Identity => "CUTENSOR_OP_IDENTITY",
Operator::Sqrt => "CUTENSOR_OP_SQRT",
Operator::Relu => "CUTENSOR_OP_RELU",
Operator::Conj => "CUTENSOR_OP_CONJ",
Operator::Rcp => "CUTENSOR_OP_RCP",
Operator::Sigmoid => "CUTENSOR_OP_SIGMOID",
Operator::Tanh => "CUTENSOR_OP_TANH",
Operator::Exp => "CUTENSOR_OP_EXP",
Operator::Log => "CUTENSOR_OP_LOG",
Operator::Abs => "CUTENSOR_OP_ABS",
Operator::Neg => "CUTENSOR_OP_NEG",
Operator::Sin => "CUTENSOR_OP_SIN",
Operator::Cos => "CUTENSOR_OP_COS",
Operator::Tan => "CUTENSOR_OP_TAN",
Operator::Sinh => "CUTENSOR_OP_SINH",
Operator::Cosh => "CUTENSOR_OP_COSH",
Operator::Asin => "CUTENSOR_OP_ASIN",
Operator::Acos => "CUTENSOR_OP_ACOS",
Operator::Atan => "CUTENSOR_OP_ATAN",
Operator::Asinh => "CUTENSOR_OP_ASINH",
Operator::Acosh => "CUTENSOR_OP_ACOSH",
Operator::Atanh => "CUTENSOR_OP_ATANH",
Operator::Ceil => "CUTENSOR_OP_CEIL",
Operator::Floor => "CUTENSOR_OP_FLOOR",
Operator::Mish => "CUTENSOR_OP_MISH",
Operator::Swish => "CUTENSOR_OP_SWISH",
Operator::SoftPlus => "CUTENSOR_OP_SOFT_PLUS",
Operator::SoftSign => "CUTENSOR_OP_SOFT_SIGN",
Operator::Add => "CUTENSOR_OP_ADD",
Operator::Mul => "CUTENSOR_OP_MUL",
Operator::Max => "CUTENSOR_OP_MAX",
Operator::Min => "CUTENSOR_OP_MIN",
Operator::Unknown => "CUTENSOR_OP_UNKNOWN",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum WorkspacePreference {
Min = sys::cutensorWorksizePreference_t::CUTENSOR_WORKSPACE_MIN as _,
Default = sys::cutensorWorksizePreference_t::CUTENSOR_WORKSPACE_DEFAULT as _,
Max = sys::cutensorWorksizePreference_t::CUTENSOR_WORKSPACE_MAX as _,
}
impl_enum_conversion!(sys::cutensorWorksizePreference_t, WorkspacePreference);
impl_enum_display!(WorkspacePreference, {
WorkspacePreference::Min => "CUTENSOR_WORKSPACE_MIN",
WorkspacePreference::Default => "CUTENSOR_WORKSPACE_DEFAULT",
WorkspacePreference::Max => "CUTENSOR_WORKSPACE_MAX",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum CacheMode {
None = sys::cutensorCacheMode_t::CUTENSOR_CACHE_MODE_NONE as _,
Pedantic = sys::cutensorCacheMode_t::CUTENSOR_CACHE_MODE_PEDANTIC as _,
}
impl_enum_conversion!(sys::cutensorCacheMode_t, CacheMode);
impl_enum_display!(CacheMode, {
CacheMode::None => "CUTENSOR_CACHE_MODE_NONE",
CacheMode::Pedantic => "CUTENSOR_CACHE_MODE_PEDANTIC",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum OperationDescriptorAttribute {
Tag = sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_TAG as _,
ScalarType =
sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_SCALAR_TYPE as _,
Flops = sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_FLOPS as _,
MovedBytes =
sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_MOVED_BYTES as _,
PaddingLeft =
sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_PADDING_LEFT
as _,
PaddingRight =
sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_PADDING_RIGHT
as _,
PaddingValue =
sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_PADDING_VALUE
as _,
}
impl_enum_conversion!(
sys::cutensorOperationDescriptorAttribute_t,
OperationDescriptorAttribute
);
impl_enum_display!(OperationDescriptorAttribute, {
OperationDescriptorAttribute::Tag => "CUTENSOR_OPERATION_DESCRIPTOR_TAG",
OperationDescriptorAttribute::ScalarType => "CUTENSOR_OPERATION_DESCRIPTOR_SCALAR_TYPE",
OperationDescriptorAttribute::Flops => "CUTENSOR_OPERATION_DESCRIPTOR_FLOPS",
OperationDescriptorAttribute::MovedBytes => "CUTENSOR_OPERATION_DESCRIPTOR_MOVED_BYTES",
OperationDescriptorAttribute::PaddingLeft => "CUTENSOR_OPERATION_DESCRIPTOR_PADDING_LEFT",
OperationDescriptorAttribute::PaddingRight => "CUTENSOR_OPERATION_DESCRIPTOR_PADDING_RIGHT",
OperationDescriptorAttribute::PaddingValue => "CUTENSOR_OPERATION_DESCRIPTOR_PADDING_VALUE",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum PlanPreferenceAttribute {
AutotuneMode =
sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_AUTOTUNE_MODE as _,
CacheMode = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_CACHE_MODE as _,
IncrementalCount =
sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_INCREMENTAL_COUNT as _,
Algorithm = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_ALGO as _,
KernelRank = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_KERNEL_RANK as _,
JitMode = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_JIT as _,
GpuArch = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_GPU_ARCH as _,
}
impl_enum_conversion!(
sys::cutensorPlanPreferenceAttribute_t,
PlanPreferenceAttribute
);
impl_enum_display!(PlanPreferenceAttribute, {
PlanPreferenceAttribute::AutotuneMode => "CUTENSOR_PLAN_PREFERENCE_AUTOTUNE_MODE",
PlanPreferenceAttribute::CacheMode => "CUTENSOR_PLAN_PREFERENCE_CACHE_MODE",
PlanPreferenceAttribute::IncrementalCount => "CUTENSOR_PLAN_PREFERENCE_INCREMENTAL_COUNT",
PlanPreferenceAttribute::Algorithm => "CUTENSOR_PLAN_PREFERENCE_ALGO",
PlanPreferenceAttribute::KernelRank => "CUTENSOR_PLAN_PREFERENCE_KERNEL_RANK",
PlanPreferenceAttribute::JitMode => "CUTENSOR_PLAN_PREFERENCE_JIT",
PlanPreferenceAttribute::GpuArch => "CUTENSOR_PLAN_PREFERENCE_GPU_ARCH",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum AutotuneMode {
None = sys::cutensorAutotuneMode_t::CUTENSOR_AUTOTUNE_MODE_NONE as _,
Incremental = sys::cutensorAutotuneMode_t::CUTENSOR_AUTOTUNE_MODE_INCREMENTAL as _,
}
impl_enum_conversion!(sys::cutensorAutotuneMode_t, AutotuneMode);
impl_enum_display!(AutotuneMode, {
AutotuneMode::None => "CUTENSOR_AUTOTUNE_MODE_NONE",
AutotuneMode::Incremental => "CUTENSOR_AUTOTUNE_MODE_INCREMENTAL",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum JitMode {
None = sys::cutensorJitMode_t::CUTENSOR_JIT_MODE_NONE as _,
Default = sys::cutensorJitMode_t::CUTENSOR_JIT_MODE_DEFAULT as _,
}
impl_enum_conversion!(sys::cutensorJitMode_t, JitMode);
impl_enum_display!(JitMode, {
JitMode::None => "CUTENSOR_JIT_MODE_NONE",
JitMode::Default => "CUTENSOR_JIT_MODE_DEFAULT",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
#[repr(u32)]
#[non_exhaustive]
pub enum PlanAttribute {
RequiredWorkspace = sys::cutensorPlanAttribute_t::CUTENSOR_PLAN_REQUIRED_WORKSPACE as _,
}
impl_enum_conversion!(sys::cutensorPlanAttribute_t, PlanAttribute);
impl_enum_display!(PlanAttribute, {
PlanAttribute::RequiredWorkspace => "CUTENSOR_PLAN_REQUIRED_WORKSPACE",
});
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(i32)]
#[non_exhaustive]
pub enum LoggerLevel {
Off = 0,
Error = 1,
PerformanceTrace = 2,
PerformanceHints = 3,
HeuristicsTrace = 4,
ApiTrace = 5,
}
impl LoggerLevel {
pub const fn as_raw(self) -> i32 {
self as i32
}
}
bitflags::bitflags! {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct LoggerMask: i32 {
const OFF = 0;
const ERROR = 1;
const PERFORMANCE_TRACE = 2;
const PERFORMANCE_HINTS = 4;
const HEURISTICS_TRACE = 8;
const API_TRACE = 16;
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[repr(transparent)]
pub struct Mode(i32);
impl Mode {
pub const fn new(raw: i32) -> Self {
Self(raw)
}
pub const fn from_char(mode: char) -> Self {
Self(mode as i32)
}
pub const fn as_raw(self) -> i32 {
self.0
}
}
impl From<char> for Mode {
fn from(value: char) -> Self {
Self::from_char(value)
}
}
impl From<i32> for Mode {
fn from(value: i32) -> Self {
Self::new(value)
}
}
impl From<Mode> for i32 {
fn from(value: Mode) -> Self {
value.as_raw()
}
}
impl Display for Mode {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
if let Some(ch) = char::from_u32(self.0 as u32)
&& !ch.is_control()
{
return write!(f, "{ch}");
}
write!(f, "{}", self.0)
}
}