1use std::fmt::{self, Display, Formatter};
2
3use singe_core::{impl_enum_conversion, impl_enum_display};
4
5use num_enum::{IntoPrimitive, TryFromPrimitive};
6use singe_cutensor_sys as sys;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
10#[repr(i32)]
11#[non_exhaustive]
12pub enum Algorithm {
13 DefaultPatient = sys::cutensorAlgo_t::CUTENSOR_ALGO_DEFAULT_PATIENT as _,
15 Gett = sys::cutensorAlgo_t::CUTENSOR_ALGO_GETT as _,
17 Tgett = sys::cutensorAlgo_t::CUTENSOR_ALGO_TGETT as _,
19 Ttgt = sys::cutensorAlgo_t::CUTENSOR_ALGO_TTGT as _,
21 Default = sys::cutensorAlgo_t::CUTENSOR_ALGO_DEFAULT as _,
23}
24
25impl_enum_conversion!(i32, sys::cutensorAlgo_t, Algorithm);
26
27impl_enum_display!(Algorithm, {
28 Algorithm::DefaultPatient => "CUTENSOR_ALGO_DEFAULT_PATIENT",
29 Algorithm::Gett => "CUTENSOR_ALGO_GETT",
30 Algorithm::Tgett => "CUTENSOR_ALGO_TGETT",
31 Algorithm::Ttgt => "CUTENSOR_ALGO_TTGT",
32 Algorithm::Default => "CUTENSOR_ALGO_DEFAULT",
33});
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
37#[repr(u32)]
38#[non_exhaustive]
39pub enum ComputeType {
40 F16 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_16F as _,
41 Bf16 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_16BF as _,
42 Tf32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_TF32 as _,
43 Tf32x3 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_3XTF32 as _,
44 F32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_32F as _,
45 F64 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_64F as _,
46 U8 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_8U as _,
47 I8 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_8I as _,
48 U32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_32U as _,
49 I32 = sys::cutensorComputeType_t::CUTENSOR_COMPUTE_32I as _,
50}
51
52impl_enum_conversion!(sys::cutensorComputeType_t, ComputeType);
53
54impl_enum_display!(ComputeType, {
55 ComputeType::F16 => "CUTENSOR_COMPUTE_16F",
56 ComputeType::Bf16 => "CUTENSOR_COMPUTE_16BF",
57 ComputeType::Tf32 => "CUTENSOR_COMPUTE_TF32",
58 ComputeType::Tf32x3 => "CUTENSOR_COMPUTE_3XTF32",
59 ComputeType::F32 => "CUTENSOR_COMPUTE_32F",
60 ComputeType::F64 => "CUTENSOR_COMPUTE_64F",
61 ComputeType::U8 => "CUTENSOR_COMPUTE_8U",
62 ComputeType::I8 => "CUTENSOR_COMPUTE_8I",
63 ComputeType::U32 => "CUTENSOR_COMPUTE_32U",
64 ComputeType::I32 => "CUTENSOR_COMPUTE_32I",
65});
66
67#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
69#[repr(u32)]
70#[non_exhaustive]
71pub enum Operator {
72 Identity = sys::cutensorOperator_t::CUTENSOR_OP_IDENTITY as _,
74 Sqrt = sys::cutensorOperator_t::CUTENSOR_OP_SQRT as _,
76 Relu = sys::cutensorOperator_t::CUTENSOR_OP_RELU as _,
78 Conj = sys::cutensorOperator_t::CUTENSOR_OP_CONJ as _,
80 Rcp = sys::cutensorOperator_t::CUTENSOR_OP_RCP as _,
82 Sigmoid = sys::cutensorOperator_t::CUTENSOR_OP_SIGMOID as _,
84 Tanh = sys::cutensorOperator_t::CUTENSOR_OP_TANH as _,
86 Exp = sys::cutensorOperator_t::CUTENSOR_OP_EXP as _,
88 Log = sys::cutensorOperator_t::CUTENSOR_OP_LOG as _,
90 Abs = sys::cutensorOperator_t::CUTENSOR_OP_ABS as _,
92 Neg = sys::cutensorOperator_t::CUTENSOR_OP_NEG as _,
94 Sin = sys::cutensorOperator_t::CUTENSOR_OP_SIN as _,
96 Cos = sys::cutensorOperator_t::CUTENSOR_OP_COS as _,
98 Tan = sys::cutensorOperator_t::CUTENSOR_OP_TAN as _,
100 Sinh = sys::cutensorOperator_t::CUTENSOR_OP_SINH as _,
102 Cosh = sys::cutensorOperator_t::CUTENSOR_OP_COSH as _,
104 Asin = sys::cutensorOperator_t::CUTENSOR_OP_ASIN as _,
106 Acos = sys::cutensorOperator_t::CUTENSOR_OP_ACOS as _,
108 Atan = sys::cutensorOperator_t::CUTENSOR_OP_ATAN as _,
110 Asinh = sys::cutensorOperator_t::CUTENSOR_OP_ASINH as _,
112 Acosh = sys::cutensorOperator_t::CUTENSOR_OP_ACOSH as _,
114 Atanh = sys::cutensorOperator_t::CUTENSOR_OP_ATANH as _,
116 Ceil = sys::cutensorOperator_t::CUTENSOR_OP_CEIL as _,
118 Floor = sys::cutensorOperator_t::CUTENSOR_OP_FLOOR as _,
120 Mish = sys::cutensorOperator_t::CUTENSOR_OP_MISH as _,
122 Swish = sys::cutensorOperator_t::CUTENSOR_OP_SWISH as _,
124 SoftPlus = sys::cutensorOperator_t::CUTENSOR_OP_SOFT_PLUS as _,
126 SoftSign = sys::cutensorOperator_t::CUTENSOR_OP_SOFT_SIGN as _,
128 Add = sys::cutensorOperator_t::CUTENSOR_OP_ADD as _,
130 Mul = sys::cutensorOperator_t::CUTENSOR_OP_MUL as _,
132 Max = sys::cutensorOperator_t::CUTENSOR_OP_MAX as _,
134 Min = sys::cutensorOperator_t::CUTENSOR_OP_MIN as _,
136 Unknown = sys::cutensorOperator_t::CUTENSOR_OP_UNKNOWN as _,
138}
139
140impl_enum_conversion!(sys::cutensorOperator_t, Operator);
141
142impl_enum_display!(Operator, {
143 Operator::Identity => "CUTENSOR_OP_IDENTITY",
144 Operator::Sqrt => "CUTENSOR_OP_SQRT",
145 Operator::Relu => "CUTENSOR_OP_RELU",
146 Operator::Conj => "CUTENSOR_OP_CONJ",
147 Operator::Rcp => "CUTENSOR_OP_RCP",
148 Operator::Sigmoid => "CUTENSOR_OP_SIGMOID",
149 Operator::Tanh => "CUTENSOR_OP_TANH",
150 Operator::Exp => "CUTENSOR_OP_EXP",
151 Operator::Log => "CUTENSOR_OP_LOG",
152 Operator::Abs => "CUTENSOR_OP_ABS",
153 Operator::Neg => "CUTENSOR_OP_NEG",
154 Operator::Sin => "CUTENSOR_OP_SIN",
155 Operator::Cos => "CUTENSOR_OP_COS",
156 Operator::Tan => "CUTENSOR_OP_TAN",
157 Operator::Sinh => "CUTENSOR_OP_SINH",
158 Operator::Cosh => "CUTENSOR_OP_COSH",
159 Operator::Asin => "CUTENSOR_OP_ASIN",
160 Operator::Acos => "CUTENSOR_OP_ACOS",
161 Operator::Atan => "CUTENSOR_OP_ATAN",
162 Operator::Asinh => "CUTENSOR_OP_ASINH",
163 Operator::Acosh => "CUTENSOR_OP_ACOSH",
164 Operator::Atanh => "CUTENSOR_OP_ATANH",
165 Operator::Ceil => "CUTENSOR_OP_CEIL",
166 Operator::Floor => "CUTENSOR_OP_FLOOR",
167 Operator::Mish => "CUTENSOR_OP_MISH",
168 Operator::Swish => "CUTENSOR_OP_SWISH",
169 Operator::SoftPlus => "CUTENSOR_OP_SOFT_PLUS",
170 Operator::SoftSign => "CUTENSOR_OP_SOFT_SIGN",
171 Operator::Add => "CUTENSOR_OP_ADD",
172 Operator::Mul => "CUTENSOR_OP_MUL",
173 Operator::Max => "CUTENSOR_OP_MAX",
174 Operator::Min => "CUTENSOR_OP_MIN",
175 Operator::Unknown => "CUTENSOR_OP_UNKNOWN",
176});
177
178#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
180#[repr(u32)]
181#[non_exhaustive]
182pub enum WorkspacePreference {
183 Min = sys::cutensorWorksizePreference_t::CUTENSOR_WORKSPACE_MIN as _,
185 Default = sys::cutensorWorksizePreference_t::CUTENSOR_WORKSPACE_DEFAULT as _,
187 Max = sys::cutensorWorksizePreference_t::CUTENSOR_WORKSPACE_MAX as _,
189}
190
191impl_enum_conversion!(sys::cutensorWorksizePreference_t, WorkspacePreference);
192
193impl_enum_display!(WorkspacePreference, {
194 WorkspacePreference::Min => "CUTENSOR_WORKSPACE_MIN",
195 WorkspacePreference::Default => "CUTENSOR_WORKSPACE_DEFAULT",
196 WorkspacePreference::Max => "CUTENSOR_WORKSPACE_MAX",
197});
198
199#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
201#[repr(u32)]
202#[non_exhaustive]
203pub enum CacheMode {
204 None = sys::cutensorCacheMode_t::CUTENSOR_CACHE_MODE_NONE as _,
206 Pedantic = sys::cutensorCacheMode_t::CUTENSOR_CACHE_MODE_PEDANTIC as _,
208}
209
210impl_enum_conversion!(sys::cutensorCacheMode_t, CacheMode);
211
212impl_enum_display!(CacheMode, {
213 CacheMode::None => "CUTENSOR_CACHE_MODE_NONE",
214 CacheMode::Pedantic => "CUTENSOR_CACHE_MODE_PEDANTIC",
215});
216
217#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
219#[repr(u32)]
220#[non_exhaustive]
221pub enum OperationDescriptorAttribute {
222 Tag = sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_TAG as _,
225 ScalarType =
227 sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_SCALAR_TYPE as _,
228 Flops = sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_FLOPS as _,
230 MovedBytes =
232 sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_MOVED_BYTES as _,
233 PaddingLeft =
235 sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_PADDING_LEFT
236 as _,
237 PaddingRight =
239 sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_PADDING_RIGHT
240 as _,
241 PaddingValue =
243 sys::cutensorOperationDescriptorAttribute_t::CUTENSOR_OPERATION_DESCRIPTOR_PADDING_VALUE
244 as _,
245}
246
247impl_enum_conversion!(
248 sys::cutensorOperationDescriptorAttribute_t,
249 OperationDescriptorAttribute
250);
251
252impl_enum_display!(OperationDescriptorAttribute, {
253 OperationDescriptorAttribute::Tag => "CUTENSOR_OPERATION_DESCRIPTOR_TAG",
254 OperationDescriptorAttribute::ScalarType => "CUTENSOR_OPERATION_DESCRIPTOR_SCALAR_TYPE",
255 OperationDescriptorAttribute::Flops => "CUTENSOR_OPERATION_DESCRIPTOR_FLOPS",
256 OperationDescriptorAttribute::MovedBytes => "CUTENSOR_OPERATION_DESCRIPTOR_MOVED_BYTES",
257 OperationDescriptorAttribute::PaddingLeft => "CUTENSOR_OPERATION_DESCRIPTOR_PADDING_LEFT",
258 OperationDescriptorAttribute::PaddingRight => "CUTENSOR_OPERATION_DESCRIPTOR_PADDING_RIGHT",
259 OperationDescriptorAttribute::PaddingValue => "CUTENSOR_OPERATION_DESCRIPTOR_PADDING_VALUE",
260});
261
262#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
264#[repr(u32)]
265#[non_exhaustive]
266pub enum PlanPreferenceAttribute {
267 AutotuneMode =
269 sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_AUTOTUNE_MODE as _,
270 CacheMode = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_CACHE_MODE as _,
272 IncrementalCount =
274 sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_INCREMENTAL_COUNT as _,
275 Algorithm = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_ALGO as _,
277 KernelRank = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_KERNEL_RANK as _,
279 JitMode = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_JIT as _,
281 GpuArch = sys::cutensorPlanPreferenceAttribute_t::CUTENSOR_PLAN_PREFERENCE_GPU_ARCH as _,
285}
286
287impl_enum_conversion!(
288 sys::cutensorPlanPreferenceAttribute_t,
289 PlanPreferenceAttribute
290);
291
292impl_enum_display!(PlanPreferenceAttribute, {
293 PlanPreferenceAttribute::AutotuneMode => "CUTENSOR_PLAN_PREFERENCE_AUTOTUNE_MODE",
294 PlanPreferenceAttribute::CacheMode => "CUTENSOR_PLAN_PREFERENCE_CACHE_MODE",
295 PlanPreferenceAttribute::IncrementalCount => "CUTENSOR_PLAN_PREFERENCE_INCREMENTAL_COUNT",
296 PlanPreferenceAttribute::Algorithm => "CUTENSOR_PLAN_PREFERENCE_ALGO",
297 PlanPreferenceAttribute::KernelRank => "CUTENSOR_PLAN_PREFERENCE_KERNEL_RANK",
298 PlanPreferenceAttribute::JitMode => "CUTENSOR_PLAN_PREFERENCE_JIT",
299 PlanPreferenceAttribute::GpuArch => "CUTENSOR_PLAN_PREFERENCE_GPU_ARCH",
300});
301
302#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
304#[repr(u32)]
305#[non_exhaustive]
306pub enum AutotuneMode {
307 None = sys::cutensorAutotuneMode_t::CUTENSOR_AUTOTUNE_MODE_NONE as _,
310 Incremental = sys::cutensorAutotuneMode_t::CUTENSOR_AUTOTUNE_MODE_INCREMENTAL as _,
316}
317
318impl_enum_conversion!(sys::cutensorAutotuneMode_t, AutotuneMode);
319
320impl_enum_display!(AutotuneMode, {
321 AutotuneMode::None => "CUTENSOR_AUTOTUNE_MODE_NONE",
322 AutotuneMode::Incremental => "CUTENSOR_AUTOTUNE_MODE_INCREMENTAL",
323});
324
325#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
327#[repr(u32)]
328#[non_exhaustive]
329pub enum JitMode {
330 None = sys::cutensorJitMode_t::CUTENSOR_JIT_MODE_NONE as _,
332 Default = sys::cutensorJitMode_t::CUTENSOR_JIT_MODE_DEFAULT as _,
335}
336
337impl_enum_conversion!(sys::cutensorJitMode_t, JitMode);
338
339impl_enum_display!(JitMode, {
340 JitMode::None => "CUTENSOR_JIT_MODE_NONE",
341 JitMode::Default => "CUTENSOR_JIT_MODE_DEFAULT",
342});
343
344#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, TryFromPrimitive, IntoPrimitive)]
346#[repr(u32)]
347#[non_exhaustive]
348pub enum PlanAttribute {
349 RequiredWorkspace = sys::cutensorPlanAttribute_t::CUTENSOR_PLAN_REQUIRED_WORKSPACE as _,
351}
352
353impl_enum_conversion!(sys::cutensorPlanAttribute_t, PlanAttribute);
354
355impl_enum_display!(PlanAttribute, {
356 PlanAttribute::RequiredWorkspace => "CUTENSOR_PLAN_REQUIRED_WORKSPACE",
357});
358
359#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
360#[repr(i32)]
361#[non_exhaustive]
362pub enum LoggerLevel {
363 Off = 0,
364 Error = 1,
365 PerformanceTrace = 2,
366 PerformanceHints = 3,
367 HeuristicsTrace = 4,
368 ApiTrace = 5,
369}
370
371impl LoggerLevel {
372 pub const fn as_raw(self) -> i32 {
373 self as i32
374 }
375}
376
377bitflags::bitflags! {
378 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
379 pub struct LoggerMask: i32 {
380 const OFF = 0;
381 const ERROR = 1;
382 const PERFORMANCE_TRACE = 2;
383 const PERFORMANCE_HINTS = 4;
384 const HEURISTICS_TRACE = 8;
385 const API_TRACE = 16;
386 }
387}
388
389#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
390#[repr(transparent)]
391pub struct Mode(i32);
392
393impl Mode {
394 pub const fn new(raw: i32) -> Self {
395 Self(raw)
396 }
397
398 pub const fn from_char(mode: char) -> Self {
399 Self(mode as i32)
400 }
401
402 pub const fn as_raw(self) -> i32 {
403 self.0
404 }
405}
406
407impl From<char> for Mode {
408 fn from(value: char) -> Self {
409 Self::from_char(value)
410 }
411}
412
413impl From<i32> for Mode {
414 fn from(value: i32) -> Self {
415 Self::new(value)
416 }
417}
418
419impl From<Mode> for i32 {
420 fn from(value: Mode) -> Self {
421 value.as_raw()
422 }
423}
424
425impl Display for Mode {
426 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
427 if let Some(ch) = char::from_u32(self.0 as u32)
428 && !ch.is_control()
429 {
430 return write!(f, "{ch}");
431 }
432
433 write!(f, "{}", self.0)
434 }
435}