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