Skip to main content

muDNN API 参考

1 引言

Moore Threads® MUSA® 深度神经网络(muDNN)库是一个用于深度神经网络中常用原语的GPU加速库。muDNN提供了高度优化的函数,可用于执行各种数学和数据处理任务,例如:

  • 张量操作:逐元素操作、矩阵操作、归约操作等。
  • 神经网络层:卷积、池化、归一化、激活等。
  • 损失函数:KLDivLoss、L2Loss、NLLLoss等。它使用户能够专注于训练神经网络和开发应用程序,而不是加速GPU性能。muDNN库提供了基于上下文的API,允许使用MUSA流轻松进行多线程。

此API参考列出了数据类型定义和函数的详细描述。

1.1 功能需求和保证

  • 每个API都可以安全且并发地访问,即它们是可重入的。用户必须确保在调用API之前,必要的输入内存分配和数据准备已经就绪。
  • 大部分API(除了有特定注释的)可以在相同的配置和输入下复现相同的结果。
  • muDNN支持最大张量维度大小为8。

2 模块索引

2.1 模块

以下是所有模块的列表:

  • 基础操作符
  • 图像操作符
  • 数学操作符
  • NN -(神经网络)操作符
  • 版本

3 类索引

3.1 类列表

  • musa::dnn::BatchMatMul 以下是类、结构体、联合体和接口及其简要描述:
  • 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 模块文档

4.1 基础操作符

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

  • #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,

类型定义

  • 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)>

枚举

  • 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 }

函数

  • 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)

变量

  • 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 详细描述

该实体包含与muDNN上下文创建和销毁、张量实用程序例程、张量核心操作、调试日志等相关的基本功能。

4.2 图像操作符

  • class musa::dnn::Interpolate

  • #define MUDNN_ITEM (x) x,

枚举

  • enum class Mode { MUDNN_ITEM }

函数

  • 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 详细描述

该实体是一系列图像处理和操作的集合,包括调整图像大小、裁剪和插值。

4.3 数学操作符

  • 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

  • #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,

枚举

  • 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 }

函数

  • 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 详细描述

该实体包含广泛的数学操作,包括:

  • 算术运算:加、减、乘、除等。
  • 指数和对数运算:exp、log、log1p等。
  • 三角函数运算:sin、cos、tan、atan等。
  • 神经网络运算:sigmoid、hardsigmoid、leaky_relu等。

4.4 神经网络(NN)操作符

  • 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

  • #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,

枚举

  • 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 }

函数

  • 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

变量

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

4.4.1 详细描述

该实体提供了构建神经网络的广泛类和函数,例如卷积、池化、激活函数等。

4.5 版本

函数

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

4.5.1 详细描述

本节描述了 muDNN 库的版本。

5 类文档

5.1 musa::dnn::BatchMatMul 类参考

musa::dnn::BatchMatMul 的继承图:

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

公共类型

  • enum class ComputeMode

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.2 musa::dnn::BatchNorm 类参考

musa::dnn::BatchNorm 的继承图:

musa::dnn::BatchNorm

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.3 musa::dnn::Binary 类参考

musa::dnn::Binary 的继承图:

musa::dnn::Binary

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.4 musa::dnn::Concat 类参考

musa::dnn::Concat 的继承图:

musa::dnn::Concat

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.5 musa::dnn::Convolution 类参考

musa::dnn::Convolution 的继承图:

musa::dnn::Convolution

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

  • 结构 FusedActivationDesc

公共类型

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

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.6 musa::dnn::CTCLoss 类参考

musa::dnn::CTCLoss 的继承图:

musa::dnn::CTCLoss

musa::dnn::ImplBase

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.7 musa::dnn::Cum 类参考

musa::dnn::Cum 的继承图:

musa::dnn::Cum

musa::dnn::ImplBase

公共类型

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

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.8 musa::dnn::Cumsum 类参考

musa::dnn::Cumsum 的继承图:

musa::dnn::Cumsum

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.9 musa::dnn::DebugInfo 类参考

公共类型

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

公共属性

  • 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

此类的文档是从以下文件生成的:

  • mudnn_base.h

5.10 musa::dnn::DeformableConv 类参考

musa::dnn::DeformableConv 的继承图:

musa::dnn::DeformableConv

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

公共类型

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

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.11 musa::dnn::Dot 类参考

musa::dnn::Dot 的继承图:

musa::dnn::Dot

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

公共类型

  • enum class ComputeMode

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.12 musa::dnn::Dropout 类参考

musa::dnn::Dropout 的继承图:

musa::dnn::Dropout

musa::dnn::ImplBase

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.13 musa::dnn::Fill 类参考

musa::dnn::Fill 的继承图:

musa::dnn::Fill

musa::dnn::ImplBase

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.14 musa::dnn::Convolution::FusedActivationDesc 结构参考

公共类型

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

公共成员函数

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

公共属性

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

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.15 musa::dnn::GatherX 类参考

musa::dnn::GatherX 的继承图:

musa::dnn::GatherX

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.16 musa::dnn::Glu 类参考

musa::dnn::Glu 的继承图:

musa::dnn::Glu

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.17 musa::dnn::GroupNorm 类参考

musa::dnn::GroupNorm 的继承图:

musa::dnn::GroupNorm

musa::dnn::ImplBase

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.18 musa::dnn::Handle 类参考

musa::dnn::Handle 的继承图:

musa::dnn::Handle

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_base.h

5.19 musa::dnn::ImplBase 类参考

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

公共成员函数

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

保护成员函数

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

保护属性

  • void∗ impl_

此类的文档是从以下文件生成的:

  • mudnn_base.h

5.20 musa::dnn::Interpolate 类参考

musa::dnn::Interpolate 的继承图:

musa::dnn::Interpolate

musa::dnn::ImplBase

公共类型

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

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_image.h

5.21 musa::dnn::KLDivLoss 类参考

musa::dnn::KLDivLoss 的继承图:

musa::dnn::KLDivLoss

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.22 musa::dnn::L2Loss 类参考

musa::dnn::L2Loss 的继承图:

musa::dnn::L2Loss

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.23 musa::dnn::LayerNorm 类参考

musa::dnn::LayerNorm 的继承图:

musa::dnn::LayerNorm

musa::dnn::ImplBase

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.24 musa::dnn::LocalResponseNorm 类参考

musa::dnn::LocalResponseNorm 的继承图:

musa::dnn::LocalResponseNorm

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.25 musa::dnn::MaskedScatter 类参考

musa::dnn::MaskedScatter 的继承图:

musa::dnn::MaskedScatter

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.26 musa::dnn::MaskedSelect 类参考

musa::dnn::MaskedSelect 的继承图:

musa::dnn::MaskedSelect

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.27 musa::dnn::MatMul 类参考

musa::dnn::MatMul 的继承图:

musa::dnn::MatMul

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

公共类型

  • enum class ComputeMode

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.28 musa::dnn::MatrixBase 类参考

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

保护类型

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

此类的文档是从以下文件生成的:

  • mudnn_base.h

5.29 musa::dnn::MultiHeadAttention 类参考

musa::dnn::MultiHeadAttention 的继承图:

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

公共类型

  • enum class ComputeMode

公共成员函数

  • 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 &maintainer) 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.30 musa::dnn::NLLLoss 类参考

musa::dnn::NLLLoss 的继承图:

musa::dnn::NLLLoss

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.31 musa::dnn::Nonzero 类参考

musa::dnn::Nonzero 的继承图:

musa::dnn::Nonzero

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.32 musa::dnn::Pad 类参考

musa::dnn::Pad 的继承图:

musa::dnn::Pad

musa::dnn::ImplBase

公共类型

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

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.33 musa::dnn::Permute 类参考

musa::dnn::Permute 的继承图:

musa::dnn::Permute

musa::dnn::ImplBase

公共成员函数

  • 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)

静态公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.34 musa::dnn::Pooling 类参考

musa::dnn::Pooling 的继承图:

musa::dnn::Pooling

musa::dnn::ImplBase

公共类型

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

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.35 musa::dnn::Reduce 类参考

musa::dnn::Reduce 的继承图:

musa::dnn::Reduce

musa::dnn::ImplBase

公共类型

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

公共成员函数

- 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 MemoryMaintainer &maintainer) const

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.36 musa::dnn::RMSNorm 类参考

musa::dnn::RMSNorm 的继承图:

musa::dnn::RMSNorm

musa::dnn::ImplBase

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.37 musa::dnn::RNN 类参考

musa::dnn::RNN 的继承图:

musa::dnn::RNN

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

公共类型

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

公共成员函数

  • 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

静态公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.38 musa::dnn::ScaledDotProductAttention 类参考

musa::dnn::ScaledDotProductAttention 的继承图:

musa::dnn::ScaledDotProductAttention

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

公共类型

  • enum class ComputeMode

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.39 musa::dnn::Scan 类参考

musa::dnn::Scan 的继承图:

musa::dnn::Scan

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.40 musa::dnn::Scatter 类参考

musa::dnn::Scatter 的继承图:

musa::dnn::Scatter

musa::dnn::ImplBase

公共类型

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

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.41 musa::dnn::ScatterND 类参考

musa::dnn::ScatterND 的继承图:

musa::dnn::ScatterND

musa::dnn::ImplBase

公共类型

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

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.42 musa::dnn::Softmax 类参考

musa::dnn::Softmax 的继承图:

musa::dnn::Softmax

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h

5.43 musa::dnn::Sort 类参考

musa::dnn::Sort 的继承图:

musa::dnn::Sort

musa::dnn::ImplBase

公共成员函数

  • 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 &maintainer) const

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.44 musa::dnn::SortByKey 类参考

musa::dnn::SortByKey 的继承图:

musa::dnn::SortByKey

musa::dnn::ImplBase

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.45 musa::dnn::Tensor 类参考

musa::dnn::Tensor 的继承图:

musa::dnn::Tensor

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

公共类型

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

公共成员函数

- **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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_base.h

5.46 musa::dnn::TensorBase 类参考

musa::dnn::TensorBase 的继承图:

musa::dnn::TensorBase

musa::dnn::Tensor

保护类型

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

此类的文档是从以下文件生成的:

  • mudnn_base.h

5.47 musa::dnn::Ternary 类参考

musa::dnn::Ternary 的继承图:

musa::dnn::Ternary

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.48 musa::dnn::TopK 类参考

musa::dnn::TopK 的继承图:

musa::dnn::TopK

musa::dnn::ImplBase

公共成员函数

- 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 &maintainer) const

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.49 musa::dnn::Unary 类参考

musa::dnn::Unary 的继承图:

musa::dnn::Unary

musa::dnn::ImplBase

公共类型

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

公共成员函数

  • 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.50 musa::dnn::Unfold 类参考

musa::dnn::Unfold 的继承图:

musa::dnn::Unfold

musa::dnn::ImplBase

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.51 musa::dnn::Unique 类参考

musa::dnn::Unique 的继承图:

musa::dnn::Unique

musa::dnn::ImplBase

公共类型

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

公共成员函数

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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_math.h

5.52 musa::dnn::WeightNorm 类参考

musa::dnn::WeightNorm 的继承图:

musa::dnn::WeightNorm

musa::dnn::ImplBase

公共成员函数

- 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

其他继承成员

此类的文档是从以下文件生成的:

  • mudnn_nn.h