Skip to main content

muDNN API Reference

1 Introduction

The Moore Threads® MUSA® Deep Neural Network(muDNN) library is a GPU-accelerated library for common- used primitives in deep neural networks. muDNN provides highly optimized functions that can be used to perform various mathematical and data processing tasks, such as:

  • Tensor operations: element-wise operations, matrix operations, reduction operations, etc.
  • Neural network layers: convolution, pooling, normalization, activation, etc.
  • Loss functions: KLDivLoss, L2Loss, NLLLoss, etc. It enables users to focus on training neural networks and developing applications rather than accelerating GPU performance. muDNN library offers a context-based API that allows for easy multithreading with MUSA streams.

This API Reference lists the data type definitions and detailed descriptions of functions.

1.1 Functional requirements and assurances

  • Each API could be accessed safely and concurrently, i.e. they are reentrant. Users must ensure that the necessary input memory allocation and data preparation are ready before calling the API.
  • Most APIs (except specific ones with annotation) could reproduce the same result with identical configuration and input.
  • muDNN supports a maximum tensor dimension of size 8.

2 Module Index

2.1 Modules

Here is a list of all modules:

  • Base Operators
  • Image Operators
  • Math Operators
  • NN -(Neural Network) Operators
  • Version

3 Class Index

3.1 Class List

  • musa::dnn::BatchMatMul Here are the classes, structs, unions and interfaces with brief descriptions:
  • musa::dnn::BatchNorm
  • musa::dnn::Binary
  • musa::dnn::Concat
  • musa::dnn::Convolution
  • musa::dnn::CTCLoss
  • musa::dnn::Cum
  • musa::dnn::Cumsum
  • musa::dnn::DebugInfo
  • musa::dnn::DeformableConv
  • musa::dnn::Dot
  • musa::dnn::Dropout
  • musa::dnn::Fill
  • musa::dnn::Convolution::FusedActivationDesc
  • musa::dnn::GatherX
  • musa::dnn::Glu
  • musa::dnn::GroupNorm
  • musa::dnn::Handle
  • musa::dnn::ImplBase
  • musa::dnn::Interpolate
  • musa::dnn::KLDivLoss
  • musa::dnn::L2Loss
  • musa::dnn::LayerNorm
  • musa::dnn::LocalResponseNorm
  • musa::dnn::MaskedScatter
  • musa::dnn::MaskedSelect
  • musa::dnn::MatMul
  • musa::dnn::MatrixBase
  • musa::dnn::MultiHeadAttention
  • musa::dnn::NLLLoss
  • musa::dnn::Nonzero
  • musa::dnn::Pad
  • musa::dnn::Permute
  • musa::dnn::Pooling
  • musa::dnn::Reduce
  • musa::dnn::RMSNorm
  • musa::dnn::RNN
  • musa::dnn::ScaledDotProductAttention
  • musa::dnn::Scan
  • musa::dnn::Scatter
  • musa::dnn::ScatterND
  • musa::dnn::Softmax
  • musa::dnn::Sort
  • musa::dnn::SortByKey
  • musa::dnn::Tensor
  • musa::dnn::TensorBase
  • musa::dnn::Ternary
  • musa::dnn::TopK
  • musa::dnn::Unary
  • musa::dnn::Unfold
  • musa::dnn::Unique
  • musa::dnn::WeightNorm

4 Module Documentation

4.1 Base Operators

Classes

  • class musa::dnn::ImplBase
  • class musa::dnn::Handle
  • class musa::dnn::TensorBase
  • class musa::dnn::Tensor
  • class musa::dnn::MatrixBase
  • class musa::dnn::DebugInfo

Macros

  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x, ...) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x, y) x,

Typedefs

- using **musa::dnn::MemoryHandler** = std::unique_ptr<void, std::function<void(void∗)>>
- using **musa::dnn::MemoryMaintainer** = std::function<MemoryHandler(size_t)>
- using **musa::dnn::CallBack** = std::function<void(DebugInfo::Severity, void∗user_data, const DebugInfo
&dbg, const char∗message)>

Enumerations

- enum class **Status** { **MUDNN_ITEM** }
- enum class **Type** { **MUDNN_ITEM** }
- enum class **Format** { **MUDNN_ITEM** }
- enum class **ComputeMode** { **MUDNN_ITEM** }
- enum class **Severity** { **MUDNN_ITEM** }

Functions

- void**musa::dnn::ImplBase::GetImpl** ()
- const void**musa::dnn::ImplBase::GetImpl** () const
- **musa::dnn::ImplBase::ImplBase** (void∗impl)
- **musa::dnn::ImplBase::ImplBase** (const ImplBase &)=delete
- **musa::dnn::ImplBase::ImplBase** (ImplBase &&)=delete
- ImplBase & **musa::dnn::ImplBase::operator=** (const ImplBase &)=delete
- ImplBase & **musa::dnn::ImplBase::operator=** (ImplBase &&)=delete
- **musa::dnn::Handle::Handle** (int device_id)
- int **musa::dnn::Handle::GetDeviceId** () const
- Status **musa::dnn::Handle::SetStream** (musaStream_t stream)
- musaStream_t **musa::dnn::Handle::GetStream** () const
- Status **musa::dnn::Handle::SetAllowTF32** (bool allow_tf32)
- bool **musa::dnn::Handle::GetAllowTF32** () const
- **musa::dnn::Tensor::Tensor** (const Tensor &)
- **musa::dnn::Tensor::Tensor** (Tensor &&)
- Tensor & **musa::dnn::Tensor::operator=** (const Tensor &)
- Tensor & **musa::dnn::Tensor::operator=** (Tensor &&)
- Status **musa::dnn::Tensor::SetAddr** (const void∗addr)
- Status **musa::dnn::Tensor::SetType** (Type t)
- Status **musa::dnn::Tensor::SetFormat** (Format f)
- Status **musa::dnn::Tensor::SetNdInfo** (std::initializer_list<int64_t>dim)
- Status **musa::dnn::Tensor::SetNdInfo** (int ndims, const int64_t∗dim)
- Status **musa::dnn::Tensor::SetNdInfo** (int64_t ndims, const int64_t∗dim)
- Status **musa::dnn::Tensor::SetNdInfo** (std::initializer_list<int64_t>dim, std::initializer_list<int64_t>
stride)
- Status **musa::dnn::Tensor::SetNdInfo** (int ndims, const int64_t∗dim, const int64_t∗stride)
- Status **musa::dnn::Tensor::SetNdInfo** (int64_t ndims, const int64_t∗dim, const int64_t∗stride)
- Status **musa::dnn::Tensor::GetNdInfo** (std::vector<int64_t>&dim)
- Status **musa::dnn::Tensor::GetNdInfo** (std::vector<int64_t>&dim, std::vector<int64_t>&stride)
- Status **musa::dnn::Tensor::SetQuantizationInfo** (int n, const float∗scales, const unsigned int∗zero_←-
points)
- Status **musa::dnn::Tensor::SetQuantizationInfo** (std::initializer_list<float>scales, std::initializer_list<
unsigned int>zero_points)
- Status **musa::dnn::Tensor::SetQuantizationInfo** (const std::vector<float>&scales)
- Status **musa::dnn::Tensor::CopyFrom** (void∗ptr, size_t bytes, int kind, Handle &h, bool sync=false)
- Status **musa::dnn::Tensor::CopyTo** (void∗ptr, size_t bytes, int kind, Handle &h, bool sync=false) const
- MUDNN_EXPORT Status **musa::dnn::SetCallBack** (DebugInfo::Severity min, void∗udata, CallBack f)

Variables

  • void∗ musa::dnn::ImplBase::impl_
  • uint32_t musa::dnn::DebugInfo::version
  • Severity musa::dnn::DebugInfo::severity
  • uint32_t musa::dnn::DebugInfo::time_sec
  • uint32_t musa::dnn::DebugInfo::time_usec
  • uint32_t musa::dnn::DebugInfo::time_delta
  • uint64_t musa::dnn::DebugInfo::tid
  • int32_t musa::dnn::DebugInfo::device_id
  • const Handle∗ musa::dnn::DebugInfo::handle

4.1.1 Detailed Description

This entity contains basic functionality related to muDNN context creation and destruction, tensor utility routines, tensor core operations, debug logs, and so on.

4.2 Image Operators

Classes

  • class musa::dnn::Interpolate

Macros

  • #define MUDNN_ITEM (x) x,

Enumerations

- enum class **Mode** { **MUDNN_ITEM** }

Functions

- Status **musa::dnn::Interpolate::SetMode** (Mode m)
- Status **musa::dnn::Interpolate::SetScaleInfo** (std::initializer_list<float>scale)
- Status **musa::dnn::Interpolate::SetScaleInfo** (int length, float∗scale)
- Status **musa::dnn::Interpolate::SetAlignCorners** (bool align_corners)
- Status **musa::dnn::Interpolate::Run** (Handle &h, Tensor &out, const Tensor &in) const
- Status **musa::dnn::Interpolate::RunBackward** (Handle &h, Tensor &out, const Tensor &in) const

4.2.1 Detailed Description

This entity is a collection of image processing and manipulation operations, including resizing images, cropping, and interpolating.

4.3 Math Operators

Classes

  • class musa::dnn::Unary
  • class musa::dnn::Binary
  • class musa::dnn::Ternary
  • class musa::dnn::Reduce
  • class musa::dnn::BatchMatMul
  • class musa::dnn::MatMul
  • class musa::dnn::Dot
  • class musa::dnn::Concat
  • class musa::dnn::Permute
  • class musa::dnn::Fill
  • class musa::dnn::Sort
  • class musa::dnn::SortByKey
  • class musa::dnn::TopK
  • class musa::dnn::Scan
  • class musa::dnn::Cumsum
  • class musa::dnn::Cum
  • class musa::dnn::GatherX
  • class musa::dnn::MaskedSelect
  • class musa::dnn::MaskedScatter
  • class musa::dnn::Unfold
  • class musa::dnn::Unique
  • class musa::dnn::Nonzero

Macros

  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,

Enumerations

- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **ScanOpType** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }

Functions

- Status **musa::dnn::Unary::SetMode** (Mode m)
- Status **musa::dnn::Unary::SetAlpha** (double alpha)
- Status **musa::dnn::Unary::SetAlpha** (int64_t alpha)
- Status **musa::dnn::Unary::SetAlpha** (const void∗alpha)
- Status **musa::dnn::Unary::SetBeta** (double beta)
- Status **musa::dnn::Unary::SetBeta** (int64_t beta)
- Status **musa::dnn::Unary::SetBeta** (const void∗beta)
- Status **musa::dnn::Unary::Run** (Handle &h, Tensor &out, const Tensor &in) const
- Status **musa::dnn::Binary::SetMode** (Mode m)
- Status **musa::dnn::Binary::SetAlpha** (double alpha)
- Status **musa::dnn::Binary::SetAlpha** (int64_t alpha)
- Status **musa::dnn::Binary::SetAlpha** (const void∗alpha)
- Status **musa::dnn::Binary::SetBeta** (double beta)
- Status **musa::dnn::Binary::SetBeta** (int64_t beta)
- Status **musa::dnn::Binary::SetBeta** (const void∗beta)
- Status **musa::dnn::Binary::Run** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r) const
- Status **musa::dnn::Ternary::SetMode** (Mode m)
- Status **musa::dnn::Ternary::SetAlpha** (double alpha)
- Status **musa::dnn::Ternary::SetAlpha** (int64_t alpha)
- Status **musa::dnn::Ternary::SetAlpha** (const void∗alpha)
- Status **musa::dnn::Ternary::SetBeta** (double beta)
- Status **musa::dnn::Ternary::SetBeta** (int64_t beta)
- Status **musa::dnn::Ternary::SetBeta** (const void∗beta)
- Status **musa::dnn::Ternary::SetGamma** (double gamma)
- Status **musa::dnn::Ternary::SetGamma** (int64_t gamma)
- Status **musa::dnn::Ternary::SetGamma** (const void∗gamma)
- Status **musa::dnn::Ternary::Run** (Handle &h, Tensor &out, const Tensor &in0, const Tensor &in1, const
Tensor &in2) const
- Status **musa::dnn::Reduce::SetMode** (Mode m)
- Status **musa::dnn::Reduce::SetDim** (std::initializer_list<int>dim)
- Status **musa::dnn::Reduce::SetDim** (int ndim, const int∗dim)
- Status **musa::dnn::Reduce::SetNormOrd** (float ord)
- Status **musa::dnn::Reduce::GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out, const
Tensor &in)
- Status **musa::dnn::Reduce::Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer
&maintainer) const
- Status **musa::dnn::Reduce::RunIndices** (Handle &h, Tensor &out, const Tensor &in, const Memory←-
Maintainer &maintainer) const
- Status **musa::dnn::Reduce::RunWithIndices** (Handle &h, Tensor &out, Tensor &indices, const Tensor &in,
const MemoryMaintainer &maintainer) const
- Status **musa::dnn::BatchMatMul::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::BatchMatMul::SetTranspose** (bool left, bool right)
- Status **musa::dnn::BatchMatMul::SetSplitK** (bool split_k)
- Status **musa::dnn::BatchMatMul::SetAlpha** (double alpha)
- Status **musa::dnn::BatchMatMul::SetBeta** (double beta)
- Status **musa::dnn::BatchMatMul::SetGamma** (double gamma)
- Status **musa::dnn::BatchMatMul::GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out,
const Tensor &l, const Tensor &r)
- Status **musa::dnn::BatchMatMul::Run** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const
MemoryMaintainer &maintainer=nullptr) const
- Status **musa::dnn::BatchMatMul::RunWithBiasAdd** (Handle &h, Tensor &out, const Tensor &l, const
Tensor &r, const Tensor &bias, const MemoryMaintainer &maintainer=nullptr) const


- Status **musa::dnn::BatchMatMul::Run** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const
int64_t batch, const int64_t m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const
int64_t ldc, const int64_t stride_a, const int64_t stride_b, const int64_t stride_c, const MemoryMaintainer
&maintainer=nullptr) const
- Status **musa::dnn::BatchMatMul::RunWithBiasAdd** (Handle &h, Tensor &out, const Tensor &l, const
Tensor &r, const Tensor &bias, const int64_t batch, const int64_t m, const int64_t n, const int64_t k, const
int64_t lda, const int64_t ldb, const int64_t ldc, const int64_t stride_a, const int64_t stride_b, const int64_t
stride_c, const MemoryMaintainer &maintainer=nullptr) const
- Status **musa::dnn::MatMul::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::MatMul::SetTranspose** (bool left, bool right)
- Status **musa::dnn::MatMul::SetSplitK** (bool split_k)
- Status **musa::dnn::MatMul::SetAlpha** (double alpha)
- Status **musa::dnn::MatMul::SetBeta** (double beta)
- Status **musa::dnn::MatMul::SetGamma** (double gamma)
- Status **musa::dnn::MatMul::GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out, const
Tensor &l, const Tensor &r)
- Status **musa::dnn::MatMul::Run** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const
MemoryMaintainer &maintainer=nullptr) const
- Status **musa::dnn::MatMul::RunWithBiasAdd** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r,
const Tensor &bias, const MemoryMaintainer &maintainer=nullptr) const
- Status **musa::dnn::MatMul::Run** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const int64_t
m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const int64_t ldc, const Memory←-
Maintainer &maintainer=nullptr) const
- Status **musa::dnn::MatMul::RunWithBiasAdd** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r,
const Tensor &bias, const int64_t m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const
int64_t ldc, const MemoryMaintainer &maintainer=nullptr) const
- Status **musa::dnn::Dot::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::Dot::SetAxis** (int axis)
- Status **musa::dnn::Dot::SetAxes** (int axis_l, int axis_r)
- Status **musa::dnn::Dot::SetSplitK** (bool split_k)
- Status **musa::dnn::Dot::GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out, const Tensor
&l, const Tensor &r)
- Status **musa::dnn::Dot::Run** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const Memory←-
Maintainer &maintainer=nullptr) const
- Status **musa::dnn::Dot::RunWithBiasAdd** (Handle &h, Tensor &out, const Tensor &l, const Tensor &r,
const Tensor &bias, const MemoryMaintainer &maintainer=nullptr) const
- Status **musa::dnn::Concat::SetAxis** (int axis)
- Status **musa::dnn::Concat::Run** (Handle &h, Tensor &out, int num_input, const Tensor∗ins) const
- Status **musa::dnn::Permute::Run** (Handle &h, Tensor &out, const Tensor &in) const
- Status **musa::dnn::Permute::SetSrcOffset** (int64_t s_offset)
- Status **musa::dnn::Permute::SetDstOffset** (int64_t d_offset)
- static Status **musa::dnn::Permute::ConfigDimStride** (Tensor &out, Tensor &in, std::initializer_list<int64←-
_t>permute_dims)
- static Status **musa::dnn::Permute::ConfigDimStride** (Tensor &out, Tensor &in, int len, const int64_←-
t∗array_dims)
- Status **musa::dnn::Permute::ConfigDimStrideForSlice** (Tensor &out, Tensor &in, const int64_t∗start)
- Status **musa::dnn::Permute::ConfigDimStrideForSlice** (Tensor &out, Tensor &in, const int64_t∗start,
const int64_t∗stride)
- Status **musa::dnn::Fill::SetValue** (double value)
- Status **musa::dnn::Fill::SetValue** (int64_t value)
- Status **musa::dnn::Fill::Run** (Handle &h, Tensor &out) const
- Status **musa::dnn::Fill::Run** (Handle &h, Tensor &out, Tensor &mask) const
- Status **musa::dnn::Sort::SetDim** (int dim)
- Status **musa::dnn::Sort::SetStable** (bool stable)
- Status **musa::dnn::Sort::SetDescending** (bool descending)


- Status **musa::dnn::Sort::Run** (Handle &h, Tensor &out, Tensor &indices, const Tensor &in, const Memory←-
Maintainer &maintainer) const
- Status **musa::dnn::SortByKey::SetDim** (int dim)
- Status **musa::dnn::SortByKey::SetStable** (bool stable)
- Status **musa::dnn::SortByKey::SetDescending** (bool descending)
- Status **musa::dnn::SortByKey::Run** (Handle &h, Tensor &key_out, Tensor &value_out, const Tensor
&key_in, const Tensor &value_in, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::TopK::SetK** (int k)
- Status **musa::dnn::TopK::SetDim** (int dim)
- Status **musa::dnn::TopK::SetLargest** (bool largest)
- Status **musa::dnn::TopK::SetSorted** (bool sorted)
- Status **musa::dnn::TopK::Run** (Handle &h, Tensor &out, Tensor &indices, const Tensor &in, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::Scan::SetMode** (Mode m)
- Status **musa::dnn::Scan::SetOpType** (ScanOpType op_type)
- Status **musa::dnn::Scan::SetInitVal** (double value)
- Status **musa::dnn::Scan::SetInitVal** (int64_t value)
- Status **musa::dnn::Scan::Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &main-
tainer) const
- Status **musa::dnn::Cumsum::SetDim** (int dim)
- Status **musa::dnn::Cumsum::Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer
&maintainer) const
- Status **musa::dnn::Cum::SetDim** (int dim)
- Status **musa::dnn::Cum::SetMode** (Mode m)
- Status **musa::dnn::Cum::Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &main-
tainer) const
- Status **musa::dnn::GatherX::SetMode** (Mode m)
- Status **musa::dnn::GatherX::SetAxis** (int axis)
- Status **musa::dnn::GatherX::SetBatchDims** (int batch_dims)
- Status **musa::dnn::GatherX::Run** (Handle &h, Tensor &out, const Tensor &index, const Tensor &in) const
- Status **musa::dnn::MaskedSelect::Run** (Handle &h, Tensor &out, const Tensor &input, const Tensor
&mask, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::MaskedScatter::Run** (Handle &h, Tensor &out, const Tensor &mask, const Tensor
&source, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::Unfold::SetAxis** (int axis)
- Status **musa::dnn::Unfold::SetSize** (int size)
- Status **musa::dnn::Unfold::SetStep** (int step)
- Status **musa::dnn::Unfold::Run** (Handle &h, Tensor &out, const Tensor &input) const
- Status **musa::dnn::Unique::SetMode** (Mode m)
- Status **musa::dnn::Unique::Run** (Handle &h, Tensor &out, Tensor &inverse_indices, Tensor &counts, const
Tensor &in, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::Nonzero::Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer
&maintainer) const

4.3.1 Detailed Description

This entity contains a wide array of math operations, including:

  • Arithmetic operations: add, subtract, multiply, divide, etc.
  • Exponential and Logarithmic operations: exp, log, log1p, etc.
  • Trigonometric operations: sin, cos, tan, atan, etc.
  • Neural networks operations: sigmoid, hardsigmoid, leaky_relu, etc.

4.4 NN(Neural Network) Operators

Classes

  • class musa::dnn::Convolution
  • struct musa::dnn::Convolution::FusedActivationDesc
  • class musa::dnn::BatchNorm
  • class musa::dnn::Pooling
  • class musa::dnn::RNN
  • class musa::dnn::DeformableConv
  • class musa::dnn::Dropout
  • class musa::dnn::Pad
  • class musa::dnn::Softmax
  • class musa::dnn::LayerNorm
  • class musa::dnn::RMSNorm
  • class musa::dnn::GroupNorm
  • class musa::dnn::WeightNorm
  • class musa::dnn::Glu
  • class musa::dnn::Scatter
  • class musa::dnn::ScatterND
  • class musa::dnn::NLLLoss
  • class musa::dnn::KLDivLoss
  • class musa::dnn::L2Loss
  • class musa::dnn::LocalResponseNorm
  • class musa::dnn::MultiHeadAttention
  • class musa::dnn::ScaledDotProductAttention
  • class musa::dnn::CTCLoss

Macros

  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,
  • #define MUDNN_ITEM (x) x,

Enumerations

- enum class **Algorithm** { **MUDNN_ITEM** }
- enum class **AlgorithmBwdData** { **MUDNN_ITEM** }
- enum class **AlgorithmBwdFilter** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Format** { **MUDNN_ITEM** }
- enum class **BiasMode** { **MUDNN_ITEM** }
- enum class **Direction** { **MUDNN_ITEM** }
- enum class **Algorithm** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Algorithm** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }

Functions

- Status **musa::dnn::Convolution::FusedActivationDesc::SetMode** (Mode m)
- Status **musa::dnn::Convolution::FusedActivationDesc::SetCoef** (double activAlpha, double activBeta,
double activGamma)
- Status **musa::dnn::Convolution::SetGroups** (int group_num)
- Status **musa::dnn::Convolution::SetNdInfo** (std::initializer_list<int>pad, std::initializer_list<int>stride,
std::initializer_list<int>dilation)
- Status **musa::dnn::Convolution::SetNdInfo** (int length, const int∗pad, const int∗stride, const int∗dilation)
- Status **musa::dnn::Convolution::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::Convolution::Run** (Handle &h, Tensor &out, const Tensor &data, const Tensor &filter,
Algorithm algo, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::Convolution::RunFusion** (Handle &h, Tensor &out, const Tensor &data, const Tensor
&filter, const Tensor &bias, const Tensor &add, const FusedActivationDesc &act, Algorithm algo, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::Convolution::RunBwdData** (Handle &h, Tensor &out, const Tensor &data, const Tensor
&filter, AlgorithmBwdData algo, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::Convolution::RunBwdFilter** (Handle &h, Tensor &out, const Tensor &data, const
Tensor &filter, AlgorithmBwdFilter algo, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::Convolution::GetForwardWorkspaceSize** (Handle &h, size_t &size_in_bytes, const
Tensor &out, const Tensor &data, const Tensor &filter, const Algorithm &algo) const
- Status **musa::dnn::Convolution::GetBackwardDataWorkspaceSize** (Handle &h, size_t &size_in_bytes,
const Tensor &out, const Tensor &data, const Tensor &filter, const AlgorithmBwdData &algo) const
- Status **musa::dnn::Convolution::GetBackwardFilterWorkspaceSize** (Handle &h, size_t &size_in_bytes,
const Tensor &out, const Tensor &data, const Tensor &filter, const AlgorithmBwdFilter &algo) const
- Status **musa::dnn::Convolution::GetRecommendForwardAlgorithm** (Handle &h, Algorithm &algo, const
Tensor &out, const Tensor &data, const Tensor &filter) const
- Status **musa::dnn::Convolution::GetRecommendBackwardDataAlgorithm** (Handle &h, AlgorithmBwd←-
Data &algo, const Tensor &out, const Tensor &data, const Tensor &filter) const
- Status **musa::dnn::Convolution::GetRecommendBackwardFilterAlgorithm** (Handle &h, AlgorithmBwd←-
Filter &algo, const Tensor &out, const Tensor &data, const Tensor &filter) const
- Status **musa::dnn::BatchNorm::SetMode** (Mode m)
- Status **musa::dnn::BatchNorm::SetEpsilon** (double epsilon)


- Status **musa::dnn::BatchNorm::SetTraining** (bool is_training)
- Status **musa::dnn::BatchNorm::RunPure** (Handle &h, Tensor &out, const Tensor &in, const Tensor &final←-
Mean, const Tensor &finalVariance, const Tensor &scale, const Tensor &bias) const
- Status **musa::dnn::BatchNorm::RunComposite** (Handle &h, Tensor &out, const Tensor &in, Tensor &acc←-
Mean, Tensor &accVariance, Tensor &freshMean, Tensor &freshVariance, const Tensor &scale, const Tensor
&bias, double momentum, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::BatchNorm::RunBwd** (Handle &h, Tensor &dx, Tensor &dm, Tensor &dv, Tensor &dg,
Tensor &db, const Tensor &x, const Tensor &dy, const Tensor &m, const Tensor &v, const Tensor &g, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::BatchNorm::GetWorkspaceSizeComposite** (Handle &h, size_t &size_in_byte, Tensor
&out, const Tensor &in, Tensor &accMean, Tensor &accVariance, Tensor &freshMean, Tensor &fresh←-
Variance, const Tensor &scale, const Tensor &bias) const
- Status **musa::dnn::BatchNorm::GetWorkspaceSizeBwd** (Handle &h, size_t &size_in_byte, Tensor &dx,
Tensor &dm, Tensor &dv, Tensor &dg, Tensor &db, const Tensor &x, const Tensor &dy, const Tensor &m,
const Tensor &v, const Tensor &g) const
- Status **musa::dnn::Pooling::SetMode** (Mode m)
- Status **musa::dnn::Pooling::SetNdInfo** (std::initializer_list<int>kernel, std::initializer_list<int>pad,
std::initializer_list<int>stride, std::initializer_list<int>dilation)
- Status **musa::dnn::Pooling::SetNdInfo** (int len, const int∗kernel, const int∗pad, const int∗stride, const int
∗dilation)
- Status **musa::dnn::Pooling::SetDivisor** (int divisor)
- Status **musa::dnn::Pooling::Run** (Handle &h, Tensor &out, const Tensor &in, Tensor &indices) const
- Status **musa::dnn::Pooling::RunBwd** (Handle &h, Tensor &out, const Tensor &in, Tensor &indices) const
- Status **musa::dnn::Pooling::RunMaxPool2dGradGrad** (Handle &h, Tensor &out, const Tensor &in, const
Tensor &indices) const
- Status **musa::dnn::RNN::SetMode** (Mode m)
- Status **musa::dnn::RNN::SetFormat** (Format m)
- Status **musa::dnn::RNN::SetBiasMode** (BiasMode bm)
- Status **musa::dnn::RNN::SetDirection** (Direction d)
- Status **musa::dnn::RNN::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::RNN::RunUnpacked** (Handle &h, Tensor &out, Tensor &hout, Tensor &cout, const
Tensor &in, const Tensor &hin, const Tensor &cin, const Tensor∗weight, Tensor &bwdHint, const Memory←-
Maintainer &maintainer) const
- Status **musa::dnn::RNN::RunUnpackedBwd** (Handle &h, Tensor &in_grad, Tensor &hin_grad, Tensor
&cin_grad, const Tensor∗weight_grad, const Tensor &in, const Tensor &hin, const Tensor &cin, const Tensor
∗weight, const Tensor &out, const Tensor &bwdHint, const Tensor &out_grad, const Tensor &hout_grad, const
Tensor &cout_grad, const MemoryMaintainer &maintainer) const
- static Status **musa::dnn::RNN::RunPackPaddedSequence** (Handle &h, Tensor &packed, Tensor &batsize,
const Tensor &padded, const Tensor &lengths, bool batch_first, bool enforce_sorted)
- static Status **musa::dnn::RNN::RunPadPackedSequence** (Handle &h, Tensor &padded, Tensor &lengths,
const Tensor &packed, const Tensor &batsize, bool batch_first, float padding_value)
- Status **musa::dnn::DeformableConv::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::DeformableConv::SetGroups** (int groups)
- Status **musa::dnn::DeformableConv::SetDeformableGroups** (int deformable_groups)
- Status **musa::dnn::DeformableConv::SetNdInfo** (std::initializer_list<int>pad, std::initializer_list<int>
stride, std::initializer_list<int>dilation)
- Status **musa::dnn::DeformableConv::SetNdInfo** (int length, const int∗pad, const int∗stride, const int
∗dilation)
- Status **musa::dnn::DeformableConv::GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out,
const Tensor &input, const Tensor &weight) const
- Status **musa::dnn::DeformableConv::RunDeformableConv** (Handle &h, Tensor &out, const Tensor &in-
put, const Tensor &weight, const Tensor &offset, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::DeformableConv::RunModulateDeformableConv** (Handle &h, Tensor &out, const
Tensor &input, const Tensor &weight, const Tensor &offset, const Tensor &mask, const MemoryMaintainer
&maintainer) const
- Status **musa::dnn::Dropout::SetP** (double p)


- Status **musa::dnn::Dropout::SetScale** (double scale)
- Status **musa::dnn::Dropout::SetSeed** (uint64_t seed)
- Status **musa::dnn::Dropout::SetOffset** (uint64_t offset)
- Status **musa::dnn::Dropout::RunDropout** (Handle &h, Tensor &out, const Tensor &input, Tensor &mask)
const
- Status **musa::dnn::Dropout::RunDropoutBwd** (Handle &h, Tensor &out, const Tensor &grad, Tensor
&mask) const
- Status **musa::dnn::Pad::SetMode** (Mode mode)
- Status **musa::dnn::Pad::SetValue** (double value)
- Status **musa::dnn::Pad::SetValue** (int64_t value)
- Status **musa::dnn::Pad::SetPaddingInfo** (std::initializer_list<int>pad)
- Status **musa::dnn::Pad::SetPaddingInfo** (int length, const int∗pad)
- Status **musa::dnn::Pad::Run** (Handle &h, Tensor &out, const Tensor &input) const
- Status **musa::dnn::Softmax::SetDim** (int d)
- Status **musa::dnn::Softmax::SetAlgorithm** (Algorithm a)
- Status **musa::dnn::Softmax::SetMode** (Mode m)
- Status **musa::dnn::Softmax::Run** (Handle &h, Tensor &out, const Tensor &input) const
- Status **musa::dnn::Softmax::RunBwd** (Handle &h, Tensor &gradInput, const Tensor &output, const Tensor
&gradOutput) const
- Status **musa::dnn::LayerNorm::SetEpsilon** (double eps)
- Status **musa::dnn::LayerNorm::SetAxis** (std::initializer_list<int>axes)
- Status **musa::dnn::LayerNorm::SetAxis** (size_t length, const int∗axes)
- Status **musa::dnn::LayerNorm::Run** (Handle &h, Tensor &out, Tensor &mean, Tensor &inv_var, const
Tensor &in, const Tensor &gamma, const Tensor &beta) const
- Status **musa::dnn::LayerNorm::RunBwd** (Handle &h, Tensor &dX, Tensor &dGamma, Tensor &dBeta,
const Tensor &dY, const Tensor &in, const Tensor &mean, const Tensor &invVar, const Tensor &gamma,
const MemoryMaintainer &maintainer) const
- Status **musa::dnn::LayerNorm::GetBackwardWorkspaceSize** (Handle &h, size_t size_in_byte, Tensor
&dX, Tensor &dGamma, Tensor &dBeta, const Tensor &dY, const Tensor &in, const Tensor &mean, const
Tensor &invVar, const Tensor &gamma) const
- Status **musa::dnn::RMSNorm::SetEpsilon** (double eps)
- Status **musa::dnn::RMSNorm::SetAxis** (std::initializer_list<int>axes)
- Status **musa::dnn::RMSNorm::SetAxis** (size_t length, const int∗axes)
- Status **musa::dnn::RMSNorm::Run** (Handle &h, Tensor &out, Tensor &mean, const Tensor &in, const
Tensor &gamma) const
- Status **musa::dnn::GroupNorm::SetEpsilon** (double eps)
- Status **musa::dnn::GroupNorm::SetAxis** (const int axis)
- Status **musa::dnn::GroupNorm::SetGroup** (const int g)
- Status **musa::dnn::GroupNorm::Run** (Handle &h, Tensor &out, Tensor &mean, Tensor &invVar, const
Tensor &in, const Tensor &gamma, const Tensor &beta) const
- Status **musa::dnn::WeightNorm::SetEpsilon** (double epsilon)
- Status **musa::dnn::WeightNorm::SetAxis** (std::initializer_list<int>axes)
- Status **musa::dnn::WeightNorm::SetAxis** (size_t length, const int∗axes)
- Status **musa::dnn::WeightNorm::Run** (Handle &h, Tensor &out, const Tensor &weightV, const Tensor
&weightG) const
- Status **musa::dnn::Glu::SetAxis** (int axis)
- Status **musa::dnn::Glu::Run** (Handle &h, Tensor &out, const Tensor &in) const
- Status **musa::dnn::Scatter::SetMode** (Mode mode)
- Status **musa::dnn::Scatter::Run** (Handle &h, Tensor &self, const Tensor &idx, const Tensor &update, int
dim, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::Scatter::Run** (Handle &h, Tensor &out, const Tensor &self, const Tensor &idx, const
Tensor &update, int dim, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::ScatterND::SetMode** (Mode mode)
- Status **musa::dnn::ScatterND::Run** (Handle &h, Tensor &self, const Tensor &idx, const Tensor &update,
const MemoryMaintainer &maintainer) const


- Status **musa::dnn::NLLLoss::SetReductionMode** (Mode mode)
- Status **musa::dnn::NLLLoss::SetIgnoreIndex** (int index)
- Status **musa::dnn::NLLLoss::Run** (Handle &h, Tensor &out, Tensor &totalWeight, const Tensor &in, const
Tensor &target, const Tensor &weight, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::NLLLoss::RunBwd** (Handle &h, Tensor &out, const Tensor &grad, const Tensor &target,
const Tensor &weight, const Tensor &totalWeight, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::KLDivLoss::SetReductionMode** (Mode mode)
- Status **musa::dnn::KLDivLoss::SetLogTarget** (bool log_target)
- Status **musa::dnn::KLDivLoss::Run** (Handle &h, Tensor &out, Tensor &input, const Tensor &target, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::KLDivLoss::RunBwd** (Handle &h, Tensor &out, const Tensor &grad, const Tensor &in-
put, const Tensor &target) const
- Status **musa::dnn::L2Loss::Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer
&maintainer) const
- Status **musa::dnn::LocalResponseNorm::SetMode** (Mode mode)
- Status **musa::dnn::LocalResponseNorm::SetParam** (unsigned n, double alpha, double beta, double k)
- Status **musa::dnn::LocalResponseNorm::Run** (Handle &h, Tensor &y, Tensor &scales, const Tensor &x)
const
- Status **musa::dnn::LocalResponseNorm::RunBwd** (Handle &h, Tensor &dx, const Tensor &y, const Tensor
&dy, const Tensor &x, const Tensor &scales) const
- Status **musa::dnn::MultiHeadAttention::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::MultiHeadAttention::SetEmbedDim** (int embed_dim)
- Status **musa::dnn::MultiHeadAttention::SetKDim** (int k_dim)
- Status **musa::dnn::MultiHeadAttention::SetVDim** (int v_dim)
- Status **musa::dnn::MultiHeadAttention::SetHeadsNum** (int num_heads)
- Status **musa::dnn::MultiHeadAttention::SetBatchFirst** (bool batch_first)
- Status **musa::dnn::MultiHeadAttention::SetDropoutP** (double p)
- Status **musa::dnn::MultiHeadAttention::SetTraining** (bool is_training)
- Status **musa::dnn::MultiHeadAttention::SetMaskMode** (bool use_pad_mask)
- Status **musa::dnn::MultiHeadAttention::SetTransLinearWeigth** (bool trans_weight)
- Status **musa::dnn::MultiHeadAttention::SetKeyFormat** (bool key_format_bhds)
- Status **musa::dnn::MultiHeadAttention::GetFwdBufferSize** (Handle &h, size_t &workspace_size_in_←-
bytes, size_t &bwd_reserve_size_in_bytes, const Tensor &q, const Tensor &k, const Tensor &v) const
- Status **musa::dnn::MultiHeadAttention::GetFwdBufferSize** (Handle &h, size_t &workspace_size_in_←-
bytes, size_t &bwd_reserve_size_in_bytes, const Tensor &in) const
- Status **musa::dnn::MultiHeadAttention::GetBwdBufferSize** (Handle &h, size_t &workspace_size_in_←-
bytes, const Tensor &q, const Tensor &k, const Tensor &v) const
- Status **musa::dnn::MultiHeadAttention::GetBwdBufferSize** (Handle &h, size_t &workspace_size_in_←-
bytes, const Tensor &in) const
- Status **musa::dnn::MultiHeadAttention::Run** (Handle &h, Tensor &out, Tensor &attn_probs, const Tensor
&q, const Tensor &k, const Tensor &v, const Tensor &weight_q, const Tensor &weight_k, const Tensor
&weight_v, const Tensor &weight_o, const Tensor &bias_q, const Tensor &bias_k, const Tensor &bias_←-
v, const Tensor &bias_o, const Tensor &mask, Tensor &bwd_reserve, const MemoryMaintainer &maintainer)
const
- Status **musa::dnn::MultiHeadAttention::Run** (Handle &h, Tensor &out, Tensor &attn_probs, const Tensor
&in, const Tensor &weight, const Tensor &weight_o, const Tensor &bias, const Tensor &bias_o, const Tensor
&mask, Tensor &bwd_reserve, const MemoryMaintainer &maintainer) const
- Status **musa::dnn::MultiHeadAttention::RunBwd** (Handle &h, Tensor &in_grad, Tensor &weight_grad,
Tensor &weight_o_grad, Tensor &bias_grad, Tensor &bias_o_grad, const Tensor &out_grad, const Tensor
&attn_probs, const Tensor &in, const Tensor &weight, const Tensor &weight_o, const Tensor &bwd_reserve,
const MemoryMaintainer &maintainer) const
- Status **musa::dnn::MultiHeadAttention::RunBwd** (Handle &h, Tensor &q_grad, Tensor &k_grad, Tensor
&v_grad, Tensor &weight_q_grad, Tensor &weight_k_grad, Tensor &weight_v_grad, Tensor &weight_o_←-
grad, Tensor &bias_q_grad, Tensor &bias_k_grad, Tensor &bias_v_grad, Tensor &bias_o_grad, const Tensor
&out_grad, const Tensor &attn_probs, const Tensor &q, const Tensor &k, const Tensor &v, const Tensor
&weight_q, const Tensor &weight_k, const Tensor &weight_v, const Tensor &weight_o, const Tensor &bwd←-
_reserve, const MemoryMaintainer &maintainer) const


- Status **musa::dnn::ScaledDotProductAttention::SetComputeMode** (ComputeMode mode)
- Status **musa::dnn::ScaledDotProductAttention::SetEmbedDim** (int embed_dim)
- Status **musa::dnn::ScaledDotProductAttention::SetHeadsNum** (int num_heads)
- Status **musa::dnn::ScaledDotProductAttention::SetDropoutP** (double p)
- Status **musa::dnn::ScaledDotProductAttention::SetTraining** (bool is_training)
- Status **musa::dnn::ScaledDotProductAttention::SetMaskMode** (bool use_pad_mask)
- Status **musa::dnn::ScaledDotProductAttention::SetKeyFormat** (bool key_format_bhds)
- Status **musa::dnn::ScaledDotProductAttention::SetCausal** (bool is_causal)
- Status **musa::dnn::ScaledDotProductAttention::RunFlash** (Handle &h, Tensor &out, Tensor &logsumexp,
const Tensor &q, const Tensor &k, const Tensor &v, const Tensor &mask, Tensor &dropout_mask, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::ScaledDotProductAttention::RunFlashBwd** (Handle &h, Tensor &q_grad, Tensor
&k_grad, Tensor &v_grad, const Tensor &out_grad, const Tensor &q, const Tensor &k, const Tensor &v,
const Tensor &mask, const Tensor &out, const Tensor &logsumexp, const Tensor &dropout_mask, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::ScaledDotProductAttention::RunMath** (Handle &h, Tensor &out, Tensor &attn_probs,
const Tensor &q, const Tensor &k, const Tensor &v, const Tensor &mask, Tensor &dropout_mask, const
MemoryMaintainer &maintainer) const
- Status **musa::dnn::ScaledDotProductAttention::RunMathBwd** (Handle &h, Tensor &q_grad, Tensor &k←-
_grad, const Tensor &v_grad, const Tensor &attn_probs_grad, const Tensor &out_grad, const Tensor &q,
const Tensor &k, const Tensor &v, const Tensor &attn_probs, const Tensor &dropout_mask, const Memory←-
Maintainer &maintainer) const
- Status **musa::dnn::CTCLoss::SetBlank** (int64_t blank)
- Status **musa::dnn::CTCLoss::SetZeroInfinity** (bool zeroInfinity)
- Status **musa::dnn::CTCLoss::SetMaxTargetLength** (int64_t maxTargetLength)
- Status **musa::dnn::CTCLoss::Run** (Handle &h, Tensor &out, Tensor &alpha, const Tensor &input, const
Tensor &targets, const Tensor &inputLengths, const Tensor &targetLength) const
- Status **musa::dnn::CTCLoss::RunBwd** (Handle &h, Tensor &out, const Tensor &grad, const Tensor &log←-
Probs, const Tensor &targets, const Tensor &inputLengths, const Tensor &targetLength, const Tensor &loss,
const Tensor &alpha, const MemoryMaintainer &maintainer) const

Variables

- Mode **musa::dnn::Convolution::FusedActivationDesc::mode** = Mode::IDENTITY
- std::array<double, 3> **musa::dnn::Convolution::FusedActivationDesc::params**

4.4.1 Detailed Description

This entity provides a wide range of classes and functions for building neural networks, such as convolution, pooling, activation functions, and many others.

4.5 Version

Functions

  • MUDNN_EXPORT void musa::dnn::PrintVersionInfo (std::ostream &os)
  • MUDNN_EXPORT size_t musa::dnn::GetVersion ()

4.5.1 Detailed Description

This section describes the version of the muDNN library.

5 Class Documentation

5.1 musa::dnn::BatchMatMul Class Reference

Inheritance diagram for musa::dnn::BatchMatMul:

musa::dnn::BatchMatMul
musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

  • enum class ComputeMode

Public Member Functions

  • Status SetComputeMode (ComputeMode mode)
  • Status SetTranspose (bool left, bool right)
  • Status SetSplitK (bool split_k)
  • Status SetAlpha (double alpha)
  • Status SetBeta (double beta)
  • Status SetGamma (double gamma)
  • Status GetWorkspaceSize (Handle &h, size_t &size_in_bytes, Tensor &out, const Tensor &l, const Tensor &r)
  • Status Run (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const MemoryMaintainer &main- tainer=nullptr) const
  • Status RunWithBiasAdd (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const Tensor &bias, const MemoryMaintainer &maintainer=nullptr) const
  • Status Run (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const int64_t batch, const int64_t m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const int64_t ldc, const int64_t stride_a, const int64_t stride_b, const int64_t stride_c, const MemoryMaintainer &maintainer=nullptr) const
  • Status RunWithBiasAdd (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const Tensor &bias, const int64_t batch, const int64_t m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const int64_t ldc, const int64_t stride_a, const int64_t stride_b, const int64_t stride_c, const MemoryMaintainer &maintainer=nullptr) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.2 musa::dnn::BatchNorm Class Reference

Inheritance diagram for musa::dnn::BatchNorm:

musa::dnn::BatchNorm

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetEpsilon (double epsilon)
  • Status SetTraining (bool is_training)
  • Status RunPure (Handle &h, Tensor &out, const Tensor &in, const Tensor &finalMean, const Tensor &final←- Variance, const Tensor &scale, const Tensor &bias) const
  • Status RunComposite (Handle &h, Tensor &out, const Tensor &in, Tensor &accMean, Tensor &accVariance, Tensor &freshMean, Tensor &freshVariance, const Tensor &scale, const Tensor &bias, double momentum, const MemoryMaintainer &maintainer) const
  • Status RunBwd (Handle &h, Tensor &dx, Tensor &dm, Tensor &dv, Tensor &dg, Tensor &db, const Tensor &x, const Tensor &dy, const Tensor &m, const Tensor &v, const Tensor &g, const MemoryMaintainer &maintainer) const
  • Status GetWorkspaceSizeComposite (Handle &h, size_t &size_in_byte, Tensor &out, const Tensor &in, Tensor &accMean, Tensor &accVariance, Tensor &freshMean, Tensor &freshVariance, const Tensor &scale, const Tensor &bias) const
  • Status GetWorkspaceSizeBwd (Handle &h, size_t &size_in_byte, Tensor &dx, Tensor &dm, Tensor &dv, Tensor &dg, Tensor &db, const Tensor &x, const Tensor &dy, const Tensor &m, const Tensor &v, const Tensor &g) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.3 musa::dnn::Binary Class Reference

Inheritance diagram for musa::dnn::Binary:

musa::dnn::Binary

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetAlpha (double alpha)
  • Status SetAlpha (int64_t alpha)
  • Status SetAlpha (const void∗alpha)
  • Status SetBeta (double beta)
  • Status SetBeta (int64_t beta)
  • Status SetBeta (const void∗beta)
  • Status Run (Handle &h, Tensor &out, const Tensor &l, const Tensor &r) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.4 musa::dnn::Concat Class Reference

Inheritance diagram for musa::dnn::Concat:

musa::dnn::Concat

musa::dnn::ImplBase

Public Member Functions

  • Status SetAxis (int axis)
  • Status Run (Handle &h, Tensor &out, int num_input, const Tensor∗ins) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.5 musa::dnn::Convolution Class Reference

Inheritance diagram for musa::dnn::Convolution:

musa::dnn::Convolution

musa::dnn::ImplBase musa::dnn::MatrixBase

Classes

  • struct FusedActivationDesc

Public Types

- enum class **Algorithm** { **MUDNN_ITEM** }
- enum class **AlgorithmBwdData** { **MUDNN_ITEM** }
- enum class **AlgorithmBwdFilter** { **MUDNN_ITEM** }
- enum class **ComputeMode**

Public Member Functions

- Status **SetGroups** (int group_num)
- Status **SetNdInfo** (std::initializer_list<int>pad, std::initializer_list<int>stride, std::initializer_list<int>
dilation)
- Status **SetNdInfo** (int length, const int∗pad, const int∗stride, const int∗dilation)
- Status **SetComputeMode** (ComputeMode mode)
- Status **Run** (Handle &h, Tensor &out, const Tensor &data, const Tensor &filter, Algorithm algo, const
MemoryMaintainer &maintainer) const
- Status **RunFusion** (Handle &h, Tensor &out, const Tensor &data, const Tensor &filter, const Tensor &bias,
const Tensor &add, const FusedActivationDesc &act, Algorithm algo, const MemoryMaintainer &maintainer)
const
- Status **RunBwdData** (Handle &h, Tensor &out, const Tensor &data, const Tensor &filter, AlgorithmBwdData
algo, const MemoryMaintainer &maintainer) const
- Status **RunBwdFilter** (Handle &h, Tensor &out, const Tensor &data, const Tensor &filter, AlgorithmBwdFilter
algo, const MemoryMaintainer &maintainer) const
- Status **GetForwardWorkspaceSize** (Handle &h, size_t &size_in_bytes, const Tensor &out, const Tensor
&data, const Tensor &filter, const Algorithm &algo) const
- Status **GetBackwardDataWorkspaceSize** (Handle &h, size_t &size_in_bytes, const Tensor &out, const
Tensor &data, const Tensor &filter, const AlgorithmBwdData &algo) const
- Status **GetBackwardFilterWorkspaceSize** (Handle &h, size_t &size_in_bytes, const Tensor &out, const
Tensor &data, const Tensor &filter, const AlgorithmBwdFilter &algo) const
- Status **GetRecommendForwardAlgorithm** (Handle &h, Algorithm &algo, const Tensor &out, const Tensor
&data, const Tensor &filter) const
- Status **GetRecommendBackwardDataAlgorithm** (Handle &h, AlgorithmBwdData &algo, const Tensor &out,
const Tensor &data, const Tensor &filter) const
- Status **GetRecommendBackwardFilterAlgorithm** (Handle &h, AlgorithmBwdFilter &algo, const Tensor
&out, const Tensor &data, const Tensor &filter) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.6 musa::dnn::CTCLoss Class Reference

Inheritance diagram for musa::dnn::CTCLoss:

musa::dnn::CTCLoss

musa::dnn::ImplBase

Public Member Functions

  • Status SetBlank (int64_t blank)
  • Status SetZeroInfinity (bool zeroInfinity)
  • Status SetMaxTargetLength (int64_t maxTargetLength)
  • Status Run (Handle &h, Tensor &out, Tensor &alpha, const Tensor &input, const Tensor &targets, const Tensor &inputLengths, const Tensor &targetLength) const
  • Status RunBwd (Handle &h, Tensor &out, const Tensor &grad, const Tensor &logProbs, const Tensor &tar- gets, const Tensor &inputLengths, const Tensor &targetLength, const Tensor &loss, const Tensor &alpha, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.7 musa::dnn::Cum Class Reference

Inheritance diagram for musa::dnn::Cum:

musa::dnn::Cum

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetDim (int dim)
  • Status SetMode (Mode m)
  • Status Run (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.8 musa::dnn::Cumsum Class Reference

Inheritance diagram for musa::dnn::Cumsum:

musa::dnn::Cumsum

musa::dnn::ImplBase

Public Member Functions

  • Status SetDim (int dim)
  • Status Run (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.9 musa::dnn::DebugInfo Class Reference

Public Types

- enum class **Severity** { **MUDNN_ITEM** }

Public Attributes

  • uint32_t version
  • Severity severity
  • uint32_t time_sec
  • uint32_t time_usec
  • uint32_t time_delta
  • uint64_t tid
  • int32_t device_id
  • const Handle∗ handle

The documentation for this class was generated from the following file:

  • mudnn_base.h

5.10 musa::dnn::DeformableConv Class Reference

Inheritance diagram for musa::dnn::DeformableConv:

musa::dnn::DeformableConv

musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

- enum class **Algorithm** { **MUDNN_ITEM** }
- enum class **ComputeMode**

Public Member Functions

- Status **SetComputeMode** (ComputeMode mode)
- Status **SetGroups** (int groups)
- Status **SetDeformableGroups** (int deformable_groups)
- Status **SetNdInfo** (std::initializer_list<int>pad, std::initializer_list<int>stride, std::initializer_list<int>
dilation)
- Status **SetNdInfo** (int length, const int∗pad, const int∗stride, const int∗dilation)
- Status **GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out, const Tensor &input, const
Tensor &weight) const
- Status **RunDeformableConv** (Handle &h, Tensor &out, const Tensor &input, const Tensor &weight, const
Tensor &offset, const MemoryMaintainer &maintainer) const
- Status **RunModulateDeformableConv** (Handle &h, Tensor &out, const Tensor &input, const Tensor &weight,
const Tensor &offset, const Tensor &mask, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.11 musa::dnn::Dot Class Reference

Inheritance diagram for musa::dnn::Dot:

musa::dnn::Dot

musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

  • enum class ComputeMode

Public Member Functions

  • Status SetComputeMode (ComputeMode mode)
  • Status SetAxis (int axis)
  • Status SetAxes (int axis_l, int axis_r)
  • Status SetSplitK (bool split_k)
  • Status GetWorkspaceSize (Handle &h, size_t &size_in_bytes, Tensor &out, const Tensor &l, const Tensor &r)
  • Status Run (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const MemoryMaintainer &main- tainer=nullptr) const
  • Status RunWithBiasAdd (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const Tensor &bias, const MemoryMaintainer &maintainer=nullptr) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.12 musa::dnn::Dropout Class Reference

Inheritance diagram for musa::dnn::Dropout:

musa::dnn::Dropout

musa::dnn::ImplBase

Public Member Functions

  • Status SetP (double p)
  • Status SetScale (double scale)
  • Status SetSeed (uint64_t seed)
  • Status SetOffset (uint64_t offset)
  • Status RunDropout (Handle &h, Tensor &out, const Tensor &input, Tensor &mask) const
  • Status RunDropoutBwd (Handle &h, Tensor &out, const Tensor &grad, Tensor &mask) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.13 musa::dnn::Fill Class Reference

Inheritance diagram for musa::dnn::Fill:

musa::dnn::Fill

musa::dnn::ImplBase

Public Member Functions

  • Status SetValue (double value)
  • Status SetValue (int64_t value)
  • Status Run (Handle &h, Tensor &out) const
  • Status Run (Handle &h, Tensor &out, Tensor &mask) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.14 musa::dnn::Convolution::FusedActivationDesc Struct Reference

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetCoef (double activAlpha, double activBeta, double activGamma)

Public Attributes

- Mode **mode** = Mode::IDENTITY
- std::array<double, 3> **params**

The documentation for this struct was generated from the following file:

  • mudnn_nn.h

5.15 musa::dnn::GatherX Class Reference

Inheritance diagram for musa::dnn::GatherX:

musa::dnn::GatherX

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetAxis (int axis)
  • Status SetBatchDims (int batch_dims)
  • Status Run (Handle &h, Tensor &out, const Tensor &index, const Tensor &in) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.16 musa::dnn::Glu Class Reference

Inheritance diagram for musa::dnn::Glu:

musa::dnn::Glu

musa::dnn::ImplBase

Public Member Functions

  • Status SetAxis (int axis)
  • Status Run (Handle &h, Tensor &out, const Tensor &in) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.17 musa::dnn::GroupNorm Class Reference

Inheritance diagram for musa::dnn::GroupNorm:

musa::dnn::GroupNorm

musa::dnn::ImplBase

Public Member Functions

  • Status SetEpsilon (double eps)
  • Status SetAxis (const int axis)
  • Status SetGroup (const int g)
  • Status Run (Handle &h, Tensor &out, Tensor &mean, Tensor &invVar, const Tensor &in, const Tensor &gamma, const Tensor &beta) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.18 musa::dnn::Handle Class Reference

Inheritance diagram for musa::dnn::Handle:

musa::dnn::Handle

musa::dnn::ImplBase

Public Member Functions

  • Handle (int device_id)
  • int GetDeviceId () const
  • Status SetStream (musaStream_t stream)
  • musaStream_t GetStream () const
  • Status SetAllowTF32 (bool allow_tf32)
  • bool GetAllowTF32 () const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_base.h

5.19 musa::dnn::ImplBase Class Reference

Inheritance diagram for musa::dnn::ImplBase:

musa::dnn::ImplBase
musa::dnn::BatchMatMul
musa::dnn::BatchNorm
musa::dnn::Binary
musa::dnn::CTCLoss
musa::dnn::Concat
musa::dnn::Convolution
musa::dnn::Cum
musa::dnn::Cumsum
musa::dnn::DeformableConv
musa::dnn::Dot
musa::dnn::Dropout
musa::dnn::Fill
musa::dnn::GatherX
musa::dnn::Glu
musa::dnn::GroupNorm
musa::dnn::Handle
musa::dnn::Interpolate
musa::dnn::KLDivLoss
musa::dnn::L2Loss
musa::dnn::LayerNorm
musa::dnn::LocalResponseNorm
musa::dnn::MaskedScatter
musa::dnn::MaskedSelect
musa::dnn::MatMul
musa::dnn::MultiHeadAttention
musa::dnn::NLLLoss
musa::dnn::Nonzero
musa::dnn::Pad
musa::dnn::Permute
musa::dnn::Pooling
musa::dnn::RMSNorm
musa::dnn::RNN
musa::dnn::Reduce
musa::dnn::ScaledDotProductAttention
musa::dnn::Scan
musa::dnn::Scatter
musa::dnn::ScatterND
musa::dnn::Softmax
musa::dnn::Sort
musa::dnn::SortByKey
musa::dnn::Tensor
musa::dnn::Ternary
musa::dnn::TopK
musa::dnn::Unary
musa::dnn::Unfold
musa::dnn::Unique
musa::dnn::WeightNorm

Public Member Functions

  • void∗ GetImpl ()
  • const void∗ GetImpl () const

Protected Member Functions

  • ImplBase (void∗impl)
  • ImplBase (const ImplBase &)=delete
  • ImplBase (ImplBase &&)=delete
  • ImplBase & operator= (const ImplBase &)=delete
  • ImplBase & operator= (ImplBase &&)=delete

Protected Attributes

  • void∗ impl_

The documentation for this class was generated from the following file:

  • mudnn_base.h

5.20 musa::dnn::Interpolate Class Reference

Inheritance diagram for musa::dnn::Interpolate:

musa::dnn::Interpolate

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

- Status **SetMode** (Mode m)
- Status **SetScaleInfo** (std::initializer_list<float>scale)
- Status **SetScaleInfo** (int length, float∗scale)
- Status **SetAlignCorners** (bool align_corners)
- Status **Run** (Handle &h, Tensor &out, const Tensor &in) const
- Status **RunBackward** (Handle &h, Tensor &out, const Tensor &in) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_image.h

5.21 musa::dnn::KLDivLoss Class Reference

Inheritance diagram for musa::dnn::KLDivLoss:

musa::dnn::KLDivLoss

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetReductionMode (Mode mode)
  • Status SetLogTarget (bool log_target)
  • Status Run (Handle &h, Tensor &out, Tensor &input, const Tensor &target, const MemoryMaintainer &main- tainer) const
  • Status RunBwd (Handle &h, Tensor &out, const Tensor &grad, const Tensor &input, const Tensor &target) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.22 musa::dnn::L2Loss Class Reference

Inheritance diagram for musa::dnn::L2Loss:

musa::dnn::L2Loss

musa::dnn::ImplBase

Public Member Functions

  • Status Run (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.23 musa::dnn::LayerNorm Class Reference

Inheritance diagram for musa::dnn::LayerNorm:

musa::dnn::LayerNorm

musa::dnn::ImplBase

Public Member Functions

- Status **SetEpsilon** (double eps)
- Status **SetAxis** (std::initializer_list<int>axes)
- Status **SetAxis** (size_t length, const int∗axes)
- Status **Run** (Handle &h, Tensor &out, Tensor &mean, Tensor &inv_var, const Tensor &in, const Tensor
&gamma, const Tensor &beta) const
- Status **RunBwd** (Handle &h, Tensor &dX, Tensor &dGamma, Tensor &dBeta, const Tensor &dY, const Tensor
&in, const Tensor &mean, const Tensor &invVar, const Tensor &gamma, const MemoryMaintainer &main-
tainer) const
- Status **GetBackwardWorkspaceSize** (Handle &h, size_t size_in_byte, Tensor &dX, Tensor &dGamma,
Tensor &dBeta, const Tensor &dY, const Tensor &in, const Tensor &mean, const Tensor &invVar, const Tensor
&gamma) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.24 musa::dnn::LocalResponseNorm Class Reference

Inheritance diagram for musa::dnn::LocalResponseNorm:

musa::dnn::LocalResponseNorm

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode mode)
  • Status SetParam (unsigned n, double alpha, double beta, double k)
  • Status Run (Handle &h, Tensor &y, Tensor &scales, const Tensor &x) const
  • Status RunBwd (Handle &h, Tensor &dx, const Tensor &y, const Tensor &dy, const Tensor &x, const Tensor &scales) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.25 musa::dnn::MaskedScatter Class Reference

Inheritance diagram for musa::dnn::MaskedScatter:

musa::dnn::MaskedScatter

musa::dnn::ImplBase

Public Member Functions

  • Status Run (Handle &h, Tensor &out, const Tensor &mask, const Tensor &source, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.26 musa::dnn::MaskedSelect Class Reference

Inheritance diagram for musa::dnn::MaskedSelect:

musa::dnn::MaskedSelect

musa::dnn::ImplBase

Public Member Functions

  • Status Run (Handle &h, Tensor &out, const Tensor &input, const Tensor &mask, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.27 musa::dnn::MatMul Class Reference

Inheritance diagram for musa::dnn::MatMul:

musa::dnn::MatMul

musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

  • enum class ComputeMode

Public Member Functions

  • Status SetComputeMode (ComputeMode mode)
  • Status SetTranspose (bool left, bool right)
  • Status SetSplitK (bool split_k)
  • Status SetAlpha (double alpha)
  • Status SetBeta (double beta)
  • Status SetGamma (double gamma)
  • Status GetWorkspaceSize (Handle &h, size_t &size_in_bytes, Tensor &out, const Tensor &l, const Tensor &r)
  • Status Run (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const MemoryMaintainer &main- tainer=nullptr) const
  • Status RunWithBiasAdd (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const Tensor &bias, const MemoryMaintainer &maintainer=nullptr) const
  • Status Run (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const int64_t m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const int64_t ldc, const MemoryMaintainer &maintainer=nullptr) const
  • Status RunWithBiasAdd (Handle &h, Tensor &out, const Tensor &l, const Tensor &r, const Tensor &bias, const int64_t m, const int64_t n, const int64_t k, const int64_t lda, const int64_t ldb, const int64_t ldc, const MemoryMaintainer &maintainer=nullptr) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.28 musa::dnn::MatrixBase Class Reference

Inheritance diagram for musa::dnn::MatrixBase:

musa::dnn::MatrixBase

musa::dnn::BatchMatMul
musa::dnn::Convolution
musa::dnn::DeformableConv
musa::dnn::Dot
musa::dnn::MatMul

musa::dnn::MultiHeadAttention

musa::dnn::RNN
musa::dnn::ScaledDotProductAttention

Protected Types

- enum class **ComputeMode** { **MUDNN_ITEM** }

The documentation for this class was generated from the following file:

  • mudnn_base.h

5.29 musa::dnn::MultiHeadAttention Class Reference

Inheritance diagram for musa::dnn::MultiHeadAttention:

musa::dnn::MultiHeadAttention
musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

  • enum class ComputeMode

Public Member Functions

  • Status SetComputeMode (ComputeMode mode)
  • Status SetEmbedDim (int embed_dim)
  • Status SetKDim (int k_dim)
  • Status SetVDim (int v_dim)
  • Status SetHeadsNum (int num_heads)
  • Status SetBatchFirst (bool batch_first)
  • Status SetDropoutP (double p)
  • Status SetTraining (bool is_training)
  • Status SetMaskMode (bool use_pad_mask)
  • Status SetTransLinearWeigth (bool trans_weight)
  • Status SetKeyFormat (bool key_format_bhds)
  • Status GetFwdBufferSize (Handle &h, size_t &workspace_size_in_bytes, size_t &bwd_reserve_size_in_←- bytes, const Tensor &q, const Tensor &k, const Tensor &v) const
  • Status GetFwdBufferSize (Handle &h, size_t &workspace_size_in_bytes, size_t &bwd_reserve_size_in_←- bytes, const Tensor &in) const
  • Status GetBwdBufferSize (Handle &h, size_t &workspace_size_in_bytes, const Tensor &q, const Tensor &k, const Tensor &v) const
  • Status GetBwdBufferSize (Handle &h, size_t &workspace_size_in_bytes, const Tensor &in) const
  • Status Run (Handle &h, Tensor &out, Tensor &attn_probs, const Tensor &q, const Tensor &k, const Tensor &v, const Tensor &weight_q, const Tensor &weight_k, const Tensor &weight_v, const Tensor &weight_o, const Tensor &bias_q, const Tensor &bias_k, const Tensor &bias_v, const Tensor &bias_o, const Tensor &mask, Tensor &bwd_reserve, const MemoryMaintainer &maintainer) const
  • Status Run (Handle &h, Tensor &out, Tensor &attn_probs, const Tensor &in, const Tensor &weight, const Tensor &weight_o, const Tensor &bias, const Tensor &bias_o, const Tensor &mask, Tensor &bwd_reserve, const MemoryMaintainer &maintainer) const
  • Status RunBwd (Handle &h, Tensor &in_grad, Tensor &weight_grad, Tensor &weight_o_grad, Tensor &bias_grad, Tensor &bias_o_grad, const Tensor &out_grad, const Tensor &attn_probs, const Tensor &in, const Tensor &weight, const Tensor &weight_o, const Tensor &bwd_reserve, const MemoryMaintainer &main- tainer) const
  • Status RunBwd (Handle &h, Tensor &q_grad, Tensor &k_grad, Tensor &v_grad, Tensor &weight_q_grad, Tensor &weight_k_grad, Tensor &weight_v_grad, Tensor &weight_o_grad, Tensor &bias_q_grad, Tensor &bias_k_grad, Tensor &bias_v_grad, Tensor &bias_o_grad, const Tensor &out_grad, const Tensor &attn_←- probs, const Tensor &q, const Tensor &k, const Tensor &v, const Tensor &weight_q, const Tensor &weight←- _k, const Tensor &weight_v, const Tensor &weight_o, const Tensor &bwd_reserve, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.30 musa::dnn::NLLLoss Class Reference

Inheritance diagram for musa::dnn::NLLLoss:

musa::dnn::NLLLoss

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetReductionMode (Mode mode)
  • Status SetIgnoreIndex (int index)
  • Status Run (Handle &h, Tensor &out, Tensor &totalWeight, const Tensor &in, const Tensor &target, const Tensor &weight, const MemoryMaintainer &maintainer) const
  • Status RunBwd (Handle &h, Tensor &out, const Tensor &grad, const Tensor &target, const Tensor &weight, const Tensor &totalWeight, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.31 musa::dnn::Nonzero Class Reference

Inheritance diagram for musa::dnn::Nonzero:

musa::dnn::Nonzero

musa::dnn::ImplBase

Public Member Functions

  • Status Run (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.32 musa::dnn::Pad Class Reference

Inheritance diagram for musa::dnn::Pad:

musa::dnn::Pad

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

- Status **SetMode** (Mode mode)
- Status **SetValue** (double value)
- Status **SetValue** (int64_t value)
- Status **SetPaddingInfo** (std::initializer_list<int>pad)
- Status **SetPaddingInfo** (int length, const int∗pad)
- Status **Run** (Handle &h, Tensor &out, const Tensor &input) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.33 musa::dnn::Permute Class Reference

Inheritance diagram for musa::dnn::Permute:

musa::dnn::Permute

musa::dnn::ImplBase

Public Member Functions

  • Status Run (Handle &h, Tensor &out, const Tensor &in) const
  • Status SetSrcOffset (int64_t s_offset)
  • Status SetDstOffset (int64_t d_offset)
  • Status ConfigDimStrideForSlice (Tensor &out, Tensor &in, const int64_t∗start)
  • Status ConfigDimStrideForSlice (Tensor &out, Tensor &in, const int64_t∗start, const int64_t∗stride)

Static Public Member Functions

- static Status **ConfigDimStride** (Tensor &out, Tensor &in, std::initializer_list<int64_t>permute_dims)
- static Status **ConfigDimStride** (Tensor &out, Tensor &in, int len, const int64_t∗array_dims)

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.34 musa::dnn::Pooling Class Reference

Inheritance diagram for musa::dnn::Pooling:

musa::dnn::Pooling

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

- Status **SetMode** (Mode m)
- Status **SetNdInfo** (std::initializer_list<int>kernel, std::initializer_list<int>pad, std::initializer_list<int>
stride, std::initializer_list<int>dilation)
- Status **SetNdInfo** (int len, const int∗kernel, const int∗pad, const int∗stride, const int∗dilation)
- Status **SetDivisor** (int divisor)
- Status **Run** (Handle &h, Tensor &out, const Tensor &in, Tensor &indices) const
- Status **RunBwd** (Handle &h, Tensor &out, const Tensor &in, Tensor &indices) const
- Status **RunMaxPool2dGradGrad** (Handle &h, Tensor &out, const Tensor &in, const Tensor &indices) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.35 musa::dnn::Reduce Class Reference

Inheritance diagram for musa::dnn::Reduce:

musa::dnn::Reduce

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

- Status **SetMode** (Mode m)
- Status **SetDim** (std::initializer_list<int>dim)
- Status **SetDim** (int ndim, const int∗dim)
- Status **SetNormOrd** (float ord)
- Status **GetWorkspaceSize** (Handle &h, size_t &size_in_bytes, Tensor &out, const Tensor &in)
- Status **Run** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const
- Status **RunIndices** (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const
- Status **RunWithIndices** (Handle &h, Tensor &out, Tensor &indices, const Tensor &in, const Memory←-
Maintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.36 musa::dnn::RMSNorm Class Reference

Inheritance diagram for musa::dnn::RMSNorm:

musa::dnn::RMSNorm

musa::dnn::ImplBase

Public Member Functions

- Status **SetEpsilon** (double eps)
- Status **SetAxis** (std::initializer_list<int>axes)
- Status **SetAxis** (size_t length, const int∗axes)
- Status **Run** (Handle &h, Tensor &out, Tensor &mean, const Tensor &in, const Tensor &gamma) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.37 musa::dnn::RNN Class Reference

Inheritance diagram for musa::dnn::RNN:

musa::dnn::RNN

musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }
- enum class **Format** { **MUDNN_ITEM** }
- enum class **BiasMode** { **MUDNN_ITEM** }
- enum class **Direction** { **MUDNN_ITEM** }
- enum class **ComputeMode**

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetFormat (Format m)
  • Status SetBiasMode (BiasMode bm)
  • Status SetDirection (Direction d)
  • Status SetComputeMode (ComputeMode mode)
  • Status RunUnpacked (Handle &h, Tensor &out, Tensor &hout, Tensor &cout, const Tensor &in, const Tensor &hin, const Tensor &cin, const Tensor∗weight, Tensor &bwdHint, const MemoryMaintainer &maintainer) const
  • Status RunUnpackedBwd (Handle &h, Tensor &in_grad, Tensor &hin_grad, Tensor &cin_grad, const Tensor ∗weight_grad, const Tensor &in, const Tensor &hin, const Tensor &cin, const Tensor∗weight, const Tensor &out, const Tensor &bwdHint, const Tensor &out_grad, const Tensor &hout_grad, const Tensor &cout_grad, const MemoryMaintainer &maintainer) const

Static Public Member Functions

  • static Status RunPackPaddedSequence (Handle &h, Tensor &packed, Tensor &batsize, const Tensor &padded, const Tensor &lengths, bool batch_first, bool enforce_sorted)
  • static Status RunPadPackedSequence (Handle &h, Tensor &padded, Tensor &lengths, const Tensor &packed, const Tensor &batsize, bool batch_first, float padding_value)

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.38 musa::dnn::ScaledDotProductAttention Class Reference

Inheritance diagram for musa::dnn::ScaledDotProductAttention:

musa::dnn::ScaledDotProductAttention

musa::dnn::ImplBase musa::dnn::MatrixBase

Public Types

  • enum class ComputeMode

Public Member Functions

- Status **SetComputeMode** (ComputeMode mode)
- Status **SetEmbedDim** (int embed_dim)
- Status **SetHeadsNum** (int num_heads)
- Status **SetDropoutP** (double p)
- Status **SetTraining** (bool is_training)
- Status **SetMaskMode** (bool use_pad_mask)
- Status **SetKeyFormat** (bool key_format_bhds)
- Status **SetCausal** (bool is_causal)
- Status **RunFlash** (Handle &h, Tensor &out, Tensor &logsumexp, const Tensor &q, const Tensor &k, const
Tensor &v, const Tensor &mask, Tensor &dropout_mask, const MemoryMaintainer &maintainer) const
- Status **RunFlashBwd** (Handle &h, Tensor &q_grad, Tensor &k_grad, Tensor &v_grad, const Tensor &out←-
_grad, const Tensor &q, const Tensor &k, const Tensor &v, const Tensor &mask, const Tensor &out, const
Tensor &logsumexp, const Tensor &dropout_mask, const MemoryMaintainer &maintainer) const
- Status **RunMath** (Handle &h, Tensor &out, Tensor &attn_probs, const Tensor &q, const Tensor &k, const
Tensor &v, const Tensor &mask, Tensor &dropout_mask, const MemoryMaintainer &maintainer) const
- Status **RunMathBwd** (Handle &h, Tensor &q_grad, Tensor &k_grad, const Tensor &v_grad, const Tensor
&attn_probs_grad, const Tensor &out_grad, const Tensor &q, const Tensor &k, const Tensor &v, const Tensor
&attn_probs, const Tensor &dropout_mask, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.39 musa::dnn::Scan Class Reference

Inheritance diagram for musa::dnn::Scan:

musa::dnn::Scan

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }
- enum class **ScanOpType** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetOpType (ScanOpType op_type)
  • Status SetInitVal (double value)
  • Status SetInitVal (int64_t value)
  • Status Run (Handle &h, Tensor &out, const Tensor &in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.40 musa::dnn::Scatter Class Reference

Inheritance diagram for musa::dnn::Scatter:

musa::dnn::Scatter

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode mode)
  • Status Run (Handle &h, Tensor &self, const Tensor &idx, const Tensor &update, int dim, const Memory←- Maintainer &maintainer) const
  • Status Run (Handle &h, Tensor &out, const Tensor &self, const Tensor &idx, const Tensor &update, int dim, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.41 musa::dnn::ScatterND Class Reference

Inheritance diagram for musa::dnn::ScatterND:

musa::dnn::ScatterND

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode mode)
  • Status Run (Handle &h, Tensor &self, const Tensor &idx, const Tensor &update, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.42 musa::dnn::Softmax Class Reference

Inheritance diagram for musa::dnn::Softmax:

musa::dnn::Softmax

musa::dnn::ImplBase

Public Types

- enum class **Algorithm** { **MUDNN_ITEM** }
- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetDim (int d)
  • Status SetAlgorithm (Algorithm a)
  • Status SetMode (Mode m)
  • Status Run (Handle &h, Tensor &out, const Tensor &input) const
  • Status RunBwd (Handle &h, Tensor &gradInput, const Tensor &output, const Tensor &gradOutput) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h

5.43 musa::dnn::Sort Class Reference

Inheritance diagram for musa::dnn::Sort:

musa::dnn::Sort

musa::dnn::ImplBase

Public Member Functions

  • Status SetDim (int dim)
  • Status SetStable (bool stable)
  • Status SetDescending (bool descending)
  • Status Run (Handle &h, Tensor &out, Tensor &indices, const Tensor &in, const MemoryMaintainer &main- tainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.44 musa::dnn::SortByKey Class Reference

Inheritance diagram for musa::dnn::SortByKey:

musa::dnn::SortByKey

musa::dnn::ImplBase

Public Member Functions

  • Status SetDim (int dim)
  • Status SetStable (bool stable)
  • Status SetDescending (bool descending)
  • Status Run (Handle &h, Tensor &key_out, Tensor &value_out, const Tensor &key_in, const Tensor &value_in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.45 musa::dnn::Tensor Class Reference

Inheritance diagram for musa::dnn::Tensor:

musa::dnn::Tensor

musa::dnn::TensorBase musa::dnn::ImplBase

Public Types

- enum class **Format** { **MUDNN_ITEM** }
- enum class **Type**

Public Member Functions

- **Tensor** (const Tensor &)
- **Tensor** (Tensor &&)
- Tensor & **operator=** (const Tensor &)
- Tensor & **operator=** (Tensor &&)
- Status **SetAddr** (const void∗addr)
- Status **SetType** (Type t)
- Status **SetFormat** (Format f)
- Status **SetNdInfo** (std::initializer_list<int64_t>dim)
- Status **SetNdInfo** (int ndims, const int64_t∗dim)
- Status **SetNdInfo** (int64_t ndims, const int64_t∗dim)
- Status **SetNdInfo** (std::initializer_list<int64_t>dim, std::initializer_list<int64_t>stride)
- Status **SetNdInfo** (int ndims, const int64_t∗dim, const int64_t∗stride)
- Status **SetNdInfo** (int64_t ndims, const int64_t∗dim, const int64_t∗stride)
- Status **GetNdInfo** (std::vector<int64_t>&dim)
- Status **GetNdInfo** (std::vector<int64_t>&dim, std::vector<int64_t>&stride)
- Status **SetQuantizationInfo** (int n, const float∗scales, const unsigned int∗zero_points)
- Status **SetQuantizationInfo** (std::initializer_list<float>scales, std::initializer_list<unsigned int>zero_←-
points)
- Status **SetQuantizationInfo** (const std::vector<float>&scales)
- Status **CopyFrom** (void∗ptr, size_t bytes, int kind, Handle &h, bool sync=false)
- Status **CopyTo** (void∗ptr, size_t bytes, int kind, Handle &h, bool sync=false) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_base.h

5.46 musa::dnn::TensorBase Class Reference

Inheritance diagram for musa::dnn::TensorBase:

musa::dnn::TensorBase

musa::dnn::Tensor

Protected Types

- enum class **Type** { **MUDNN_ITEM** }

The documentation for this class was generated from the following file:

  • mudnn_base.h

5.47 musa::dnn::Ternary Class Reference

Inheritance diagram for musa::dnn::Ternary:

musa::dnn::Ternary

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetAlpha (double alpha)
  • Status SetAlpha (int64_t alpha)
  • Status SetAlpha (const void∗alpha)
  • Status SetBeta (double beta)
  • Status SetBeta (int64_t beta)
  • Status SetBeta (const void∗beta)
  • Status SetGamma (double gamma)
  • Status SetGamma (int64_t gamma)
  • Status SetGamma (const void∗gamma)
  • Status Run (Handle &h, Tensor &out, const Tensor &in0, const Tensor &in1, const Tensor &in2) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.48 musa::dnn::TopK Class Reference

Inheritance diagram for musa::dnn::TopK:

musa::dnn::TopK

musa::dnn::ImplBase

Public Member Functions

- Status **SetK** (int k)
- Status **SetDim** (int dim)
- Status **SetLargest** (bool largest)
- Status **SetSorted** (bool sorted)
- Status **Run** (Handle &h, Tensor &out, Tensor &indices, const Tensor &in, const MemoryMaintainer &main-
tainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.49 musa::dnn::Unary Class Reference

Inheritance diagram for musa::dnn::Unary:

musa::dnn::Unary

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status SetAlpha (double alpha)
  • Status SetAlpha (int64_t alpha)
  • Status SetAlpha (const void∗alpha)
  • Status SetBeta (double beta)
  • Status SetBeta (int64_t beta)
  • Status SetBeta (const void∗beta)
  • Status Run (Handle &h, Tensor &out, const Tensor &in) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.50 musa::dnn::Unfold Class Reference

Inheritance diagram for musa::dnn::Unfold:

musa::dnn::Unfold

musa::dnn::ImplBase

Public Member Functions

  • Status SetAxis (int axis)
  • Status SetSize (int size)
  • Status SetStep (int step)
  • Status Run (Handle &h, Tensor &out, const Tensor &input) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.51 musa::dnn::Unique Class Reference

Inheritance diagram for musa::dnn::Unique:

musa::dnn::Unique

musa::dnn::ImplBase

Public Types

- enum class **Mode** { **MUDNN_ITEM** }

Public Member Functions

  • Status SetMode (Mode m)
  • Status Run (Handle &h, Tensor &out, Tensor &inverse_indices, Tensor &counts, const Tensor &in, const MemoryMaintainer &maintainer) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_math.h

5.52 musa::dnn::WeightNorm Class Reference

Inheritance diagram for musa::dnn::WeightNorm:

musa::dnn::WeightNorm

musa::dnn::ImplBase

Public Member Functions

- Status **SetEpsilon** (double epsilon)
- Status **SetAxis** (std::initializer_list<int>axes)
- Status **SetAxis** (size_t length, const int∗axes)
- Status **Run** (Handle &h, Tensor &out, const Tensor &weightV, const Tensor &weightG) const

Additional Inherited Members

The documentation for this class was generated from the following file:

  • mudnn_nn.h