mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-20 19:12:24 +00:00
onboard MoE (#18279)
### Description <!-- Describe your changes. --> 1. Introduce MoE CUDA op to ORT based on FT implementation. 2. Upgrade cutlass to 3.1.0 to avoid some build failures on Windows. Remove patch file for cutlass 3.0.0. 3. Sharded MoE implementation will come with another PR limitation: __CUDA_ARCH__ >= 700 ### Motivation and Context <!-- - Why is this change required? What problem does it solve? - If it fixes an open issue, please link to the issue here. -->
This commit is contained in:
parent
27d068569a
commit
f9af94009b
32 changed files with 4219 additions and 97 deletions
|
|
@ -286,7 +286,7 @@
|
|||
"component": {
|
||||
"type": "git",
|
||||
"git": {
|
||||
"commitHash": "c4f6b8c6bc94ff69048492fb34df0dfaf1983933",
|
||||
"commitHash": "6f47420213f757831fae65c686aa471749fa8d60",
|
||||
"repositoryUrl": "https://github.com/NVIDIA/cutlass.git"
|
||||
},
|
||||
"comments": "cutlass"
|
||||
|
|
|
|||
|
|
@ -51,7 +51,7 @@ pytorch_cpuinfo;https://github.com/pytorch/cpuinfo/archive/959002f82d7962a473d8b
|
|||
re2;https://github.com/google/re2/archive/refs/tags/2022-06-01.zip;aa77313b76e91b531ee7f3e45f004c6a502a5374
|
||||
safeint;https://github.com/dcleblanc/SafeInt/archive/refs/tags/3.0.28.zip;23f252040ff6cb9f1fd18575b32fa8fb5928daac
|
||||
tensorboard;https://github.com/tensorflow/tensorboard/archive/373eb09e4c5d2b3cc2493f0949dc4be6b6a45e81.zip;67b833913605a4f3f499894ab11528a702c2b381
|
||||
cutlass;https://github.com/NVIDIA/cutlass/archive/refs/tags/v3.0.0.zip;0f95b3c1fc1bd1175c4a90b2c9e39074d1bccefd
|
||||
cutlass;https://github.com/NVIDIA/cutlass/archive/refs/tags/v3.1.0.zip;757f90a795034a89d4f48a79d1f009f7a04c8dee
|
||||
utf8_range;https://github.com/protocolbuffers/utf8_range/archive/72c943dea2b9240cd09efde15191e144bc7c7d38.zip;9925739c9debc0efa2adcb194d371a35b6a03156
|
||||
extensions;https://github.com/microsoft/onnxruntime-extensions/archive/94142d8391c9791ec71c38336436319a2d4ac7a0.zip;4365ac5140338b4cb75a39944a4be276e3829b3c
|
||||
composable_kernel;https://github.com/ROCmSoftwarePlatform/composable_kernel/archive/a4f72a314a85732ed67d5aa8d1088d207a7e0e61.zip;f57357ab6d300e207a632d034ebc8aa036a090d9
|
||||
|
|
|
|||
1
cmake/external/cutlass.cmake
vendored
1
cmake/external/cutlass.cmake
vendored
|
|
@ -4,7 +4,6 @@ if (onnxruntime_USE_FLASH_ATTENTION OR onnxruntime_USE_MEMORY_EFFICIENT_ATTENTIO
|
|||
cutlass
|
||||
URL ${DEP_URL_cutlass}
|
||||
URL_HASH SHA1=${DEP_SHA1_cutlass}
|
||||
PATCH_COMMAND ${Patch_EXECUTABLE} --binary --ignore-whitespace -p1 < ${PROJECT_SOURCE_DIR}/patches/cutlass/cutlass.patch
|
||||
)
|
||||
|
||||
FetchContent_GetProperties(cutlass)
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@ set(contrib_ops_excluded_files
|
|||
"math/gemm_float8.cc"
|
||||
"math/gemm_float8.cu"
|
||||
"math/gemm_float8.h"
|
||||
"moe/*"
|
||||
"quantization/attention_quantization.cc"
|
||||
"quantization/attention_quantization.h"
|
||||
"quantization/attention_quantization_impl.cu"
|
||||
|
|
|
|||
|
|
@ -1,92 +0,0 @@
|
|||
diff --git a/include/cute/numeric/complex.hpp b/include/cute/numeric/complex.hpp
|
||||
index 3790ebd3..cf727d09 100644
|
||||
--- a/include/cute/numeric/complex.hpp
|
||||
+++ b/include/cute/numeric/complex.hpp
|
||||
@@ -41,10 +41,14 @@
|
||||
// With CUDA 11.4, builds show spurious "-Wconversion" warnings
|
||||
// on line 656 of thrust/detail/type_traits.h.
|
||||
// These pragmas suppress the warnings.
|
||||
+#ifdef __GNUC__
|
||||
#pragma GCC diagnostic push
|
||||
#pragma GCC diagnostic ignored "-Wconversion"
|
||||
+#endif
|
||||
#include <thrust/complex.h>
|
||||
+#ifdef __GNUC__
|
||||
#pragma GCC diagnostic pop
|
||||
+#endif
|
||||
|
||||
#include <cute/config.hpp>
|
||||
|
||||
diff --git a/include/cutlass/functional.h b/include/cutlass/functional.h
|
||||
index 59aec46a..8f2a913a 100644
|
||||
--- a/include/cutlass/functional.h
|
||||
+++ b/include/cutlass/functional.h
|
||||
@@ -89,7 +89,7 @@ struct multiplies {
|
||||
}
|
||||
};
|
||||
|
||||
-#if defined(__CUDA_ARCH__)
|
||||
+#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 530)
|
||||
/// Partial specializations needed when __CUDA_NO_HALF2_OPERATORS__ is set
|
||||
template<>
|
||||
struct plus<__half2> {
|
||||
@@ -143,12 +143,12 @@ struct multiplies<__half> {
|
||||
|
||||
|
||||
// Maximum with nan propogation
|
||||
-// To propgate the NANs, the "max" of a two element that contains NaNs should also return a NaN
|
||||
+// To propgate the NANs, the "max" of a two element that contains NaNs should also return a NaN
|
||||
template <typename T>
|
||||
struct maximum_with_nan_propogation {
|
||||
CUTLASS_HOST_DEVICE
|
||||
T operator()(T const &lhs, T const &rhs) const {
|
||||
- return lhs > rhs or std::isnan(lhs) ? lhs : rhs;
|
||||
+ return lhs > rhs or isnan(lhs) ? lhs : rhs;
|
||||
}
|
||||
};
|
||||
|
||||
@@ -160,7 +160,7 @@ struct maximum_with_nan_propogation<float> {
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800)
|
||||
asm volatile("max.NaN.f32 %0, %1, %2;\n" : "=f"(res) : "f"(lhs), "f"(rhs));
|
||||
#else
|
||||
- res = lhs > rhs or std::isnan(lhs) ? lhs : rhs;
|
||||
+ res = lhs > rhs or isnan(lhs) ? lhs : rhs;
|
||||
#endif
|
||||
return res;
|
||||
}
|
||||
@@ -233,7 +233,7 @@ struct negate {
|
||||
}
|
||||
};
|
||||
|
||||
-/// Greater equal
|
||||
+/// Greater equal
|
||||
template <typename T>
|
||||
struct greater_equal {
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -242,7 +242,7 @@ struct greater_equal {
|
||||
}
|
||||
};
|
||||
|
||||
-/// Greater
|
||||
+/// Greater
|
||||
template <typename T>
|
||||
struct greater {
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -251,7 +251,7 @@ struct greater {
|
||||
}
|
||||
};
|
||||
|
||||
-/// Less equal
|
||||
+/// Less equal
|
||||
template <typename T>
|
||||
struct less_equal {
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -260,7 +260,7 @@ struct less_equal {
|
||||
}
|
||||
};
|
||||
|
||||
-/// Less
|
||||
+/// Less
|
||||
template <typename T>
|
||||
struct less {
|
||||
CUTLASS_HOST_DEVICE
|
||||
|
|
@ -54,6 +54,7 @@ Do not modify directly.*
|
|||
* <a href="#com.microsoft.MatMulIntegerToFloat">com.microsoft.MatMulIntegerToFloat</a>
|
||||
* <a href="#com.microsoft.MatMulNBits">com.microsoft.MatMulNBits</a>
|
||||
* <a href="#com.microsoft.MaxpoolWithMask">com.microsoft.MaxpoolWithMask</a>
|
||||
* <a href="#com.microsoft.MoE">com.microsoft.MoE</a>
|
||||
* <a href="#com.microsoft.MulInteger">com.microsoft.MulInteger</a>
|
||||
* <a href="#com.microsoft.MultiHeadAttention">com.microsoft.MultiHeadAttention</a>
|
||||
* <a href="#com.microsoft.MurmurHash3">com.microsoft.MurmurHash3</a>
|
||||
|
|
@ -2904,6 +2905,58 @@ This version of the operator has been available since version 1 of the 'com.micr
|
|||
</dl>
|
||||
|
||||
|
||||
### <a name="com.microsoft.MoE"></a><a name="com.microsoft.moe">**com.microsoft.MoE**</a>
|
||||
|
||||
Mixture of experts. Examples: Switch transformer(https://arxiv.org/pdf/2101.03961.pdf) use top 1,
|
||||
GLaM(https://arxiv.org/abs/2112.06905) activates top 2 FFN, and Vision MOE(https://arxiv.org/pdf/2106.05974.pdf)
|
||||
usually uses top 32 experts.
|
||||
|
||||
|
||||
#### Version
|
||||
|
||||
This version of the operator has been available since version 1 of the 'com.microsoft' operator set.
|
||||
|
||||
#### Attributes
|
||||
|
||||
<dl>
|
||||
<dt><tt>activation_type</tt> : string</dt>
|
||||
<dd>Activation function to use. Choose from relu, gelu, silu and identity. Default is relu</dd>
|
||||
<dt><tt>k</tt> : int</dt>
|
||||
<dd>Number of top experts to select from expert pool</dd>
|
||||
</dl>
|
||||
|
||||
#### Inputs (4 - 6)
|
||||
|
||||
<dl>
|
||||
<dt><tt>input</tt> : T</dt>
|
||||
<dd>2D input tensor with shape (num_rows, hidden_size) or 3D input tensor with shape (batch_size, sequence_length, hidden_size)</dd>
|
||||
<dt><tt>router_probs</tt> : T</dt>
|
||||
<dd>2D input tensor with shape (num_rows, num_experts)</dd>
|
||||
<dt><tt>fc1_experts_weights</tt> : T</dt>
|
||||
<dd>3D input tensor with shape (num_experts, hidden_size, inter_size)</dd>
|
||||
<dt><tt>fc2_experts_weights</tt> : T</dt>
|
||||
<dd>3D input tensor with shape (num_experts, inter_size, hidden_size)</dd>
|
||||
<dt><tt>fc1_experts_bias</tt> (optional) : T</dt>
|
||||
<dd>2D optional input tensor with shape (num_experts, inter_size)</dd>
|
||||
<dt><tt>fc2_experts_bias</tt> (optional) : T</dt>
|
||||
<dd>2D optional input tensor with shape (num_experts, hidden_size)</dd>
|
||||
</dl>
|
||||
|
||||
#### Outputs
|
||||
|
||||
<dl>
|
||||
<dt><tt>output</tt> : T</dt>
|
||||
<dd>2D input tensor with shape (num_rows, hidden_size) or 3D input tensor with shape (batch_size, sequence_length, hidden_size)</dd>
|
||||
</dl>
|
||||
|
||||
#### Type Constraints
|
||||
|
||||
<dl>
|
||||
<dt><tt>T</tt> : tensor(float), tensor(float16)</dt>
|
||||
<dd>Constrain input and output types to float or float16 tensors.</dd>
|
||||
</dl>
|
||||
|
||||
|
||||
### <a name="com.microsoft.MulInteger"></a><a name="com.microsoft.mulinteger">**com.microsoft.MulInteger**</a>
|
||||
|
||||
Performs element-wise binary quantized multiplication (with Numpy-style broadcasting support).
|
||||
|
|
|
|||
|
|
@ -842,6 +842,7 @@ Do not modify directly.*
|
|||
|LongformerAttention|*in* input:**T**<br> *in* weight:**T**<br> *in* bias:**T**<br> *in* mask:**T**<br> *in* global_weight:**T**<br> *in* global_bias:**T**<br> *in* global:**G**<br> *out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
|
||||
|MatMulBnb4|*in* A:**T1**<br> *in* B:**T2**<br> *in* absmax:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)|
|
||||
|MatMulNBits|*in* A:**T1**<br> *in* B:**T2**<br> *in* scales:**T1**<br> *in* zero_points:**T2**<br> *out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)|
|
||||
|MoE|*in* input:**T**<br> *in* router_probs:**T**<br> *in* fc1_experts_weights:**T**<br> *in* fc2_experts_weights:**T**<br> *in* fc1_experts_bias:**T**<br> *in* fc2_experts_bias:**T**<br> *out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
|
||||
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* relative_position_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**|1+|**T** = tensor(float), tensor(float16)|
|
||||
|NGramRepeatBlock|*in* input_ids:**Tid**<br> *in* scores:**T**<br> *out* scores_out:**T**|1+|**T** = tensor(float)<br/> **Tid** = tensor(int64)|
|
||||
|NhwcConv|*in* X:**T**<br> *in* W:**T**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|
||||
|
|
|
|||
|
|
@ -70,6 +70,8 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float, Crop);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double, Crop);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, MLFloat16, Crop);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, MoE);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, MoE);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, MultiHeadAttention);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, MultiHeadAttention);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, GroupQueryAttention);
|
||||
|
|
@ -260,6 +262,8 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, float, Crop)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, double, Crop)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, MLFloat16, Crop)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, MoE)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, MoE)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, MultiHeadAttention)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, MultiHeadAttention)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, GroupQueryAttention)>,
|
||||
|
|
|
|||
51
onnxruntime/contrib_ops/cuda/moe/ft_moe/compute_occupancy.h
Normal file
51
onnxruntime/contrib_ops/cuda/moe/ft_moe/compute_occupancy.h
Normal file
|
|
@ -0,0 +1,51 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include <cuda_runtime_api.h>
|
||||
|
||||
#include "core/providers/cuda/shared_inc/cuda_call.h"
|
||||
#include "cutlass/device_kernel.h"
|
||||
|
||||
using namespace onnxruntime;
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
template <typename GemmKernel>
|
||||
inline int compute_occupancy_for_kernel() {
|
||||
int smem_size = int(sizeof(typename GemmKernel::SharedStorage));
|
||||
|
||||
if (smem_size > (48 << 10)) {
|
||||
cudaError_t status =
|
||||
cudaFuncSetAttribute(cutlass::Kernel<GemmKernel>, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
|
||||
if (status == cudaError::cudaErrorInvalidValue) {
|
||||
// Clear the error bit since we can ignore this.
|
||||
// This should mean that smem_size > cudaDevAttrMaxSharedMemoryPerBlockOptin. In that case, we return an
|
||||
// occupancy of 0. This will cause the heuristic to ignore this configuration.
|
||||
status = cudaGetLastError();
|
||||
return 0;
|
||||
}
|
||||
CUDA_CALL_THROW(status);
|
||||
}
|
||||
|
||||
int max_active_blocks = -1;
|
||||
CUDA_CALL_THROW(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&max_active_blocks, cutlass::Kernel<GemmKernel>,
|
||||
GemmKernel::kThreadCount, smem_size));
|
||||
|
||||
return max_active_blocks;
|
||||
}
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
187
onnxruntime/contrib_ops/cuda/moe/ft_moe/cutlass_heuristic.cc
Normal file
187
onnxruntime/contrib_ops/cuda/moe/ft_moe/cutlass_heuristic.cc
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#include "cutlass_heuristic.h"
|
||||
|
||||
#include <cuda_runtime_api.h>
|
||||
#include <vector>
|
||||
#include <stdexcept>
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
struct TileShape {
|
||||
int m;
|
||||
int n;
|
||||
};
|
||||
|
||||
TileShape get_cta_shape_for_config(CutlassTileConfig tile_config) {
|
||||
switch (tile_config) {
|
||||
case CutlassTileConfig::CtaShape32x128x64_WarpShape32x32x64:
|
||||
return TileShape{32, 128};
|
||||
case CutlassTileConfig::CtaShape64x128x64_WarpShape32x64x64:
|
||||
case CutlassTileConfig::CtaShape64x128x64_WarpShape64x32x64:
|
||||
return TileShape{64, 128};
|
||||
case CutlassTileConfig::CtaShape128x128x8_WarpShape64x64x8:
|
||||
case CutlassTileConfig::CtaShape128x128x64_WarpShape64x32x64:
|
||||
case CutlassTileConfig::CtaShape128x128x64_WarpShape128x32x64:
|
||||
return TileShape{128, 128};
|
||||
default:
|
||||
ORT_THROW("[FT Error][get_grid_shape_for_config] Invalid config");
|
||||
}
|
||||
}
|
||||
|
||||
bool is_valid_split_k_factor(const int64_t m, const int64_t n, const int64_t k, const TileShape tile_shape,
|
||||
const int split_k_factor, const size_t workspace_bytes, const bool is_weight_only) {
|
||||
// All tile sizes have a k_tile of 64.
|
||||
static constexpr int k_tile = 64;
|
||||
|
||||
// For weight-only quant, we need k and k_elements_per_split to be a multiple of cta_k
|
||||
if (is_weight_only) {
|
||||
if ((k % k_tile) != 0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if ((k % split_k_factor) != 0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const int k_elements_per_split = static_cast<int>(k / split_k_factor);
|
||||
if ((k_elements_per_split % k_tile) != 0) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
// Check that the workspace has sufficient space for this split-k factor
|
||||
const int ctas_in_m_dim = static_cast<int>((m + tile_shape.m - 1) / tile_shape.m);
|
||||
const int ctas_in_n_dim = static_cast<int>((n + tile_shape.n - 1) / tile_shape.n);
|
||||
const int required_ws_bytes = split_k_factor == 1 ? 0 : sizeof(int) * ctas_in_m_dim * ctas_in_n_dim;
|
||||
|
||||
if (required_ws_bytes > workspace_bytes) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
std::vector<CutlassTileConfig> get_candidate_tiles(const bool is_weight_only, const bool simt_configs_only) {
|
||||
std::vector<CutlassTileConfig> simt_configs{CutlassTileConfig::CtaShape128x128x8_WarpShape64x64x8};
|
||||
|
||||
std::vector<CutlassTileConfig> square_configs{CutlassTileConfig::CtaShape32x128x64_WarpShape32x32x64,
|
||||
CutlassTileConfig::CtaShape64x128x64_WarpShape32x64x64,
|
||||
CutlassTileConfig::CtaShape128x128x64_WarpShape64x32x64};
|
||||
|
||||
std::vector<CutlassTileConfig> quant_B_configs{CutlassTileConfig::CtaShape32x128x64_WarpShape32x32x64,
|
||||
CutlassTileConfig::CtaShape64x128x64_WarpShape64x32x64,
|
||||
CutlassTileConfig::CtaShape128x128x64_WarpShape128x32x64};
|
||||
|
||||
const std::vector<CutlassTileConfig> allowed_configs = is_weight_only ? quant_B_configs : square_configs;
|
||||
return simt_configs_only ? simt_configs : allowed_configs;
|
||||
}
|
||||
|
||||
std::vector<CutlassGemmConfig> get_candidate_configs(int sm, const bool is_weight_only, const bool simt_configs_only) {
|
||||
std::vector<CutlassTileConfig> tiles = get_candidate_tiles(is_weight_only, simt_configs_only);
|
||||
|
||||
std::vector<CutlassGemmConfig> candidate_configs;
|
||||
const int min_stages = 2;
|
||||
const int max_stages = sm >= 80 ? 4 : 2;
|
||||
|
||||
for (const auto& tile_config : tiles) {
|
||||
for (int stages = min_stages; stages <= max_stages; ++stages) {
|
||||
CutlassGemmConfig config{tile_config, SplitKStyle::NO_SPLIT_K, 1, stages};
|
||||
candidate_configs.push_back(config);
|
||||
}
|
||||
}
|
||||
|
||||
return candidate_configs;
|
||||
}
|
||||
|
||||
CutlassGemmConfig estimate_best_config_from_occupancies(const std::vector<CutlassGemmConfig>& candidate_configs,
|
||||
const std::vector<int>& occupancies, const int64_t m,
|
||||
const int64_t n, const int64_t k, const int64_t,
|
||||
const int split_k_limit, const size_t workspace_bytes,
|
||||
const int multi_processor_count, const int is_weight_only) {
|
||||
if (occupancies.size() != candidate_configs.size()) {
|
||||
ORT_THROW(
|
||||
"[FT Error][estimate_best_config_from_occupancies] occpancies and "
|
||||
"candidate configs vectors must have equal length.");
|
||||
}
|
||||
|
||||
CutlassGemmConfig best_config;
|
||||
// Score will be [0, 1]. The objective is to minimize this score.
|
||||
// It represents the fraction of SM resources unused in the last wave.
|
||||
float config_score = 1.0f;
|
||||
int config_waves = INT_MAX;
|
||||
int current_m_tile = 0;
|
||||
|
||||
const int max_split_k = n >= multi_processor_count * 256 ? 1 : split_k_limit;
|
||||
for (int ii = 0; ii < candidate_configs.size(); ++ii) {
|
||||
CutlassGemmConfig candidate_config = candidate_configs[ii];
|
||||
TileShape tile_shape = get_cta_shape_for_config(candidate_config.tile_config);
|
||||
int occupancy = occupancies[ii];
|
||||
|
||||
if (occupancy == 0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Keep small tile sizes when possible.
|
||||
if (best_config.tile_config != CutlassTileConfig::ChooseWithHeuristic && m < current_m_tile &&
|
||||
current_m_tile < tile_shape.m) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int ctas_in_m_dim = static_cast<int>((m + tile_shape.m - 1) / tile_shape.m);
|
||||
const int ctas_in_n_dim = static_cast<int>((n + tile_shape.n - 1) / tile_shape.n);
|
||||
|
||||
for (int split_k_factor = 1; split_k_factor <= max_split_k; ++split_k_factor) {
|
||||
if (is_valid_split_k_factor(m, n, k, tile_shape, split_k_factor, workspace_bytes, is_weight_only)) {
|
||||
const int ctas_per_wave = occupancy * multi_processor_count;
|
||||
const int ctas_for_problem = ctas_in_m_dim * ctas_in_n_dim * split_k_factor;
|
||||
|
||||
const int num_waves_total = (ctas_for_problem + ctas_per_wave - 1) / ctas_per_wave;
|
||||
const float num_waves_fractional = ctas_for_problem / float(ctas_per_wave);
|
||||
const float current_score = float(num_waves_total) - num_waves_fractional;
|
||||
|
||||
const float score_slack = 0.1f;
|
||||
if (current_score < config_score ||
|
||||
((config_waves > num_waves_total) && (current_score < config_score + score_slack))) {
|
||||
config_score = current_score;
|
||||
config_waves = num_waves_total;
|
||||
SplitKStyle split_style = split_k_factor > 1 ? SplitKStyle::SPLIT_K_SERIAL : SplitKStyle::NO_SPLIT_K;
|
||||
best_config =
|
||||
CutlassGemmConfig{candidate_config.tile_config, split_style, split_k_factor, candidate_config.stages};
|
||||
current_m_tile = tile_shape.m;
|
||||
} else if (current_score == config_score &&
|
||||
(best_config.stages < candidate_config.stages || split_k_factor < best_config.split_k_factor ||
|
||||
current_m_tile < tile_shape.m)) {
|
||||
// Prefer deeper pipeline or smaller split-k
|
||||
SplitKStyle split_style = split_k_factor > 1 ? SplitKStyle::SPLIT_K_SERIAL : SplitKStyle::NO_SPLIT_K;
|
||||
best_config =
|
||||
CutlassGemmConfig{candidate_config.tile_config, split_style, split_k_factor, candidate_config.stages};
|
||||
current_m_tile = tile_shape.m;
|
||||
config_waves = num_waves_total;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (best_config.tile_config == CutlassTileConfig::ChooseWithHeuristic) {
|
||||
ORT_THROW("[FT Error] Heurisitc failed to find a valid config.");
|
||||
}
|
||||
|
||||
return best_config;
|
||||
}
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
39
onnxruntime/contrib_ops/cuda/moe/ft_moe/cutlass_heuristic.h
Normal file
39
onnxruntime/contrib_ops/cuda/moe/ft_moe/cutlass_heuristic.h
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "ft_gemm_configs.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <vector>
|
||||
|
||||
#include "core/common/common.h"
|
||||
|
||||
using namespace onnxruntime;
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
std::vector<CutlassGemmConfig> get_candidate_configs(int sm, const bool is_weight_only, const bool simt_configs_only);
|
||||
|
||||
CutlassGemmConfig estimate_best_config_from_occupancies(const std::vector<CutlassGemmConfig>& candidate_configs,
|
||||
const std::vector<int>& occupancies, const int64_t m,
|
||||
const int64_t n, const int64_t k, const int64_t num_experts,
|
||||
const int split_k_limit, const size_t workspace_bytes,
|
||||
const int multi_processor_count, const int is_weight_only);
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
133
onnxruntime/contrib_ops/cuda/moe/ft_moe/epilogue_helpers.h
Normal file
133
onnxruntime/contrib_ops/cuda/moe/ft_moe/epilogue_helpers.h
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
/**
|
||||
* @file epilogue_helpers.h
|
||||
*
|
||||
* This file includes types for the epilogues. The empty structs exist so we can signal to template
|
||||
* code the type of epilogue we want to run, and let the underlying code specify the details such as
|
||||
* element types, accumulator type and elements per vector access.
|
||||
*
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/epilogue/thread/activation.h"
|
||||
#include "cutlass/epilogue/thread/scale_type.h"
|
||||
#include "cutlass/functional.h"
|
||||
#include "cutlass/half.h"
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination_generic.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination_relu.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination_silu.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace thread {
|
||||
|
||||
__forceinline__ __device__ float copysignf_pos(float a, float b) {
|
||||
float r;
|
||||
r = __int_as_float(__float_as_int(a) | (__float_as_int(b) & 0x80000000));
|
||||
return r;
|
||||
}
|
||||
|
||||
__forceinline__ __device__ float tanh_opt(float x) {
|
||||
#if (__CUDACC_VER_MAJOR__ < 11) || (__CUDA_ARCH__ < 750)
|
||||
const float exp_val = -1.f * fabs(2 * x);
|
||||
return copysignf_pos((1.0f - __expf(exp_val)) / (__expf(exp_val) + 1.0f), x);
|
||||
#else
|
||||
return fast_tanh(x);
|
||||
#endif
|
||||
}
|
||||
|
||||
template <>
|
||||
struct GELU_taylor<float> {
|
||||
static const bool kIsHeavy = true;
|
||||
CUTLASS_DEVICE
|
||||
float operator()(float const& z) const {
|
||||
float k0 = float(0.7978845608028654);
|
||||
float k1 = float(0.044715);
|
||||
|
||||
return float(
|
||||
cutlass::constants::half<float>() * z *
|
||||
(cutlass::constants::one<float>() + tanh_opt(k0 * z * (cutlass::constants::one<float>() + k1 * z * z))));
|
||||
}
|
||||
|
||||
using Params = LinearCombinationGenericParams<float>;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
float operator()(float const& scalar, Params const& params_) const { return this->operator()(scalar); }
|
||||
};
|
||||
|
||||
} // namespace thread
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
struct EpilogueOpBiasSilu {};
|
||||
|
||||
struct EpilogueOpBiasReLU {};
|
||||
|
||||
struct EpilogueOpBiasFtGelu {};
|
||||
|
||||
struct EpilogueOpBias {};
|
||||
|
||||
struct EpilogueOpNoBias {};
|
||||
|
||||
template <typename ElementType, int ElementsPerVectorAccess, typename ElementAccumulator, typename Op>
|
||||
struct Epilogue {};
|
||||
|
||||
template <typename ElementType, int ElementsPerVectorAccess, typename ElementAccumulator>
|
||||
struct Epilogue<ElementType, ElementsPerVectorAccess, ElementAccumulator, EpilogueOpBiasSilu> {
|
||||
using Op = cutlass::epilogue::thread::LinearCombinationSilu<ElementType, ElementsPerVectorAccess, ElementAccumulator,
|
||||
ElementAccumulator,
|
||||
cutlass::epilogue::thread::ScaleType::NoBetaScaling>;
|
||||
};
|
||||
|
||||
template <typename ElementType, int ElementsPerVectorAccess, typename ElementAccumulator>
|
||||
struct Epilogue<ElementType, ElementsPerVectorAccess, ElementAccumulator, EpilogueOpBiasReLU> {
|
||||
using Op = cutlass::epilogue::thread::LinearCombinationRelu<ElementType, ElementsPerVectorAccess, ElementAccumulator,
|
||||
ElementAccumulator,
|
||||
cutlass::epilogue::thread::ScaleType::NoBetaScaling>;
|
||||
};
|
||||
|
||||
template <typename ElementType, int ElementsPerVectorAccess, typename ElementAccumulator>
|
||||
struct Epilogue<ElementType, ElementsPerVectorAccess, ElementAccumulator, EpilogueOpBiasFtGelu> {
|
||||
using Op = cutlass::epilogue::thread::LinearCombinationGeneric<
|
||||
cutlass::epilogue::thread::GELU_taylor, ElementType, ElementsPerVectorAccess, ElementAccumulator,
|
||||
ElementAccumulator, cutlass::epilogue::thread::ScaleType::NoBetaScaling,
|
||||
cutlass::FloatRoundStyle::round_to_nearest, true>;
|
||||
};
|
||||
|
||||
template <typename ElementType, int ElementsPerVectorAccess, typename ElementAccumulator>
|
||||
struct Epilogue<ElementType, ElementsPerVectorAccess, ElementAccumulator, EpilogueOpBias> {
|
||||
using Op = cutlass::epilogue::thread::LinearCombination<ElementType, ElementsPerVectorAccess, ElementAccumulator,
|
||||
ElementAccumulator,
|
||||
cutlass::epilogue::thread::ScaleType::NoBetaScaling>;
|
||||
};
|
||||
|
||||
template <typename ElementType, int ElementsPerVectorAccess, typename ElementAccumulator>
|
||||
struct Epilogue<ElementType, ElementsPerVectorAccess, ElementAccumulator, EpilogueOpNoBias> {
|
||||
using Op =
|
||||
cutlass::epilogue::thread::LinearCombination<ElementType, ElementsPerVectorAccess, ElementAccumulator,
|
||||
ElementAccumulator, cutlass::epilogue::thread::ScaleType::Default>;
|
||||
};
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
58
onnxruntime/contrib_ops/cuda/moe/ft_moe/ft_gemm_configs.h
Normal file
58
onnxruntime/contrib_ops/cuda/moe/ft_moe/ft_gemm_configs.h
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
// Note: The shapes are in the format MxNxK. The K shape of the runtime config MUST match the K shape
|
||||
// in the kernel layout details when doing weight only quantization.
|
||||
enum class CutlassTileConfig {
|
||||
// Signals that we should run heuristics do choose a config
|
||||
Undefined,
|
||||
|
||||
// Signals that we should run heuristics do choose a config
|
||||
ChooseWithHeuristic,
|
||||
|
||||
// SiMT config
|
||||
CtaShape128x128x8_WarpShape64x64x8,
|
||||
|
||||
// TensorCore configs CTA_N = 128, CTA_K = 64
|
||||
// Warp configs for M=32
|
||||
CtaShape32x128x64_WarpShape32x32x64,
|
||||
|
||||
// Warp configs for M=64
|
||||
CtaShape64x128x64_WarpShape32x64x64,
|
||||
CtaShape64x128x64_WarpShape64x32x64,
|
||||
|
||||
// Warp configs for M=128
|
||||
CtaShape128x128x64_WarpShape64x32x64,
|
||||
CtaShape128x128x64_WarpShape128x32x64
|
||||
};
|
||||
|
||||
enum class SplitKStyle {
|
||||
NO_SPLIT_K,
|
||||
SPLIT_K_SERIAL,
|
||||
// SPLIT_K_PARALLEL // Not supported yet
|
||||
};
|
||||
|
||||
struct CutlassGemmConfig {
|
||||
CutlassTileConfig tile_config = CutlassTileConfig::ChooseWithHeuristic;
|
||||
SplitKStyle split_k_style = SplitKStyle::NO_SPLIT_K;
|
||||
int split_k_factor = -1;
|
||||
int stages = -1;
|
||||
};
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
|
|
@ -0,0 +1,79 @@
|
|||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Scheduler for grouped GEMM
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/kernel/gemm_grouped_problem_visitor.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
|
||||
#include "moe_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/// Visitor class to abstract away the algorithm for iterating over tiles
|
||||
template <typename ThreadblockShape, GroupScheduleMode GroupScheduleMode_, int PrefetchTileCount, int ThreadCount,
|
||||
bool Transposed = false>
|
||||
struct GemmMoeProblemVisitor
|
||||
: public MoeProblemVisitor<detail::GemmGroupedProblemSizeHelper<ThreadblockShape, Transposed>, ThreadblockShape,
|
||||
GroupScheduleMode_, PrefetchTileCount, ThreadCount> {
|
||||
static bool const kTransposed = Transposed;
|
||||
|
||||
using ProblemSizeHelper = detail::GemmGroupedProblemSizeHelper<ThreadblockShape, Transposed>;
|
||||
using Base =
|
||||
MoeProblemVisitor<ProblemSizeHelper, ThreadblockShape, GroupScheduleMode_, PrefetchTileCount, ThreadCount>;
|
||||
using Params = typename Base::Params;
|
||||
using SharedStorage = typename Base::SharedStorage;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
GemmMoeProblemVisitor(Params const& params_, SharedStorage& shared_storage_, int32_t block_idx)
|
||||
: Base(params_, shared_storage_, block_idx) {}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
153
onnxruntime/contrib_ops/cuda/moe/ft_moe/layout_traits_helper.h
Normal file
153
onnxruntime/contrib_ops/cuda/moe/ft_moe/layout_traits_helper.h
Normal file
|
|
@ -0,0 +1,153 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
/*
|
||||
This file exists so that we use the same weight layout for MoE grouped gemm and regular gemm when the weight is
|
||||
quantized. The preprocessing code reads this template to know how to organize the quantized weight matrices
|
||||
to be consumed by CUTLASS.
|
||||
|
||||
Note that for int4, ThreadBlockK MUST be 64.
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/arch/mma.h"
|
||||
#include "cutlass/platform/platform.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
template <typename TypeB, typename Arch, typename Enable = void>
|
||||
struct LayoutDetailsB {};
|
||||
|
||||
// Volta specialiations. Volta will dequantize before STS, so we need a different operator
|
||||
template <typename TypeB>
|
||||
struct LayoutDetailsB<TypeB, arch::Sm70> {
|
||||
static constexpr int ThreadblockK = 64;
|
||||
using Layout = layout::RowMajor;
|
||||
static constexpr int ElementsPerAccess = 8;
|
||||
using Operator = cutlass::arch::OpMultiplyAdd;
|
||||
};
|
||||
|
||||
// Specializations for Turing+ when B is FP16. These are currently only used for MoE networks.
|
||||
// TODO - Switch this to column major for weights since gemms should be more performant.
|
||||
template <typename Arch>
|
||||
struct LayoutDetailsB<half_t, Arch, typename platform::enable_if<Arch::kMinComputeCapability >= 75>::type> {
|
||||
static constexpr int ThreadblockK = 64;
|
||||
using Layout = layout::RowMajor;
|
||||
static constexpr int ElementsPerAccess = 128 / cutlass::sizeof_bits<half_t>::value;
|
||||
using Operator = cutlass::arch::OpMultiplyAdd;
|
||||
};
|
||||
|
||||
template <typename TypeA, typename TypeB, typename arch, typename Enable = void>
|
||||
struct MixedGemmArchTraits {};
|
||||
|
||||
template <typename arch>
|
||||
struct MixedGemmArchTraits<float, float, arch> {
|
||||
static constexpr int Stages = 2;
|
||||
using OperatorClass = cutlass::arch::OpClassSimt;
|
||||
using AccType = float;
|
||||
using LayoutB = cutlass::layout::RowMajor;
|
||||
|
||||
static constexpr int ElementsPerAccessA = 1;
|
||||
static constexpr int ElementsPerAccessB = 1;
|
||||
static constexpr int ElementsPerAccessC = 1;
|
||||
static constexpr int ThreadblockK = 8;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<1, 1, 1>;
|
||||
|
||||
using Operator = cutlass::arch::OpMultiplyAdd;
|
||||
};
|
||||
|
||||
// ========================= Volta Traits ===========================
|
||||
// Volta will always dequantize after the global memory load.
|
||||
// This will instantiate any HMMA tensorcore kernels for Volta.
|
||||
template <typename TypeA, typename TypeB>
|
||||
struct MixedGemmArchTraits<
|
||||
TypeA, TypeB, cutlass::arch::Sm70,
|
||||
typename cutlass::platform::enable_if<cutlass::platform::is_same<TypeA, cutlass::half_t>::value>::type> {
|
||||
private:
|
||||
using LayoutDetails = LayoutDetailsB<TypeB, cutlass::arch::Sm70>;
|
||||
|
||||
public:
|
||||
static constexpr int ThreadblockK = LayoutDetails::ThreadblockK;
|
||||
|
||||
using OperatorClass = cutlass::arch::OpClassTensorOp;
|
||||
using AccType = float;
|
||||
using LayoutB = typename LayoutDetails::Layout;
|
||||
|
||||
static constexpr int ElementsPerAccessA = 128 / cutlass::sizeof_bits<TypeA>::value;
|
||||
static constexpr int ElementsPerAccessB = LayoutDetails::ElementsPerAccess;
|
||||
static constexpr int ElementsPerAccessC = 128 / cutlass::sizeof_bits<TypeA>::value;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<8, 8, 4>;
|
||||
|
||||
using Operator = typename LayoutDetails::Operator;
|
||||
};
|
||||
|
||||
// ======================= Turing Traits ==============================
|
||||
template <typename TypeA, typename TypeB>
|
||||
struct MixedGemmArchTraits<
|
||||
TypeA, TypeB, cutlass::arch::Sm75,
|
||||
typename cutlass::platform::enable_if<cutlass::platform::is_same<TypeA, cutlass::half_t>::value>::type> {
|
||||
private:
|
||||
using LayoutDetails = LayoutDetailsB<TypeB, cutlass::arch::Sm75>;
|
||||
|
||||
public:
|
||||
static constexpr int ThreadblockK = LayoutDetails::ThreadblockK;
|
||||
|
||||
using OperatorClass = cutlass::arch::OpClassTensorOp;
|
||||
using AccType = float;
|
||||
using LayoutB = typename LayoutDetails::Layout;
|
||||
|
||||
static constexpr int ElementsPerAccessA = 128 / cutlass::sizeof_bits<TypeA>::value;
|
||||
static constexpr int ElementsPerAccessB = LayoutDetails::ElementsPerAccess;
|
||||
static constexpr int ElementsPerAccessC = 128 / cutlass::sizeof_bits<TypeA>::value;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 8>;
|
||||
|
||||
using Operator = typename LayoutDetails::Operator;
|
||||
};
|
||||
|
||||
// ======================= Ampere Traits ==============================
|
||||
template <typename TypeA, typename TypeB>
|
||||
struct MixedGemmArchTraits<
|
||||
TypeA, TypeB, cutlass::arch::Sm80,
|
||||
typename cutlass::platform::enable_if<cutlass::platform::is_same<TypeA, cutlass::half_t>::value>::type> {
|
||||
private:
|
||||
using LayoutDetails = LayoutDetailsB<TypeB, cutlass::arch::Sm80>;
|
||||
|
||||
public:
|
||||
static constexpr int ThreadblockK = LayoutDetails::ThreadblockK;
|
||||
|
||||
using OperatorClass = cutlass::arch::OpClassTensorOp;
|
||||
using AccType = float;
|
||||
using LayoutB = typename LayoutDetails::Layout;
|
||||
|
||||
static constexpr int ElementsPerAccessA = 128 / cutlass::sizeof_bits<TypeA>::value;
|
||||
static constexpr int ElementsPerAccessB = LayoutDetails::ElementsPerAccess;
|
||||
static constexpr int ElementsPerAccessC = 128 / cutlass::sizeof_bits<TypeA>::value;
|
||||
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>;
|
||||
|
||||
using Operator = typename LayoutDetails::Operator;
|
||||
};
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
463
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_cutlass_kernel.h
Normal file
463
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_cutlass_kernel.h
Normal file
|
|
@ -0,0 +1,463 @@
|
|||
/***************************************************************************************************
|
||||
* Copyright (c) 2017-2021, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without modification, are permitted
|
||||
* provided that the following conditions are met:
|
||||
* * Redistributions of source code must retain the above copyright notice, this list of
|
||||
* conditions and the following disclaimer.
|
||||
* * Redistributions in binary form must reproduce the above copyright notice, this list of
|
||||
* conditions and the following disclaimer in the documentation and/or other materials
|
||||
* provided with the distribution.
|
||||
* * Neither the name of the NVIDIA CORPORATION nor the names of its contributors may be used
|
||||
* to endorse or promote products derived from this software without specific prior written
|
||||
* permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR
|
||||
* IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND
|
||||
* FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
|
||||
* BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS;
|
||||
* OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT,
|
||||
* STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
|
||||
#include "cutlass/gemm/kernel/gemm_transpose_operands.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
#include "gemm_moe_problem_visitor.h"
|
||||
#include "tile_interleaved_layout.h"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// This section exists to that we can use the same kernel code for regular gemm and dequantizing gemms.
|
||||
// It will dispatch to the dequantizing gemm if the Mma type has an Iterator for scales in global.
|
||||
template <typename...>
|
||||
using void_t = void;
|
||||
|
||||
template <typename Mma, typename = void>
|
||||
struct use_dq_gemm : platform::false_type {};
|
||||
|
||||
template <typename Mma>
|
||||
struct use_dq_gemm<Mma, void_t<typename Mma::IteratorScale>> : platform::true_type {};
|
||||
|
||||
// SFINAE overload for dequantizing gemm
|
||||
template <typename Mma, typename ElementScale, typename platform::enable_if<use_dq_gemm<Mma>::value, bool>::type = true>
|
||||
CUTLASS_DEVICE static void run_mma(Mma mma, int gemm_k_iterations, typename Mma::FragmentC& accum,
|
||||
typename Mma::IteratorA iterator_A, typename Mma::IteratorB iterator_B,
|
||||
typename Mma::FragmentC const& src_accum, ElementScale* weight_scale_ptr,
|
||||
MatrixCoord scale_extent, const int thread_idx, MatrixCoord tb_offset_scale) {
|
||||
typename Mma::IteratorScale iterator_scale(Mma::IteratorScale::Layout(scale_extent.column()), weight_scale_ptr,
|
||||
scale_extent, thread_idx, tb_offset_scale);
|
||||
|
||||
mma(gemm_k_iterations, accum, iterator_A, iterator_B, iterator_scale, src_accum);
|
||||
}
|
||||
|
||||
// SFINAE overload for normal gemm. This completely ignores the scale parameters
|
||||
template <typename Mma, typename ElementScale,
|
||||
typename platform::enable_if<!use_dq_gemm<Mma>::value, bool>::type = true>
|
||||
CUTLASS_DEVICE static void run_mma(Mma mma, int gemm_k_iterations, typename Mma::FragmentC& accum,
|
||||
typename Mma::IteratorA iterator_A, typename Mma::IteratorB iterator_B,
|
||||
typename Mma::FragmentC const& src_accum, ElementScale* weight_scale_ptr,
|
||||
MatrixCoord scale_extent, const int thread_idx, MatrixCoord tb_offset_scale) {
|
||||
mma(gemm_k_iterations, accum, iterator_A, iterator_B, src_accum);
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
typename KernelArch, ///! The Architecture this kernel is compiled for. Used since SIMT kernels lose
|
||||
/// top-level
|
||||
/// arch.
|
||||
GroupScheduleMode GroupScheduleMode_ ///! Type of scheduling to perform
|
||||
>
|
||||
struct MoeFCGemm {
|
||||
public:
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
static GroupScheduleMode const kGroupScheduleMode = GroupScheduleMode_;
|
||||
static bool const kTransposed = false;
|
||||
|
||||
// Optional transpose
|
||||
using MapArguments =
|
||||
kernel::detail::MapArguments<typename Mma::IteratorA::Element, typename Mma::IteratorA::Layout, Mma::kTransformA,
|
||||
Mma::IteratorA::AccessType::kElements, typename Mma::IteratorB::Element,
|
||||
typename Mma::IteratorB::Layout, Mma::kTransformB,
|
||||
Mma::IteratorB::AccessType::kElements, typename Mma::LayoutC, kTransposed>;
|
||||
|
||||
// Public-facing type definitions related to operand element type, layout, and complex conjugate
|
||||
// operation. Must interact with the 'kTransposed' notion.
|
||||
static_assert(!kTransposed, "Transpose problem not supported");
|
||||
using ElementA = typename MapArguments::ElementA;
|
||||
using LayoutA = typename MapArguments::LayoutA;
|
||||
using ElementB = typename MapArguments::ElementB;
|
||||
using LayoutB = typename MapArguments::LayoutB;
|
||||
using ElementC = typename Epilogue::OutputTileIterator::Element;
|
||||
using LayoutC = typename MapArguments::LayoutC;
|
||||
using ElementScale = ElementC;
|
||||
|
||||
static ComplexTransform const kTransformA = MapArguments::kTransformA;
|
||||
static ComplexTransform const kTransformB = MapArguments::kTransformB;
|
||||
|
||||
// Type definitions about the mainloop.
|
||||
using Operator = typename Mma::Operator;
|
||||
using OperatorClass = typename Mma::Operator::OperatorClass;
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename Mma::Operator::Shape;
|
||||
using InstructionShape = typename Mma::Policy::Operator::InstructionShape;
|
||||
using ArchTag = typename Mma::ArchTag;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static int const kAlignmentA = MapArguments::kAlignmentA;
|
||||
static int const kAlignmentB = MapArguments::kAlignmentB;
|
||||
static int const kAlignmentC = Epilogue::OutputTileIterator::kElementsPerAccess;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using ProblemVisitor =
|
||||
GemmMoeProblemVisitor<ThreadblockShape, kGroupScheduleMode, kThreadCount, kThreadCount, kTransposed>;
|
||||
|
||||
//
|
||||
// Structures
|
||||
//
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
int problem_count;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
|
||||
ElementA* ptr_A;
|
||||
ElementB* ptr_B;
|
||||
ElementScale* weight_scales;
|
||||
ElementC* ptr_C;
|
||||
ElementC* ptr_D;
|
||||
|
||||
int64_t* total_rows_before_expert;
|
||||
int64_t gemm_n;
|
||||
int64_t gemm_k;
|
||||
|
||||
// Only used by device-level operator
|
||||
GemmCoord* host_problem_sizes;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments()
|
||||
: problem_count(0),
|
||||
threadblock_count(0),
|
||||
ptr_A(nullptr),
|
||||
ptr_B(nullptr),
|
||||
weight_scales(nullptr),
|
||||
ptr_C(nullptr),
|
||||
ptr_D(nullptr),
|
||||
total_rows_before_expert(nullptr),
|
||||
gemm_n(0),
|
||||
gemm_k(0),
|
||||
host_problem_sizes(nullptr) {}
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(int problem_count, int threadblock_count, typename EpilogueOutputOp::Params output_op,
|
||||
const ElementA* ptr_A, const ElementB* ptr_B, const ElementScale* weight_scales, const ElementC* ptr_C,
|
||||
ElementC* ptr_D, int64_t* total_rows_before_expert, int64_t gemm_n, int64_t gemm_k,
|
||||
GemmCoord* host_problem_sizes = nullptr)
|
||||
: problem_count(problem_count),
|
||||
threadblock_count(threadblock_count),
|
||||
output_op(output_op),
|
||||
ptr_A(const_cast<ElementA*>(ptr_A)),
|
||||
ptr_B(const_cast<ElementB*>(ptr_B)),
|
||||
weight_scales(const_cast<ElementScale*>(weight_scales)),
|
||||
ptr_C(const_cast<ElementC*>(ptr_C)),
|
||||
ptr_D(ptr_D),
|
||||
total_rows_before_expert(total_rows_before_expert),
|
||||
gemm_n(gemm_n),
|
||||
gemm_k(gemm_k),
|
||||
host_problem_sizes(nullptr) {
|
||||
if (platform::is_same<uint8_t, ElementB>::value || platform::is_same<uint4b_t, ElementB>::value) {
|
||||
assert(weight_scales);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// Structure for precomputing values in host memory and passing to kernels
|
||||
//
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
typename ProblemVisitor::Params problem_visitor;
|
||||
int threadblock_count;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
|
||||
ElementA* ptr_A;
|
||||
ElementB* ptr_B;
|
||||
ElementScale* weight_scales;
|
||||
ElementC* ptr_C;
|
||||
ElementC* ptr_D;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params() : ptr_A(nullptr), ptr_B(nullptr), weight_scales(nullptr), ptr_C(nullptr), ptr_D(nullptr) {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(Arguments const& args, void* workspace = nullptr, int tile_count = 0)
|
||||
: problem_visitor(args.total_rows_before_expert, args.gemm_n, args.gemm_k, args.problem_count, workspace,
|
||||
tile_count),
|
||||
threadblock_count(args.threadblock_count),
|
||||
output_op(args.output_op),
|
||||
ptr_A(args.ptr_A),
|
||||
ptr_B(args.ptr_B),
|
||||
weight_scales(args.weight_scales),
|
||||
ptr_C(args.ptr_C),
|
||||
ptr_D(args.ptr_D) {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void update(Arguments const& args, void* workspace = nullptr, int tile_count = 0) {
|
||||
problem_visitor = typename ProblemVisitor::Params(args.total_rows_before_expert, args.gemm_n, args.gemm_k,
|
||||
args.problem_count, workspace, tile_count);
|
||||
threadblock_count = args.threadblock_count;
|
||||
output_op = args.output_op;
|
||||
ptr_A = args.ptr_A;
|
||||
ptr_B = args.ptr_B;
|
||||
weight_scales = args.weight_scales;
|
||||
ptr_C = args.ptr_C;
|
||||
ptr_D = args.ptr_D;
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
union SharedStorage {
|
||||
typename ProblemVisitor::SharedStorage problem_visitor;
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
};
|
||||
|
||||
public:
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE
|
||||
MoeFCGemm() {}
|
||||
|
||||
/// Determines whether kernel satisfies alignment
|
||||
static Status can_implement(cutlass::gemm::GemmCoord const& problem_size) { return Status::kSuccess; }
|
||||
|
||||
static Status can_implement(Arguments const& args) {
|
||||
if (args.weight_scales != nullptr) {
|
||||
CUTLASS_TRACE_HOST(
|
||||
"MoeFCGemm::can_implement() - weight scales are ignored for all types except uint8_t and uint4b_t");
|
||||
return Status::kInvalid;
|
||||
}
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
static size_t get_extra_workspace_size(Arguments const& args, cutlass::gemm::GemmCoord const& grid_tiled_shape) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
// The dummy template parameter is not used and exists so that we can compile this code using
|
||||
// a standard earlier than C++17. Prior to C++17, fully specialized templates HAD to exists in
|
||||
// a namespace
|
||||
template <bool B, typename dummy = void>
|
||||
struct KernelRunner {
|
||||
CUTLASS_DEVICE
|
||||
static void run_kernel(Params const& params, SharedStorage& shared_storage) { CUTLASS_NOT_IMPLEMENTED(); }
|
||||
};
|
||||
|
||||
template <typename dummy>
|
||||
struct KernelRunner<true, dummy> {
|
||||
CUTLASS_DEVICE
|
||||
static void run_kernel(Params const& params, SharedStorage& shared_storage) {
|
||||
//
|
||||
// These types shadow the type-level definitions and support the ability to implement
|
||||
// a 'transposed' GEMM that computes the transposed problems.
|
||||
//
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using ElementC = typename Epilogue::OutputTileIterator::Element;
|
||||
using LayoutC = typename Epilogue::OutputTileIterator::Layout;
|
||||
static constexpr int kInterleave = Mma::IteratorB::Shape::kRow / Mma::Shape::kK;
|
||||
static_assert(platform::is_same<LayoutB, layout::RowMajor>::value && kInterleave == 1 ||
|
||||
platform::is_same<LayoutB, layout::ColumnMajor>::value && kInterleave >= 1,
|
||||
"B must be row major/col major OR col major interleaved.");
|
||||
|
||||
//
|
||||
// Problem visitor.
|
||||
//
|
||||
ProblemVisitor problem_visitor(params.problem_visitor, shared_storage.problem_visitor, blockIdx.x);
|
||||
|
||||
const int64_t gemm_k = params.problem_visitor.gemm_k;
|
||||
const int64_t gemm_n = params.problem_visitor.gemm_n;
|
||||
int64_t bytes_per_expert_matrix = (gemm_k * gemm_n / 8) * cutlass::sizeof_bits<ElementB>::value;
|
||||
|
||||
// Outer 'persistent' loop to iterate over tiles
|
||||
while (problem_visitor.next_tile()) {
|
||||
GemmCoord problem_size = problem_visitor.problem_size();
|
||||
int32_t problem_idx = problem_visitor.problem_index();
|
||||
int32_t cta_idx = int32_t(problem_visitor.threadblock_idx());
|
||||
|
||||
GemmCoord grid_shape = problem_visitor.grid_shape(problem_size);
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_offset(int(cta_idx / grid_shape.n()) * Mma::Shape::kM,
|
||||
int(cta_idx % grid_shape.n()) * Mma::Shape::kN, 0);
|
||||
|
||||
// Load element pointers. Exchange pointers and strides if working on the transpose
|
||||
const int64_t rows_to_jump =
|
||||
problem_idx == 0 ? 0 : params.problem_visitor.last_row_for_problem[problem_idx - 1];
|
||||
ElementA* ptr_A = reinterpret_cast<ElementA*>(params.ptr_A) + rows_to_jump * gemm_k;
|
||||
typename LayoutA::LongIndex ldm_A = gemm_k;
|
||||
|
||||
char* byte_ptr_B = ((char*)params.ptr_B) + problem_idx * bytes_per_expert_matrix;
|
||||
ElementB* ptr_B = reinterpret_cast<ElementB*>(byte_ptr_B);
|
||||
typename LayoutB::LongIndex ldm_B =
|
||||
platform::is_same<layout::RowMajor, LayoutB>::value ? gemm_n : gemm_k * kInterleave;
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
cutlass::MatrixCoord tb_offset_A{
|
||||
threadblock_offset.m(),
|
||||
0,
|
||||
};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_B{0, threadblock_offset.n() / kInterleave};
|
||||
|
||||
cutlass::MatrixCoord tb_offset_scale{0, threadblock_offset.n()};
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(LayoutA(ldm_A), ptr_A, {problem_size.m(), problem_size.k()}, thread_idx,
|
||||
tb_offset_A);
|
||||
|
||||
typename Mma::IteratorB iterator_B(LayoutB(ldm_B), ptr_B,
|
||||
{problem_size.k() * kInterleave, problem_size.n() / kInterleave}, thread_idx,
|
||||
tb_offset_B);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Matrix multiply phase
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
int gemm_k_iterations = (problem_size.k() + Mma::Shape::kK - 1) / Mma::Shape::kK;
|
||||
|
||||
// Wait for all threads to finish their epilogue phases from the previous tile.
|
||||
__syncthreads();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
ElementScale* weight_scale_ptr = params.weight_scales + problem_idx * problem_size.n();
|
||||
run_mma<Mma>(mma, gemm_k_iterations, accumulators, iterator_A, iterator_B, accumulators, weight_scale_ptr,
|
||||
{1, problem_size.n()}, thread_idx, tb_offset_scale);
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
ElementC* ptr_C = reinterpret_cast<ElementC*>(params.ptr_C) + problem_idx * gemm_n;
|
||||
ElementC* ptr_D = reinterpret_cast<ElementC*>(params.ptr_D) + rows_to_jump * gemm_n;
|
||||
|
||||
LayoutC layout_C(0);
|
||||
LayoutC layout_D(gemm_n);
|
||||
|
||||
typename Epilogue::OutputTileIterator::Params params_C(layout_C);
|
||||
typename Epilogue::OutputTileIterator::Params params_D(layout_D);
|
||||
|
||||
// Tile iterator loading from source tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_C(params_C, ptr_C, problem_size.mn(), thread_idx,
|
||||
threadblock_offset.mn());
|
||||
|
||||
// Tile iterator writing to destination tensor.
|
||||
typename Epilogue::OutputTileIterator iterator_D(params_D, ptr_D, problem_size.mn(), thread_idx,
|
||||
threadblock_offset.mn());
|
||||
|
||||
Epilogue epilogue(shared_storage.epilogue, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
epilogue(output_op, iterator_D, accumulators, iterator_C);
|
||||
|
||||
// Next tile
|
||||
problem_visitor.advance(gridDim.x);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/*
|
||||
To improve compilation speed, we do not compile the device operator if the CUDA_ARCH does not correspond
|
||||
to the ArchTag of the cutlass kernel operator.
|
||||
*/
|
||||
/// Executes one GEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const& params, SharedStorage& shared_storage) {
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 700) && (__CUDA_ARCH__ < 750)
|
||||
static constexpr bool compile_needed = platform::is_same<KernelArch, arch::Sm70>::value;
|
||||
KernelRunner<compile_needed>::run_kernel(params, shared_storage);
|
||||
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 750) && (__CUDA_ARCH__ < 800)
|
||||
static constexpr bool compile_needed = platform::is_same<KernelArch, arch::Sm75>::value;
|
||||
KernelRunner<compile_needed>::run_kernel(params, shared_storage);
|
||||
#elif defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800) && (__CUDA_ARCH__ < 900)
|
||||
static constexpr bool compile_needed = platform::is_same<KernelArch, arch::Sm80>::value;
|
||||
KernelRunner<compile_needed>::run_kernel(params, shared_storage);
|
||||
#else
|
||||
CUTLASS_NOT_IMPLEMENTED();
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
64
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_gemm_kernels.h
Normal file
64
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_gemm_kernels.h
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda_runtime_api.h>
|
||||
#include "ft_gemm_configs.h"
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
enum class ActivationType { Gelu,
|
||||
Relu,
|
||||
Silu,
|
||||
GeGLU,
|
||||
ReGLU,
|
||||
SiGLU,
|
||||
Identity,
|
||||
InvalidType };
|
||||
|
||||
template <typename T, /*The type used for activations/scales/compute*/
|
||||
typename WeightType /* The type for the MoE weights */>
|
||||
class MoeGemmRunner {
|
||||
public:
|
||||
MoeGemmRunner();
|
||||
|
||||
void initialize(int sm);
|
||||
|
||||
void moe_gemm_bias_act(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, ActivationType activation_type, cudaStream_t stream);
|
||||
|
||||
void moe_gemm(const T* A, const WeightType* B, const T* weight_scales, T* C, int64_t* total_rows_before_expert,
|
||||
int64_t total_rows, int64_t gemm_n, int64_t gemm_k, int num_experts, cudaStream_t stream);
|
||||
|
||||
private:
|
||||
template <typename EpilogueTag>
|
||||
void dispatch_to_arch(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, CutlassGemmConfig gemm_config, cudaStream_t stream, int* occupancy = nullptr);
|
||||
|
||||
template <typename EpilogueTag>
|
||||
void run_gemm(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
cudaStream_t stream);
|
||||
|
||||
private:
|
||||
int sm_;
|
||||
int multi_processor_count_;
|
||||
};
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
|
|
@ -0,0 +1,21 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#include "moe_gemm_kernels_template.h"
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
template class MoeGemmRunner<half, half>;
|
||||
} // namespace ort_fastertransformer
|
||||
|
|
@ -0,0 +1,21 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#include "moe_gemm_kernels_template.h"
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
template class MoeGemmRunner<float, float>;
|
||||
} // namespace ort_fastertransformer
|
||||
|
|
@ -0,0 +1,428 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
// Ignore CUTLASS warnings about type punning
|
||||
#ifdef __GNUC__
|
||||
#pragma GCC diagnostic push
|
||||
#pragma GCC diagnostic ignored "-Wstrict-aliasing"
|
||||
#endif
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
#include "cutlass/layout/matrix.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/gemm/device/gemm_grouped.h"
|
||||
#include "cutlass/gemm/kernel/default_gemm_grouped.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/epilogue/thread/linear_combination_relu.h"
|
||||
|
||||
#include "compute_occupancy.h"
|
||||
#include "epilogue_helpers.h"
|
||||
#include "layout_traits_helper.h"
|
||||
#include "moe_cutlass_kernel.h"
|
||||
|
||||
#ifdef __GNUC__
|
||||
#pragma GCC diagnostic pop
|
||||
#endif
|
||||
|
||||
#include "cutlass_heuristic.h"
|
||||
#include "moe_gemm_kernels.h"
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <math.h>
|
||||
#include <sstream>
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
// ============================= Variable batched Gemm things ===========================
|
||||
template <typename T, typename WeightType, typename arch, typename EpilogueTag, typename ThreadblockShape,
|
||||
typename WarpShape, int Stages>
|
||||
void generic_moe_gemm_kernelLauncher(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
CutlassGemmConfig gemm_config, const int multi_processor_count,
|
||||
cudaStream_t stream, int* kernel_occupancy = nullptr) {
|
||||
if (gemm_config.split_k_style != SplitKStyle::NO_SPLIT_K) {
|
||||
ORT_THROW("[FT Error][MoeGemm] Grouped gemm does not support split-k");
|
||||
}
|
||||
|
||||
static_assert(cutlass::platform::is_same<T, half>::value || cutlass::platform::is_same<T, float>::value,
|
||||
"Specialized for half, float");
|
||||
|
||||
static_assert(cutlass::platform::is_same<T, WeightType>::value ||
|
||||
cutlass::platform::is_same<WeightType, uint8_t>::value ||
|
||||
cutlass::platform::is_same<WeightType, cutlass::uint4b_t>::value,
|
||||
"");
|
||||
|
||||
// The cutlass type for the input elements. This is needed to convert to cutlass::half_t if necessary.
|
||||
using ElementType_ =
|
||||
typename cutlass::platform::conditional<cutlass::platform::is_same<T, half>::value, cutlass::half_t, T>::type;
|
||||
using ElementType = ElementType_;
|
||||
|
||||
using CutlassWeightType_ =
|
||||
typename cutlass::platform::conditional<cutlass::platform::is_same<WeightType, half>::value, cutlass::half_t,
|
||||
WeightType>::type;
|
||||
using CutlassWeightType = CutlassWeightType_;
|
||||
|
||||
// We need separate config for each architecture since we will target different tensorcore instructions. For float,
|
||||
// we do not target TCs.
|
||||
using MixedGemmArchTraits = cutlass::gemm::kernel::MixedGemmArchTraits<ElementType, CutlassWeightType, arch>;
|
||||
using ElementAccumulator = typename MixedGemmArchTraits::AccType;
|
||||
|
||||
using EpilogueOp =
|
||||
typename Epilogue<ElementType, MixedGemmArchTraits::ElementsPerAccessC, ElementAccumulator, EpilogueTag>::Op;
|
||||
|
||||
// Finally, set up the kernel.
|
||||
using GemmKernel_ = typename cutlass::gemm::kernel::DefaultGemmGrouped<
|
||||
ElementType, cutlass::layout::RowMajor, cutlass::ComplexTransform::kNone, MixedGemmArchTraits::ElementsPerAccessA,
|
||||
CutlassWeightType, typename MixedGemmArchTraits::LayoutB, cutlass::ComplexTransform::kNone,
|
||||
MixedGemmArchTraits::ElementsPerAccessB, ElementType, cutlass::layout::RowMajor, ElementAccumulator,
|
||||
typename MixedGemmArchTraits::OperatorClass, arch, ThreadblockShape, WarpShape,
|
||||
typename MixedGemmArchTraits::InstructionShape, EpilogueOp,
|
||||
cutlass::gemm::threadblock::GemmBatchedIdentityThreadblockSwizzle, Stages,
|
||||
cutlass::gemm::kernel::GroupScheduleMode::kDeviceOnly, typename MixedGemmArchTraits::Operator>::GemmKernel;
|
||||
|
||||
using GemmKernel = cutlass::gemm::kernel::MoeFCGemm<typename GemmKernel_::Mma, typename GemmKernel_::Epilogue,
|
||||
typename GemmKernel_::ThreadblockSwizzle,
|
||||
arch, // Ensure top level arch is used for dispatch
|
||||
GemmKernel_::kGroupScheduleMode>;
|
||||
|
||||
using GemmGrouped = cutlass::gemm::device::GemmGrouped<GemmKernel>;
|
||||
|
||||
if (kernel_occupancy != nullptr) {
|
||||
*kernel_occupancy = compute_occupancy_for_kernel<GemmKernel>();
|
||||
return;
|
||||
}
|
||||
int occupancy = std::min(2, GemmGrouped::maximum_active_blocks());
|
||||
if (occupancy == 0) {
|
||||
ORT_THROW("[FT Error][MoE Runner] GPU lacks the shared memory resources to run GroupedGEMM kernel");
|
||||
}
|
||||
const int threadblock_count = multi_processor_count * occupancy;
|
||||
|
||||
typename EpilogueOp::Params epilogue_op(ElementAccumulator(1.f), ElementAccumulator(0.f));
|
||||
|
||||
typename GemmGrouped::Arguments args(
|
||||
num_experts, threadblock_count, epilogue_op, reinterpret_cast<const ElementType*>(A),
|
||||
reinterpret_cast<const CutlassWeightType*>(B), reinterpret_cast<const ElementType*>(weight_scales),
|
||||
reinterpret_cast<const ElementType*>(biases), reinterpret_cast<ElementType*>(C), total_rows_before_expert, gemm_n,
|
||||
gemm_k);
|
||||
|
||||
GemmGrouped gemm;
|
||||
|
||||
auto can_implement = gemm.can_implement(args);
|
||||
if (can_implement != cutlass::Status::kSuccess) {
|
||||
std::string err_msg =
|
||||
"MoEFC kernel will fail for params. Error: " + std::string(cutlassGetStatusString(can_implement));
|
||||
ORT_THROW("[FT Error][MoE Runner] " + err_msg);
|
||||
}
|
||||
|
||||
auto init_status = gemm.initialize(args);
|
||||
if (init_status != cutlass::Status::kSuccess) {
|
||||
std::string err_msg = "Failed to initialize cutlass variable batched gemm. Error: " +
|
||||
std::string(cutlassGetStatusString(init_status));
|
||||
ORT_THROW("[FT Error][MoE Runner] " + err_msg);
|
||||
}
|
||||
|
||||
auto run_status = gemm.run(stream);
|
||||
if (run_status != cutlass::Status::kSuccess) {
|
||||
std::string err_msg =
|
||||
"Failed to run cutlass variable batched gemm. Error: " + std::string(cutlassGetStatusString(run_status));
|
||||
ORT_THROW("[FT Error][MoE Runner] " + err_msg);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename arch, typename EpilogueTag, typename ThreadblockShape,
|
||||
typename WarpShape, int Stages, typename Enable = void>
|
||||
struct dispatch_stages {
|
||||
static void dispatch(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
CutlassGemmConfig gemm_config, int multi_processor_count, cudaStream_t stream,
|
||||
int* occupancy = nullptr) {
|
||||
std::string err_msg = "Cutlass fpA_intB gemm. Not instantiates for arch " +
|
||||
std::to_string(arch::kMinComputeCapability) + " with stages set to " + std::to_string(Stages);
|
||||
ORT_THROW("[FT Error][dispatch_stages::dispatch] " + err_msg);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, typename WeightType, typename arch, typename EpilogueTag, typename ThreadblockShape,
|
||||
typename WarpShape>
|
||||
struct dispatch_stages<T, WeightType, arch, EpilogueTag, ThreadblockShape, WarpShape, 2> {
|
||||
static void dispatch(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
CutlassGemmConfig gemm_config, int multi_processor_count, cudaStream_t stream,
|
||||
int* occupancy = nullptr) {
|
||||
generic_moe_gemm_kernelLauncher<T, WeightType, arch, EpilogueTag, ThreadblockShape, WarpShape, 2>(
|
||||
A, B, weight_scales, biases, C, total_rows_before_expert, gemm_n, gemm_k, num_experts, gemm_config,
|
||||
multi_processor_count, stream, occupancy);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, typename WeightType, typename EpilogueTag, typename ThreadblockShape, typename WarpShape,
|
||||
int Stages>
|
||||
struct dispatch_stages<T, WeightType, cutlass::arch::Sm80, EpilogueTag, ThreadblockShape, WarpShape, Stages,
|
||||
typename std::enable_if<(Stages > 2)>::type> {
|
||||
static void dispatch(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
CutlassGemmConfig gemm_config, int multi_processor_count, cudaStream_t stream,
|
||||
int* occupancy = nullptr) {
|
||||
generic_moe_gemm_kernelLauncher<T, WeightType, cutlass::arch::Sm80, EpilogueTag, ThreadblockShape, WarpShape,
|
||||
Stages>(A, B, weight_scales, biases, C, total_rows_before_expert, gemm_n, gemm_k,
|
||||
num_experts, gemm_config, multi_processor_count, stream, occupancy);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, typename WeightType, typename arch, typename EpilogueTag, typename ThreadblockShape,
|
||||
typename WarpShape>
|
||||
void dispatch_gemm_config(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
CutlassGemmConfig gemm_config, int multi_processor_count, cudaStream_t stream,
|
||||
int* occupancy = nullptr) {
|
||||
switch (gemm_config.stages) {
|
||||
case 2:
|
||||
using DispatcherStages2 = dispatch_stages<T, WeightType, arch, EpilogueTag, ThreadblockShape, WarpShape, 2>;
|
||||
DispatcherStages2::dispatch(A, B, weight_scales, biases, C, total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case 3:
|
||||
using DispatcherStages3 = dispatch_stages<T, WeightType, arch, EpilogueTag, ThreadblockShape, WarpShape, 3>;
|
||||
DispatcherStages3::dispatch(A, B, weight_scales, biases, C, total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case 4:
|
||||
using DispatcherStages4 = dispatch_stages<T, WeightType, arch, EpilogueTag, ThreadblockShape, WarpShape, 4>;
|
||||
DispatcherStages4::dispatch(A, B, weight_scales, biases, C, total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
default:
|
||||
std::string err_msg = "dispatch_gemm_config does not support stages " + std::to_string(gemm_config.stages);
|
||||
ORT_THROW("[FT Error][MoE][dispatch_gemm_config] " + err_msg);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// This overload will handle tensorop gemms. It is disabled via SFINAE for fp32.
|
||||
// This overload is only enabled when T == WeightType.
|
||||
template <
|
||||
typename T, typename WeightType, typename arch, typename EpilogueTag,
|
||||
typename std::enable_if<!std::is_same<T, float>::value && std::is_same<T, WeightType>::value>::type* = nullptr>
|
||||
void dispatch_moe_gemm_to_cutlass(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, CutlassGemmConfig gemm_config, int sm_version,
|
||||
int multi_processor_count, cudaStream_t stream, int* occupancy = nullptr) {
|
||||
switch (gemm_config.tile_config) {
|
||||
case CutlassTileConfig::CtaShape32x128x64_WarpShape32x32x64:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<32, 128, 64>,
|
||||
cutlass::gemm::GemmShape<32, 32, 64>>(A, B, weight_scales, biases, C,
|
||||
total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::CtaShape64x128x64_WarpShape32x64x64:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<64, 128, 64>,
|
||||
cutlass::gemm::GemmShape<32, 64, 64>>(A, B, weight_scales, biases, C,
|
||||
total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::CtaShape128x128x64_WarpShape64x32x64:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<128, 128, 64>,
|
||||
cutlass::gemm::GemmShape<64, 32, 64>>(A, B, weight_scales, biases, C,
|
||||
total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::Undefined:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass] gemm config undefined.");
|
||||
break;
|
||||
case CutlassTileConfig::ChooseWithHeuristic:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass] gemm config should have already been set by heuristic.");
|
||||
break;
|
||||
default:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass] Config is invalid for same type MoE tensorop GEMM.");
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Tensorop GEMM overload
|
||||
// Overload for quantize MoE GEMMs. We disable some warp configs here since they will not be used and we can improve
|
||||
// compile time
|
||||
template <
|
||||
typename T, typename WeightType, typename arch, typename EpilogueTag,
|
||||
typename std::enable_if<!std::is_same<T, float>::value && !std::is_same<T, WeightType>::value>::type* = nullptr>
|
||||
void dispatch_moe_gemm_to_cutlass(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, CutlassGemmConfig gemm_config, int sm_version,
|
||||
int multi_processor_count, cudaStream_t stream, int* occupancy = nullptr) {
|
||||
switch (gemm_config.tile_config) {
|
||||
case CutlassTileConfig::CtaShape32x128x64_WarpShape32x32x64:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<32, 128, 64>,
|
||||
cutlass::gemm::GemmShape<32, 32, 64>>(A, B, weight_scales, biases, C,
|
||||
total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::CtaShape64x128x64_WarpShape64x32x64:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<64, 128, 64>,
|
||||
cutlass::gemm::GemmShape<64, 32, 64>>(A, B, weight_scales, biases, C,
|
||||
total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::CtaShape128x128x64_WarpShape128x32x64:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<128, 128, 64>,
|
||||
cutlass::gemm::GemmShape<128, 32, 64>>(
|
||||
A, B, weight_scales, biases, C, total_rows_before_expert, gemm_n, gemm_k, num_experts, gemm_config,
|
||||
multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::Undefined:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass] gemm config undefined.");
|
||||
break;
|
||||
case CutlassTileConfig::ChooseWithHeuristic:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass] gemm config should have already been set by heuristic.");
|
||||
break;
|
||||
default:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass] Config is invalid for mixed type tensorop GEMM.");
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// This overload will handle simt gemms. It is disabled via SFINAE for tensorop.
|
||||
template <typename T, typename WeightType, typename arch, typename EpilogueTag,
|
||||
typename std::enable_if<std::is_same<T, float>::value>::type* = nullptr>
|
||||
void dispatch_moe_gemm_to_cutlass(const T* A, const WeightType* B, const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, CutlassGemmConfig gemm_config, int sm_version,
|
||||
int multi_processor_count, cudaStream_t stream, int* occupancy = nullptr) {
|
||||
switch (gemm_config.tile_config) {
|
||||
case CutlassTileConfig::CtaShape128x128x8_WarpShape64x64x8:
|
||||
dispatch_gemm_config<T, WeightType, arch, EpilogueTag, cutlass::gemm::GemmShape<128, 128, 8>,
|
||||
cutlass::gemm::GemmShape<64, 64, 8>>(A, B, weight_scales, biases, C,
|
||||
total_rows_before_expert, gemm_n, gemm_k, num_experts,
|
||||
gemm_config, multi_processor_count, stream, occupancy);
|
||||
break;
|
||||
case CutlassTileConfig::Undefined:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass][SIMT] gemm config undefined.");
|
||||
break;
|
||||
case CutlassTileConfig::ChooseWithHeuristic:
|
||||
ORT_THROW(
|
||||
"[FT Error][dispatch_moe_gemm_to_cutlass][SIMT] gemm config should have already been set by heuristic.");
|
||||
break;
|
||||
default:
|
||||
ORT_THROW("[FT Error][dispatch_moe_gemm_to_cutlass][SIMT] Unsupported config for float MoE gemm.");
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType>
|
||||
MoeGemmRunner<T, WeightType>::MoeGemmRunner() {}
|
||||
|
||||
template <typename T, typename WeightType>
|
||||
void MoeGemmRunner<T, WeightType>::initialize(int sm_version) {
|
||||
int device{-1};
|
||||
cudaGetDevice(&device);
|
||||
sm_ = sm_version;
|
||||
cudaDeviceGetAttribute(&multi_processor_count_, cudaDevAttrMultiProcessorCount, device);
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType>
|
||||
template <typename EpilogueTag>
|
||||
void MoeGemmRunner<T, WeightType>::dispatch_to_arch<EpilogueTag>(const T* A, const WeightType* B,
|
||||
const T* weight_scales, const T* biases, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows,
|
||||
int64_t gemm_n, int64_t gemm_k, int num_experts,
|
||||
CutlassGemmConfig gemm_config, cudaStream_t stream,
|
||||
int* occupancy) {
|
||||
if (sm_ >= 70 && sm_ < 75) {
|
||||
dispatch_moe_gemm_to_cutlass<T, WeightType, cutlass::arch::Sm70, EpilogueTag>(
|
||||
A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k, num_experts, gemm_config,
|
||||
sm_, multi_processor_count_, stream, occupancy);
|
||||
} else if (sm_ >= 75 && sm_ < 80) {
|
||||
dispatch_moe_gemm_to_cutlass<T, WeightType, cutlass::arch::Sm75, EpilogueTag>(
|
||||
A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k, num_experts, gemm_config,
|
||||
sm_, multi_processor_count_, stream, occupancy);
|
||||
} else if (sm_ >= 80 && sm_ < 90) {
|
||||
dispatch_moe_gemm_to_cutlass<T, WeightType, cutlass::arch::Sm80, EpilogueTag>(
|
||||
A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k, num_experts, gemm_config,
|
||||
sm_, multi_processor_count_, stream, occupancy);
|
||||
} else {
|
||||
ORT_THROW("[FT Error][MoE][GEMM Dispatch] Arch unsupported for MoE GEMM");
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType>
|
||||
template <typename EpilogueTag>
|
||||
void MoeGemmRunner<T, WeightType>::run_gemm<EpilogueTag>(const T* A, const WeightType* B, const T* weight_scales,
|
||||
const T* biases, T* C, int64_t* total_rows_before_expert,
|
||||
int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, cudaStream_t stream) {
|
||||
static constexpr bool is_weight_only = !std::is_same<T, WeightType>::value;
|
||||
static constexpr bool only_simt_configs = std::is_same<T, float>::value;
|
||||
std::vector<CutlassGemmConfig> candidate_configs = get_candidate_configs(sm_, is_weight_only, only_simt_configs);
|
||||
std::vector<int> occupancies(candidate_configs.size());
|
||||
|
||||
for (size_t ii = 0; ii < candidate_configs.size(); ++ii) {
|
||||
dispatch_to_arch<EpilogueTag>(A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k,
|
||||
num_experts, candidate_configs[ii], stream, &occupancies[ii]);
|
||||
}
|
||||
|
||||
static constexpr int workspace_bytes = 0; // No workspace for MoE GEMMs.
|
||||
static constexpr int split_k_limit = 1; // MoE GEMM does not support split-k.
|
||||
CutlassGemmConfig chosen_config =
|
||||
estimate_best_config_from_occupancies(candidate_configs, occupancies, total_rows, gemm_n, gemm_k, num_experts,
|
||||
split_k_limit, workspace_bytes, multi_processor_count_, is_weight_only);
|
||||
|
||||
dispatch_to_arch<EpilogueTag>(A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k,
|
||||
num_experts, chosen_config, stream);
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType>
|
||||
void MoeGemmRunner<T, WeightType>::moe_gemm_bias_act(const T* A, const WeightType* B, const T* weight_scales,
|
||||
const T* biases, T* C, int64_t* total_rows_before_expert,
|
||||
int64_t total_rows, int64_t gemm_n, int64_t gemm_k,
|
||||
int num_experts, ActivationType activation_type,
|
||||
cudaStream_t stream) {
|
||||
switch (activation_type) {
|
||||
case ActivationType::Relu:
|
||||
run_gemm<EpilogueOpBiasReLU>(A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k,
|
||||
num_experts, stream);
|
||||
break;
|
||||
case ActivationType::Gelu:
|
||||
run_gemm<EpilogueOpBiasFtGelu>(A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n,
|
||||
gemm_k, num_experts, stream);
|
||||
break;
|
||||
case ActivationType::Silu:
|
||||
run_gemm<EpilogueOpBiasSilu>(A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k,
|
||||
num_experts, stream);
|
||||
break;
|
||||
case ActivationType::Identity:
|
||||
run_gemm<EpilogueOpBias>(A, B, weight_scales, biases, C, total_rows_before_expert, total_rows, gemm_n, gemm_k,
|
||||
num_experts, stream);
|
||||
break;
|
||||
case ActivationType::InvalidType:
|
||||
ORT_THROW("[FT Error][MoE Runner] Invalid activation type for MoE GEMM");
|
||||
break;
|
||||
default: {
|
||||
ORT_THROW("[FT Error][MoE Runner] Invalid activation type for MoE GEMM");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType>
|
||||
void MoeGemmRunner<T, WeightType>::moe_gemm(const T* A, const WeightType* B, const T* weight_scales, T* C,
|
||||
int64_t* total_rows_before_expert, int64_t total_rows, int64_t gemm_n,
|
||||
int64_t gemm_k, int num_experts, cudaStream_t stream) {
|
||||
run_gemm<EpilogueOpNoBias>(A, B, weight_scales, nullptr, C, total_rows_before_expert, total_rows, gemm_n, gemm_k,
|
||||
num_experts, stream);
|
||||
}
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
830
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_kernel.cu
Normal file
830
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_kernel.cu
Normal file
|
|
@ -0,0 +1,830 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#include <cuda.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <math.h>
|
||||
#include <sstream>
|
||||
#include <algorithm>
|
||||
|
||||
// Ignore CUTLASS warnings about type punning
|
||||
#ifdef __GNUC__
|
||||
#pragma GCC diagnostic push
|
||||
#pragma GCC diagnostic ignored "-Wstrict-aliasing"
|
||||
#endif
|
||||
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/numeric_conversion.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
|
||||
#ifdef __GNUC__
|
||||
#pragma GCC diagnostic pop
|
||||
#endif
|
||||
|
||||
#include "moe_kernel.h"
|
||||
|
||||
#if CUDA_VERSION >= 11000
|
||||
#include <cub/cub.cuh>
|
||||
#include <cub/device/device_radix_sort.cuh>
|
||||
#include <cub/util_type.cuh>
|
||||
#else
|
||||
#include "cub/cub.cuh"
|
||||
#include "cub/device/device_radix_sort.cuh"
|
||||
#include "cub/util_type.cuh"
|
||||
#endif
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
static constexpr int WARP_SIZE = 32;
|
||||
|
||||
// ====================== Softmax things ===============================
|
||||
// We have our own implementation of softmax here so we can support transposing the output
|
||||
// in the softmax kernel when we extend this module to support expert-choice routing.
|
||||
template <typename T, int TPB>
|
||||
__launch_bounds__(TPB) __global__
|
||||
void moe_softmax(const T* input, const bool* finished, T* output, const int num_cols) {
|
||||
using BlockReduce = cub::BlockReduce<float, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||
|
||||
__shared__ float normalizing_factor;
|
||||
__shared__ float float_max;
|
||||
|
||||
const int thread_row_offset = blockIdx.x * num_cols;
|
||||
|
||||
cub::Sum sum;
|
||||
float threadData(-FLT_MAX);
|
||||
|
||||
// Don't touch finished rows.
|
||||
if ((finished != nullptr) && finished[blockIdx.x]) {
|
||||
return;
|
||||
}
|
||||
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
threadData = max(static_cast<float>(input[idx]), threadData);
|
||||
}
|
||||
|
||||
const float maxElem = BlockReduce(tmpStorage).Reduce(threadData, cub::Max());
|
||||
if (threadIdx.x == 0) {
|
||||
float_max = maxElem;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
threadData = 0;
|
||||
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
threadData += exp((static_cast<float>(input[idx]) - float_max));
|
||||
}
|
||||
|
||||
const auto Z = BlockReduce(tmpStorage).Reduce(threadData, sum);
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
normalizing_factor = 1.f / Z;
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||
const int idx = thread_row_offset + ii;
|
||||
const float val = exp((static_cast<float>(input[idx]) - float_max)) * normalizing_factor;
|
||||
output[idx] = T(val);
|
||||
}
|
||||
}
|
||||
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ < 530
|
||||
template <typename T, int TPB>
|
||||
__launch_bounds__(TPB) __global__ void moe_top_k(const T*, const bool*, T*, int*, int*, int, const int) {
|
||||
// Does not support pre-Kepler architectures
|
||||
;
|
||||
}
|
||||
#else
|
||||
template <typename T, int TPB>
|
||||
__launch_bounds__(TPB) __global__ void moe_top_k(const T* inputs_after_softmax, const bool* finished, T* output,
|
||||
int* indices, int* source_rows, int num_experts, int k) {
|
||||
using cub_kvp = cub::KeyValuePair<int, T>;
|
||||
using BlockReduce = cub::BlockReduce<cub_kvp, TPB>;
|
||||
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||
|
||||
cub_kvp thread_kvp;
|
||||
cub::ArgMax arg_max;
|
||||
|
||||
int num_rows = gridDim.x;
|
||||
const int block_row = blockIdx.x;
|
||||
|
||||
const bool should_process_row = finished ? !finished[block_row] : true;
|
||||
const int thread_read_offset = blockIdx.x * num_experts;
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
thread_kvp.key = 0;
|
||||
thread_kvp.value = T(-1.f); // This is OK because inputs are probabilities
|
||||
|
||||
cub_kvp inp_kvp;
|
||||
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||
const int idx = thread_read_offset + expert;
|
||||
inp_kvp.key = expert;
|
||||
inp_kvp.value = inputs_after_softmax[idx];
|
||||
|
||||
for (int prior_k = 0; prior_k < k_idx; ++prior_k) {
|
||||
const int prior_winning_expert = indices[k * block_row + prior_k];
|
||||
|
||||
if (prior_winning_expert == expert) {
|
||||
inp_kvp = thread_kvp;
|
||||
}
|
||||
}
|
||||
|
||||
thread_kvp = arg_max(inp_kvp, thread_kvp);
|
||||
}
|
||||
|
||||
const cub_kvp result_kvp = BlockReduce(tmpStorage).Reduce(thread_kvp, arg_max);
|
||||
if (threadIdx.x == 0) {
|
||||
const int idx = k * block_row + k_idx;
|
||||
output[idx] = result_kvp.value;
|
||||
indices[idx] = should_process_row ? result_kvp.key : num_experts;
|
||||
source_rows[idx] = k_idx * num_rows + block_row;
|
||||
}
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
// ====================== TopK softmax things ===============================
|
||||
|
||||
/*
|
||||
A Top-K gating softmax written to exploit when the number of experts in the MoE layers
|
||||
are a small power of 2. This allows us to cleanly share the rows among the threads in
|
||||
a single warp and eliminate communication between warps (so no need to use shared mem).
|
||||
|
||||
It fuses the softmax, max and argmax into a single kernel.
|
||||
|
||||
Limitations:
|
||||
1) This implementation is intended for when the number of experts is a small power of 2.
|
||||
2) This implementation assumes k is small, but will work for any k.
|
||||
*/
|
||||
|
||||
template <typename T, int VPT, int NUM_EXPERTS, int WARPS_PER_CTA, int BYTES_PER_LDG>
|
||||
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE) __global__
|
||||
void topk_gating_softmax(const T* input, const bool* finished, T* output, int num_rows, int* indices,
|
||||
int* source_rows, int k) {
|
||||
// We begin by enforcing compile time assertions and setting up compile time constants.
|
||||
static_assert(VPT == (VPT & -VPT), "VPT must be power of 2");
|
||||
static_assert(NUM_EXPERTS == (NUM_EXPERTS & -NUM_EXPERTS), "NUM_EXPERTS must be power of 2");
|
||||
static_assert(BYTES_PER_LDG == (BYTES_PER_LDG & -BYTES_PER_LDG), "BYTES_PER_LDG must be power of 2");
|
||||
static_assert(BYTES_PER_LDG <= 16, "BYTES_PER_LDG must be leq 16");
|
||||
|
||||
// Number of bytes each thread pulls in per load
|
||||
static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T);
|
||||
static constexpr int ELTS_PER_ROW = NUM_EXPERTS;
|
||||
static constexpr int THREADS_PER_ROW = ELTS_PER_ROW / VPT;
|
||||
static constexpr int LDG_PER_THREAD = VPT / ELTS_PER_LDG;
|
||||
|
||||
// Restrictions based on previous section.
|
||||
static_assert(VPT % ELTS_PER_LDG == 0, "The elements per thread must be a multiple of the elements per ldg");
|
||||
static_assert(WARP_SIZE % THREADS_PER_ROW == 0, "The threads per row must cleanly divide the threads per warp");
|
||||
static_assert(THREADS_PER_ROW == (THREADS_PER_ROW & -THREADS_PER_ROW), "THREADS_PER_ROW must be power of 2");
|
||||
static_assert(THREADS_PER_ROW <= WARP_SIZE, "THREADS_PER_ROW can be at most warp size");
|
||||
|
||||
// We have NUM_EXPERTS elements per row. We specialize for small #experts
|
||||
static constexpr int ELTS_PER_WARP = WARP_SIZE * VPT;
|
||||
static constexpr int ROWS_PER_WARP = ELTS_PER_WARP / ELTS_PER_ROW;
|
||||
static constexpr int ROWS_PER_CTA = WARPS_PER_CTA * ROWS_PER_WARP;
|
||||
|
||||
// Restrictions for previous section.
|
||||
static_assert(ELTS_PER_WARP % ELTS_PER_ROW == 0, "The elts per row must cleanly divide the total elt per warp");
|
||||
|
||||
// ===================== From this point, we finally start computing run-time variables. ========================
|
||||
|
||||
// Compute CTA and warp rows. We pack multiple rows into a single warp, and a block contains WARPS_PER_CTA warps.
|
||||
// This, each block processes a chunk of rows. We start by computing the start row for each block.
|
||||
const int cta_base_row = blockIdx.x * ROWS_PER_CTA;
|
||||
|
||||
// Now, using the base row per thread block, we compute the base row per warp.
|
||||
const int warp_base_row = cta_base_row + threadIdx.y * ROWS_PER_WARP;
|
||||
|
||||
// The threads in a warp are split into sub-groups that will work on a row.
|
||||
// We compute row offset for each thread sub-group
|
||||
const int thread_row_in_warp = threadIdx.x / THREADS_PER_ROW;
|
||||
const int thread_row = warp_base_row + thread_row_in_warp;
|
||||
|
||||
// Threads with indices out of bounds should early exit here.
|
||||
if (thread_row >= num_rows) return;
|
||||
const bool should_process_row = finished ? !finished[thread_row] : true;
|
||||
|
||||
// We finally start setting up the read pointers for each thread. First, each thread jumps to the start of the
|
||||
// row it will read.
|
||||
const T* thread_row_ptr = input + thread_row * ELTS_PER_ROW;
|
||||
|
||||
// Now, we compute the group each thread belong to in order to determine the first column to start loads.
|
||||
const int thread_group_idx = threadIdx.x % THREADS_PER_ROW;
|
||||
const int first_elt_read_by_thread = thread_group_idx * ELTS_PER_LDG;
|
||||
const T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread;
|
||||
|
||||
// Determine the pointer type to use to read in the data depending on the BYTES_PER_LDG template param. In theory,
|
||||
// this can support all powers of 2 up to 16.
|
||||
using AccessType = cutlass::AlignedArray<T, ELTS_PER_LDG>;
|
||||
|
||||
// Finally, we pull in the data from global mem
|
||||
cutlass::Array<T, VPT> row_chunk_input;
|
||||
AccessType* row_chunk_vec_ptr = reinterpret_cast<AccessType*>(&row_chunk_input);
|
||||
const AccessType* vec_thread_read_ptr = reinterpret_cast<const AccessType*>(thread_read_ptr);
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < LDG_PER_THREAD; ++ii) {
|
||||
row_chunk_vec_ptr[ii] = vec_thread_read_ptr[ii * THREADS_PER_ROW];
|
||||
}
|
||||
|
||||
using ComputeType = float;
|
||||
using Converter = cutlass::NumericArrayConverter<ComputeType, T, VPT>;
|
||||
Converter compute_type_converter;
|
||||
cutlass::Array<ComputeType, VPT> row_chunk = compute_type_converter(row_chunk_input);
|
||||
|
||||
// First, we perform a max reduce within the thread. We can do the max in fp16 safely (I think) and just
|
||||
// convert to float afterwards for the exp + sum reduction.
|
||||
ComputeType thread_max = row_chunk[0];
|
||||
#pragma unroll
|
||||
for (int ii = 1; ii < VPT; ++ii) {
|
||||
thread_max = max(thread_max, row_chunk[ii]);
|
||||
}
|
||||
|
||||
// Now, we find the max within the thread group and distribute among the threads. We use a butterfly reduce.
|
||||
#pragma unroll
|
||||
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||
thread_max = max(thread_max, __shfl_xor_sync(0xFFFFFFFF, thread_max, mask, THREADS_PER_ROW));
|
||||
}
|
||||
|
||||
// From this point, thread max in all the threads have the max within the row.
|
||||
// Now, we subtract the max from each element in the thread and take the exp. We also compute the thread local sum.
|
||||
float row_sum = 0;
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
row_chunk[ii] = expf(row_chunk[ii] - thread_max);
|
||||
row_sum += row_chunk[ii];
|
||||
}
|
||||
|
||||
// Now, we perform the sum reduce within each thread group. Similar to the max reduce, we use a bufferfly pattern.
|
||||
#pragma unroll
|
||||
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||
row_sum += __shfl_xor_sync(0xFFFFFFFF, row_sum, mask, THREADS_PER_ROW);
|
||||
}
|
||||
|
||||
// From this point, all threads have the max and the sum for their rows in the thread_max and thread_sum variables
|
||||
// respectively. Finally, we can scale the rows for the softmax. Technically, for top-k gating we don't need to
|
||||
// compute the entire softmax row. We can likely look at the maxes and only compute for the top-k values in the row.
|
||||
// However, this kernel will likely not be a bottle neck and it seems better to closer match torch and find the
|
||||
// argmax after computing the softmax.
|
||||
const float reciprocal_row_sum = 1.f / row_sum;
|
||||
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < VPT; ++ii) {
|
||||
row_chunk[ii] = row_chunk[ii] * reciprocal_row_sum;
|
||||
}
|
||||
|
||||
// Now, softmax_res contains the softmax of the row chunk. Now, I want to find the topk elements in each row, along
|
||||
// with the max index.
|
||||
int start_col = first_elt_read_by_thread;
|
||||
static constexpr int COLS_PER_GROUP_LDG = ELTS_PER_LDG * THREADS_PER_ROW;
|
||||
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
// First, each thread does the local argmax
|
||||
float max_val = row_chunk[0];
|
||||
int expert = start_col;
|
||||
#pragma unroll
|
||||
for (int ldg = 0, col = start_col; ldg < LDG_PER_THREAD; ++ldg, col += COLS_PER_GROUP_LDG) {
|
||||
#pragma unroll
|
||||
for (int ii = 0; ii < ELTS_PER_LDG; ++ii) {
|
||||
float val = row_chunk[ldg * ELTS_PER_LDG + ii];
|
||||
|
||||
// No check on the experts here since columns with the smallest index are processed first and only
|
||||
// updated if > (not >=)
|
||||
if (val > max_val) {
|
||||
max_val = val;
|
||||
expert = col + ii;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Now, we perform the argmax reduce. We use the butterfly pattern so threads reach consensus about the max.
|
||||
// This will be useful for K > 1 so that the threads can agree on "who" had the max value. That thread can
|
||||
// then blank out their max with -inf and the warp can run more iterations...
|
||||
#pragma unroll
|
||||
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||
float other_max = __shfl_xor_sync(0xFFFFFFFF, max_val, mask, THREADS_PER_ROW);
|
||||
int other_expert = __shfl_xor_sync(0xFFFFFFFF, expert, mask, THREADS_PER_ROW);
|
||||
|
||||
// We want lower indices to "win" in every thread so we break ties this way
|
||||
if (other_max > max_val || (other_max == max_val && other_expert < expert)) {
|
||||
max_val = other_max;
|
||||
expert = other_expert;
|
||||
}
|
||||
}
|
||||
|
||||
// Write the max for this k iteration to global memory.
|
||||
if (thread_group_idx == 0) {
|
||||
// The lead thread from each sub-group will write out the final results to global memory. (This will be a
|
||||
// single) thread per row of the input/output matrices.
|
||||
const int idx = k * thread_row + k_idx;
|
||||
output[idx] = T(max_val);
|
||||
indices[idx] = should_process_row ? expert : NUM_EXPERTS;
|
||||
source_rows[idx] = k_idx * num_rows + thread_row;
|
||||
}
|
||||
|
||||
// Finally, we clear the value in the thread with the current max if there is another iteration to run.
|
||||
if (k_idx + 1 < k) {
|
||||
const int ldg_group_for_expert = expert / COLS_PER_GROUP_LDG;
|
||||
const int thread_to_clear_in_group = (expert / ELTS_PER_LDG) % THREADS_PER_ROW;
|
||||
|
||||
// Only the thread in the group which produced the max will reset the "winning" value to -inf.
|
||||
if (thread_group_idx == thread_to_clear_in_group) {
|
||||
const int offset_for_expert = expert % ELTS_PER_LDG;
|
||||
// Safe to set to any negative value since row_chunk values must be between 0 and 1.
|
||||
row_chunk[ldg_group_for_expert * ELTS_PER_LDG + offset_for_expert] = ComputeType(-10000.f);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
namespace detail {
|
||||
// Constructs some constants needed to partition the work across threads at compile time.
|
||||
template <typename T, int EXPERTS, int BYTES_PER_LDG>
|
||||
struct TopkConstants {
|
||||
static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T);
|
||||
static_assert(EXPERTS / (ELTS_PER_LDG * WARP_SIZE) == 0 || EXPERTS % (ELTS_PER_LDG * WARP_SIZE) == 0, "");
|
||||
static constexpr int VECs_PER_THREAD = std::max(1, (int)EXPERTS / (ELTS_PER_LDG * WARP_SIZE));
|
||||
static constexpr int VPT = VECs_PER_THREAD * ELTS_PER_LDG;
|
||||
static constexpr int THREADS_PER_ROW = EXPERTS / VPT;
|
||||
static constexpr int ROWS_PER_WARP = WARP_SIZE / THREADS_PER_ROW;
|
||||
};
|
||||
} // namespace detail
|
||||
|
||||
template <typename T, int EXPERTS, int WARPS_PER_TB>
|
||||
void topk_gating_softmax_launcher_helper(const T* input, const bool* finished, T* output, int* indices, int* source_row,
|
||||
int num_rows, int num_experts, int k, cudaStream_t stream) {
|
||||
static constexpr unsigned long MAX_BYTES_PER_LDG = 16;
|
||||
|
||||
static constexpr int BYTES_PER_LDG = std::min((int)MAX_BYTES_PER_LDG, (int)sizeof(T) * EXPERTS);
|
||||
using Constants = detail::TopkConstants<T, EXPERTS, BYTES_PER_LDG>;
|
||||
static constexpr int VPT = Constants::VPT;
|
||||
static constexpr int ROWS_PER_WARP = Constants::ROWS_PER_WARP;
|
||||
const int num_warps = (num_rows + ROWS_PER_WARP - 1) / ROWS_PER_WARP;
|
||||
const int num_blocks = (num_warps + WARPS_PER_TB - 1) / WARPS_PER_TB;
|
||||
|
||||
dim3 block_dim(WARP_SIZE, WARPS_PER_TB);
|
||||
topk_gating_softmax<T, VPT, EXPERTS, WARPS_PER_TB, BYTES_PER_LDG>
|
||||
<<<num_blocks, block_dim, 0, stream>>>(input, finished, output, num_rows, indices, source_row, k);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void topk_gating_softmax_kernelLauncher(const T* input, const bool* finished, T* output, T* softmax_temp_output,
|
||||
int* indices, int* source_row, int num_rows, int num_experts,
|
||||
int k, cudaStream_t stream) {
|
||||
static constexpr int WARPS_PER_TB = 4;
|
||||
|
||||
switch (num_experts) {
|
||||
case 2: {
|
||||
topk_gating_softmax_launcher_helper<T, 2, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 4: {
|
||||
topk_gating_softmax_launcher_helper<T, 4, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 8: {
|
||||
topk_gating_softmax_launcher_helper<T, 8, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 16: {
|
||||
topk_gating_softmax_launcher_helper<T, 16, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 32: {
|
||||
topk_gating_softmax_launcher_helper<T, 32, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 64: {
|
||||
topk_gating_softmax_launcher_helper<T, 64, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 128: {
|
||||
topk_gating_softmax_launcher_helper<T, 128, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
case 256: {
|
||||
topk_gating_softmax_launcher_helper<T, 256, WARPS_PER_TB>(input, finished, output, indices, source_row, num_rows,
|
||||
num_experts, k, stream);
|
||||
break;
|
||||
}
|
||||
default: {
|
||||
static constexpr int TPB = 256;
|
||||
moe_softmax<T, TPB><<<num_rows, TPB, 0, stream>>>(input, finished, softmax_temp_output, num_experts);
|
||||
moe_top_k<T, TPB>
|
||||
<<<num_rows, TPB, 0, stream>>>(softmax_temp_output, finished, output, indices, source_row, num_experts, k);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ========================== CUB Sorting things ====================================
|
||||
CubKeyValueSorter::CubKeyValueSorter() : num_experts_(0), num_bits_(sizeof(int) * 8) {}
|
||||
|
||||
CubKeyValueSorter::CubKeyValueSorter(int num_experts)
|
||||
: num_experts_(num_experts), num_bits_((int)log2(num_experts) + 1) {}
|
||||
|
||||
void CubKeyValueSorter::update_num_experts(int num_experts) {
|
||||
num_experts_ = num_experts;
|
||||
num_bits_ = (int)log2(num_experts) + 1;
|
||||
}
|
||||
|
||||
size_t CubKeyValueSorter::getWorkspaceSize(const size_t num_key_value_pairs) {
|
||||
num_key_value_pairs_ = num_key_value_pairs;
|
||||
size_t required_storage = 0;
|
||||
int* null_int = nullptr;
|
||||
cub::DeviceRadixSort::SortPairs(NULL, required_storage, null_int, null_int, null_int, null_int,
|
||||
(int)num_key_value_pairs, 0, num_bits_);
|
||||
return required_storage;
|
||||
}
|
||||
|
||||
void CubKeyValueSorter::run(void* workspace, const size_t workspace_size, const int* keys_in, int* keys_out,
|
||||
const int* values_in, int* values_out, const size_t num_key_value_pairs,
|
||||
cudaStream_t stream) {
|
||||
size_t expected_ws_size = getWorkspaceSize(num_key_value_pairs);
|
||||
size_t actual_ws_size = workspace_size;
|
||||
|
||||
if (expected_ws_size > workspace_size) {
|
||||
ORT_THROW("Error. The allocated workspace is too small to run this problem. Expected workspace size of at least ",
|
||||
expected_ws_size, " but got problem size ", workspace_size, "\n");
|
||||
}
|
||||
cub::DeviceRadixSort::SortPairs(workspace, actual_ws_size, keys_in, keys_out, values_in, values_out,
|
||||
(int)num_key_value_pairs, 0, num_bits_, stream);
|
||||
}
|
||||
|
||||
// ============================== Infer GEMM sizes =================================
|
||||
__device__ inline int find_total_elts_leq_target(const int* sorted_indices, const int arr_length, const int target) {
|
||||
int64_t low = 0, high = arr_length - 1, target_location = -1;
|
||||
while (low <= high) {
|
||||
int64_t mid = (low + high) / 2;
|
||||
|
||||
if (sorted_indices[mid] > target) {
|
||||
high = mid - 1;
|
||||
} else {
|
||||
low = mid + 1;
|
||||
target_location = mid;
|
||||
}
|
||||
}
|
||||
return target_location + 1;
|
||||
}
|
||||
|
||||
// Sets up the gemm assuming the inputs, experts and outputs are stored in row major order.
|
||||
// Assumes we want to perform output = matmul(inputs, experts) + bias
|
||||
__global__ void compute_total_rows_before_expert_kernel(const int* sorted_experts, const int sorted_experts_len,
|
||||
const int64_t num_experts, int64_t* total_rows_before_expert) {
|
||||
// First, compute the global tid. We only need 1 thread per expert.
|
||||
const int expert = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
if (expert >= num_experts) return;
|
||||
|
||||
// This should construct the last index where each expert occurs.
|
||||
total_rows_before_expert[expert] = find_total_elts_leq_target(sorted_experts, sorted_experts_len, expert);
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename Enable>
|
||||
CutlassMoeFCRunner<T, WeightType, Enable>::CutlassMoeFCRunner(int sm_version) {
|
||||
moe_gemm_runner_.initialize(sm_version);
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename Enable>
|
||||
size_t CutlassMoeFCRunner<T, WeightType, Enable>::getWorkspaceSize(int num_rows, const int hidden_size,
|
||||
const int inter_size, int num_experts,
|
||||
int k) {
|
||||
const int buf_size = static_cast<int>(pad_to_multiple_of_16(k * num_rows * hidden_size));
|
||||
const int interbuf_size = static_cast<int>(pad_to_multiple_of_16(k * num_rows * inter_size));
|
||||
const int padded_experts = static_cast<int>(pad_to_multiple_of_16(num_experts));
|
||||
const int num_moe_inputs = static_cast<int>(pad_to_multiple_of_16(k * num_rows));
|
||||
int num_softmax_outs = 0;
|
||||
|
||||
const bool is_pow_2 = (num_experts != 0) && ((num_experts & (num_experts - 1)) == 0);
|
||||
if (!is_pow_2 || num_experts > 256) {
|
||||
num_softmax_outs = static_cast<int>(pad_to_multiple_of_16(num_rows * num_experts));
|
||||
}
|
||||
|
||||
// softmax output, permuted_rows and permuted_experts have moved to outside of moe kernel, allocate them
|
||||
// in Encoder or Decoder before invoking FfnLayer forward.
|
||||
size_t total_ws_bytes = 3 * num_moe_inputs * sizeof(int); // source_rows_, permuted_rows_, permuted_experts_
|
||||
total_ws_bytes += buf_size * sizeof(T); // permuted_data
|
||||
total_ws_bytes += padded_experts * sizeof(int64_t); // Hold total_rows_before_expert_
|
||||
total_ws_bytes += num_softmax_outs * sizeof(T);
|
||||
const int bytes_for_fc1_result = interbuf_size * sizeof(T);
|
||||
const int sorter_ws_size_bytes = static_cast<int>(pad_to_multiple_of_16(sorter_.getWorkspaceSize(num_rows)));
|
||||
sorter_.update_num_experts(num_experts);
|
||||
|
||||
int bytes_for_intermediate_and_sorting = bytes_for_fc1_result;
|
||||
if (sorter_ws_size_bytes > bytes_for_fc1_result) {
|
||||
int remaining_bytes = static_cast<int>(pad_to_multiple_of_16(sorter_ws_size_bytes - bytes_for_fc1_result));
|
||||
bytes_for_intermediate_and_sorting += remaining_bytes;
|
||||
}
|
||||
|
||||
total_ws_bytes += bytes_for_intermediate_and_sorting; // intermediate (fc1) output + cub sorting workspace
|
||||
return total_ws_bytes;
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename Enable>
|
||||
void CutlassMoeFCRunner<T, WeightType, Enable>::configure_ws_ptrs(char* ws_ptr, int num_rows,
|
||||
const int hidden_size, const int inter_size,
|
||||
int num_experts, int k) {
|
||||
const int buf_size = static_cast<int>(pad_to_multiple_of_16(k * num_rows * hidden_size));
|
||||
const int interbuf_size = static_cast<int>(pad_to_multiple_of_16(k * num_rows * inter_size));
|
||||
const int padded_experts = static_cast<int>(pad_to_multiple_of_16(num_experts));
|
||||
const int num_moe_inputs = static_cast<int>(pad_to_multiple_of_16(k * num_rows));
|
||||
// const int num_softmax_outs = pad_to_multiple_of_16(num_rows * num_experts);
|
||||
|
||||
source_rows_ = (int*)ws_ptr;
|
||||
permuted_rows_ = source_rows_ + num_moe_inputs;
|
||||
permuted_experts_ = permuted_rows_ + num_moe_inputs;
|
||||
permuted_data_ = (T*)(permuted_experts_ + num_moe_inputs);
|
||||
|
||||
total_rows_before_expert_ = (int64_t*)(permuted_data_ + buf_size);
|
||||
|
||||
fc1_result_ = (T*)(total_rows_before_expert_ + padded_experts);
|
||||
|
||||
const bool is_pow_2 = (num_experts != 0) && ((num_experts & (num_experts - 1)) == 0);
|
||||
if (!is_pow_2 || num_experts > 256) {
|
||||
softmax_out_ = (T*)(fc1_result_ + interbuf_size);
|
||||
} else {
|
||||
softmax_out_ = nullptr;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename Enable>
|
||||
void CutlassMoeFCRunner<T, WeightType, Enable>::run_moe_fc(
|
||||
const T* input_activations, const T* gating_output, const WeightType* fc1_expert_weights, const T* fc1_scales,
|
||||
const T* fc1_expert_biases, ActivationType fc1_activation_type, const WeightType* fc2_expert_weights,
|
||||
const T* fc2_scales, int num_rows, const int hidden_size, const int inter_size, int num_experts,
|
||||
int k, char* workspace_ptr, T* fc2_result, const bool* finished, int active_rows, T* expert_scales,
|
||||
int* expanded_source_row_to_expanded_dest_row, int* expert_for_source_row, cudaStream_t stream) {
|
||||
static constexpr bool scales_required =
|
||||
std::is_same<WeightType, uint8_t>::value || std::is_same<WeightType, cutlass::uint4b_t>::value;
|
||||
|
||||
if (scales_required) {
|
||||
if (fc1_scales == nullptr) {
|
||||
ORT_THROW("[FT Error][Run MoE FC] Scales expected but scale for first matmul is a null pointer");
|
||||
} else if (fc2_scales == nullptr) {
|
||||
ORT_THROW("[FT Error][Run MoE FC] Scales expected but scale for second matmul is a null pointer");
|
||||
}
|
||||
} else {
|
||||
if (fc1_scales != nullptr) {
|
||||
ORT_THROW("[FT Error][Run MoE FC] Scales are ignored for fp32/fp16/bf16 but received scale for FC1");
|
||||
} else if (fc2_scales != nullptr) {
|
||||
ORT_THROW("[FT Error][Run MoE FC] Scales are ignored for fp32/fp16/bf16 but received scale for FC2");
|
||||
}
|
||||
}
|
||||
|
||||
configure_ws_ptrs(workspace_ptr, num_rows, hidden_size, inter_size, num_experts, k);
|
||||
topk_gating_softmax_kernelLauncher<T>(gating_output, finished, expert_scales, softmax_out_, expert_for_source_row,
|
||||
source_rows_, num_rows, num_experts, k, stream);
|
||||
|
||||
const int sorter_ws_size_bytes = static_cast<int>(pad_to_multiple_of_16(sorter_.getWorkspaceSize(k * num_rows)));
|
||||
sorter_.run((void*)fc1_result_, sorter_ws_size_bytes, expert_for_source_row, permuted_experts_, source_rows_,
|
||||
permuted_rows_, k * num_rows, stream);
|
||||
|
||||
initialize_moe_routing_kernelLauncher(input_activations, permuted_data_, permuted_rows_,
|
||||
expanded_source_row_to_expanded_dest_row, num_rows, active_rows, hidden_size, k,
|
||||
stream);
|
||||
|
||||
const int expanded_active_expert_rows = k * active_rows;
|
||||
compute_total_rows_before_expert(permuted_experts_, expanded_active_expert_rows, num_experts,
|
||||
total_rows_before_expert_, stream);
|
||||
|
||||
moe_gemm_runner_.moe_gemm_bias_act(permuted_data_, fc1_expert_weights, fc1_scales, fc1_expert_biases, fc1_result_,
|
||||
total_rows_before_expert_, expanded_active_expert_rows, inter_size, hidden_size,
|
||||
num_experts, fc1_activation_type, stream);
|
||||
|
||||
moe_gemm_runner_.moe_gemm(fc1_result_, fc2_expert_weights, fc2_scales, fc2_result, total_rows_before_expert_,
|
||||
expanded_active_expert_rows, hidden_size, inter_size, num_experts, stream);
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename Enable>
|
||||
void CutlassMoeFCRunner<T, WeightType, Enable>::run_moe_fc(
|
||||
const T* input_activations, const T* gating_output, const WeightType* fc1_expert_weights, const T* fc1_scales,
|
||||
const T* fc1_expert_biases, ActivationType fc1_activation_type, const WeightType* fc2_expert_weights,
|
||||
const T* fc2_scales, int num_rows, const int hidden_size, const int inter_size, int num_experts,
|
||||
int k, char* workspace_ptr, T* fc2_result, T* expert_scales, int* expanded_source_row_to_expanded_dest_row,
|
||||
int* expert_for_source_row, cudaStream_t stream) {
|
||||
run_moe_fc(input_activations, gating_output, fc1_expert_weights, fc1_scales, fc1_expert_biases, fc1_activation_type,
|
||||
fc2_expert_weights, fc2_scales, num_rows, hidden_size, inter_size, num_experts, k, workspace_ptr,
|
||||
fc2_result, nullptr, num_rows, expert_scales, expanded_source_row_to_expanded_dest_row,
|
||||
expert_for_source_row, stream);
|
||||
}
|
||||
|
||||
template <typename T, typename WeightType, typename Enable>
|
||||
void CutlassMoeFCRunner<T, WeightType, Enable>::compute_total_rows_before_expert(const int* sorted_indices,
|
||||
const int total_indices,
|
||||
int num_experts,
|
||||
int64_t* total_rows_before_expert,
|
||||
cudaStream_t stream) {
|
||||
const int threads = std::min(1024, num_experts);
|
||||
const int blocks = (num_experts + threads - 1) / threads;
|
||||
|
||||
compute_total_rows_before_expert_kernel<<<blocks, threads, 0, stream>>>(sorted_indices, total_indices, num_experts,
|
||||
total_rows_before_expert);
|
||||
}
|
||||
|
||||
// ========================== Permutation things =======================================
|
||||
|
||||
// Duplicated and permutes rows for MoE. In addition, reverse the permutation map to help with finalizing routing.
|
||||
|
||||
// "expanded_x_row" simply means that the number of values is num_rows x k. It is "expanded" since we will have to
|
||||
// duplicate some rows in the input matrix to match the dimensions. Duplicates will always get routed to separate
|
||||
// experts in the end.
|
||||
|
||||
// Note that the expanded_dest_row_to_expanded_source_row map referred to here has indices in the range (0,
|
||||
// k*rows_in_input - 1). However, it is set up so that index 0, rows_in_input, 2*rows_in_input ... (k-1)*rows_in_input
|
||||
// all map to row 0 in the original matrix. Thus, to know where to read in the source matrix, we simply take the modulus
|
||||
// of the expanded index.
|
||||
|
||||
template <typename T>
|
||||
__global__ void initialize_moe_routing_kernel(const T* unpermuted_input, T* permuted_output,
|
||||
const int* expanded_dest_row_to_expanded_source_row,
|
||||
int* expanded_source_row_to_expanded_dest_row, int num_rows,
|
||||
int active_rows, int cols) {
|
||||
// Reverse permutation map.
|
||||
// I do this so that later, we can use the source -> dest map to do the k-way reduction and unpermuting. I need the
|
||||
// reverse map for that reduction to allow each threadblock to do 1 k-way reduce without atomics later in MoE. 1
|
||||
// thread block will be responsible for all k summations.
|
||||
const int expanded_dest_row = blockIdx.x;
|
||||
const int expanded_source_row = expanded_dest_row_to_expanded_source_row[expanded_dest_row];
|
||||
if (threadIdx.x == 0) {
|
||||
expanded_source_row_to_expanded_dest_row[expanded_source_row] = expanded_dest_row;
|
||||
}
|
||||
|
||||
if (blockIdx.x < active_rows) {
|
||||
// Duplicate and permute rows
|
||||
const int source_row = expanded_source_row % num_rows;
|
||||
|
||||
const T* source_row_ptr = unpermuted_input + source_row * cols;
|
||||
T* dest_row_ptr = permuted_output + expanded_dest_row * cols;
|
||||
|
||||
for (int tid = threadIdx.x; tid < cols; tid += blockDim.x) {
|
||||
dest_row_ptr[tid] = source_row_ptr[tid];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void initialize_moe_routing_kernelLauncher(const T* unpermuted_input, T* permuted_output,
|
||||
const int* expanded_dest_row_to_expanded_source_row,
|
||||
int* expanded_source_row_to_expanded_dest_row, int num_rows,
|
||||
int active_rows, int cols, int k, cudaStream_t stream) {
|
||||
const int blocks = num_rows * k;
|
||||
const int threads = std::min(cols, 1024);
|
||||
initialize_moe_routing_kernel<T>
|
||||
<<<blocks, threads, 0, stream>>>(unpermuted_input, permuted_output, expanded_dest_row_to_expanded_source_row,
|
||||
expanded_source_row_to_expanded_dest_row, num_rows, k * active_rows, cols);
|
||||
}
|
||||
|
||||
// Final kernel to unpermute and scale
|
||||
// This kernel unpermutes the original data, does the k-way reduction and performs the final skip connection.
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ < 530
|
||||
template <typename T, int RESIDUAL_NUM>
|
||||
__global__ void finalize_moe_routing_kernel(const T*, T*, const T*, const T*, const T*, const T*, const int*,
|
||||
const int*, int, const int) {
|
||||
// Does not support pre-Kepler architectures
|
||||
;
|
||||
}
|
||||
#else
|
||||
template <typename T, int RESIDUAL_NUM>
|
||||
__global__ void finalize_moe_routing_kernel(const T* expanded_permuted_rows, T* reduced_unpermuted_output,
|
||||
const T* skip_1, const T* skip_2, const T* bias, const T* scales,
|
||||
const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int cols, int k) {
|
||||
const int original_row = blockIdx.x;
|
||||
int num_rows = gridDim.x;
|
||||
T* reduced_row_ptr = reduced_unpermuted_output + original_row * cols;
|
||||
|
||||
const T* skip_1_row_ptr = nullptr;
|
||||
if (RESIDUAL_NUM == 1) {
|
||||
skip_1_row_ptr = skip_1 + original_row * cols;
|
||||
}
|
||||
const T* skip_2_row_ptr = nullptr;
|
||||
if (RESIDUAL_NUM == 2) {
|
||||
skip_2_row_ptr = skip_2 + original_row * cols;
|
||||
}
|
||||
|
||||
for (int tid = threadIdx.x; tid < cols; tid += blockDim.x) {
|
||||
T thread_output;
|
||||
if (RESIDUAL_NUM == 0) {
|
||||
thread_output = T(0);
|
||||
} else if (RESIDUAL_NUM == 1) {
|
||||
thread_output = skip_1_row_ptr[tid];
|
||||
} else if (RESIDUAL_NUM == 2) {
|
||||
thread_output = skip_1_row_ptr[tid] + skip_2_row_ptr[tid];
|
||||
}
|
||||
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||
const int expanded_original_row = original_row + k_idx * num_rows;
|
||||
const int expanded_permuted_row = expanded_source_row_to_expanded_dest_row[expanded_original_row];
|
||||
|
||||
const int64_t k_offset = original_row * k + k_idx;
|
||||
const T row_scale = scales[k_offset];
|
||||
const T* expanded_permuted_rows_row_ptr = expanded_permuted_rows + expanded_permuted_row * cols;
|
||||
|
||||
const int expert_idx = expert_for_source_row[k_offset];
|
||||
const T* bias_ptr = bias + expert_idx * cols;
|
||||
|
||||
thread_output = thread_output + row_scale * (expanded_permuted_rows_row_ptr[tid] + bias_ptr[tid]);
|
||||
}
|
||||
reduced_row_ptr[tid] = thread_output;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
|
||||
template <typename T>
|
||||
void finalize_moe_routing_kernelLauncher(const T* expanded_permuted_rows, T* reduced_unpermuted_output, const T* bias,
|
||||
const T* scales, const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int num_rows, int cols,
|
||||
int k, cudaStream_t stream) {
|
||||
const int blocks = num_rows;
|
||||
const int threads = std::min(cols, 1024);
|
||||
finalize_moe_routing_kernel<T, 0><<<blocks, threads, 0, stream>>>(
|
||||
expanded_permuted_rows, reduced_unpermuted_output, nullptr, nullptr, bias, scales,
|
||||
expanded_source_row_to_expanded_dest_row, expert_for_source_row, cols, k);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void finalize_moe_routing_kernelLauncher(const T* expanded_permuted_rows, T* reduced_unpermuted_output, const T* skip,
|
||||
const T* bias, const T* scales,
|
||||
const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int num_rows, int cols,
|
||||
int k, cudaStream_t stream) {
|
||||
const int blocks = num_rows;
|
||||
const int threads = std::min(cols, 1024);
|
||||
finalize_moe_routing_kernel<T, 1>
|
||||
<<<blocks, threads, 0, stream>>>(expanded_permuted_rows, reduced_unpermuted_output, skip, nullptr, bias, scales,
|
||||
expanded_source_row_to_expanded_dest_row, expert_for_source_row, cols, k);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void finalize_moe_routing_kernelLauncher(const T* expanded_permuted_rows, T* reduced_unpermuted_output, const T* skip_1,
|
||||
const T* skip_2, const T* bias, const T* scales,
|
||||
const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int num_rows, int cols,
|
||||
int k, cudaStream_t stream) {
|
||||
const int blocks = num_rows;
|
||||
const int threads = std::min(cols, 1024);
|
||||
if (skip_2 == nullptr) {
|
||||
finalize_moe_routing_kernel<T, 1><<<blocks, threads, 0, stream>>>(
|
||||
expanded_permuted_rows, reduced_unpermuted_output, skip_1, skip_2, bias, scales,
|
||||
expanded_source_row_to_expanded_dest_row, expert_for_source_row, cols, k);
|
||||
} else {
|
||||
finalize_moe_routing_kernel<T, 2><<<blocks, threads, 0, stream>>>(
|
||||
expanded_permuted_rows, reduced_unpermuted_output, skip_1, skip_2, bias, scales,
|
||||
expanded_source_row_to_expanded_dest_row, expert_for_source_row, cols, k);
|
||||
}
|
||||
}
|
||||
|
||||
// ========================= TopK Softmax specializations ===========================
|
||||
template void topk_gating_softmax_kernelLauncher(const float*, const bool*, float*, float*, int*, int*, int,
|
||||
int, int, cudaStream_t);
|
||||
template void topk_gating_softmax_kernelLauncher(const half*, const bool*, half*, half*, int*, int*, int,
|
||||
int, int, cudaStream_t);
|
||||
|
||||
// ==================== Variable batched GEMM specializations ==================================
|
||||
template class CutlassMoeFCRunner<float, float>;
|
||||
template class CutlassMoeFCRunner<half, half>;
|
||||
|
||||
// ===================== Specializations for init routing =========================
|
||||
template void initialize_moe_routing_kernelLauncher(const float*, float*, const int*, int*, int, int,
|
||||
int, int, cudaStream_t);
|
||||
template void initialize_moe_routing_kernelLauncher(const half*, half*, const int*, int*, int, int,
|
||||
int, int, cudaStream_t);
|
||||
|
||||
// ==================== Specializations for final routing ===================================
|
||||
template void finalize_moe_routing_kernelLauncher(const float*, float*, const float*, const float*, const int*,
|
||||
const int*, int, int, int, cudaStream_t);
|
||||
template void finalize_moe_routing_kernelLauncher(const half*, half*, const half*, const half*, const int*, const int*,
|
||||
int, int, int, cudaStream_t);
|
||||
template void finalize_moe_routing_kernelLauncher(const float*, float*, const float*, const float*, const float*,
|
||||
const int*, const int*, int, int, int,
|
||||
cudaStream_t);
|
||||
template void finalize_moe_routing_kernelLauncher(const half*, half*, const half*, const half*, const half*, const int*,
|
||||
const int*, int, int, int, cudaStream_t);
|
||||
template void finalize_moe_routing_kernelLauncher(const float*, float*, const float*, const float*, const float*,
|
||||
const float*, const int*, const int*, int, int, int,
|
||||
cudaStream_t);
|
||||
template void finalize_moe_routing_kernelLauncher(const half*, half*, const half*, const half*, const half*,
|
||||
const half*, const int*, const int*, int, int, int,
|
||||
cudaStream_t);
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
158
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_kernel.h
Normal file
158
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_kernel.h
Normal file
|
|
@ -0,0 +1,158 @@
|
|||
/*
|
||||
* Copyright (c) 2020-2023, NVIDIA CORPORATION. All rights reserved.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "moe_gemm_kernels.h"
|
||||
#include <cuda_runtime_api.h>
|
||||
|
||||
#include "core/common/common.h"
|
||||
|
||||
using namespace onnxruntime;
|
||||
|
||||
namespace ort_fastertransformer {
|
||||
|
||||
static inline size_t pad_to_multiple_of_16(size_t input) {
|
||||
static constexpr int ALIGNMENT = 16;
|
||||
return ALIGNMENT * ((input + ALIGNMENT - 1) / ALIGNMENT);
|
||||
}
|
||||
|
||||
/*
|
||||
Launches the topk gating softmax required for the MoE layers.
|
||||
|
||||
Params:
|
||||
input - a [num_rows x num_experts]
|
||||
finished - [num_rows] vector with 1 if the sentence at this row is done translating and 0 otherwise.
|
||||
output - a buffer of shape [num_rows x k] containing the top-k values of the softmax for each row.
|
||||
indices - a matrix of shape [num_rows x k] containing the top-k experts each row should get routed to.
|
||||
source_rows - a matrix of shape [num_rows x k] used internally for permuting. source_rows[row][k] = k * num_rows +
|
||||
row. It is constructed like this so we can track where each of the original rows end up in order to perform the
|
||||
"k-way" reduction later in the routing.
|
||||
|
||||
num_rows - The number of rows in the matrix
|
||||
num_experts - The number of expert layers present
|
||||
k - k value in topk
|
||||
*/
|
||||
template <typename T>
|
||||
void topk_gating_softmax_kernelLauncher(const T* input, const bool* finished, T* output, T* softmax_temp_out,
|
||||
int* indices, int* source_row, int num_rows, int num_experts,
|
||||
int k, cudaStream_t stream);
|
||||
|
||||
class CubKeyValueSorter {
|
||||
public:
|
||||
CubKeyValueSorter();
|
||||
|
||||
CubKeyValueSorter(int num_experts);
|
||||
|
||||
void update_num_experts(int num_experts);
|
||||
|
||||
size_t getWorkspaceSize(const size_t num_key_value_pairs);
|
||||
|
||||
void run(void* workspace, const size_t workspace_size, const int* keys_in, int* keys_out, const int* values_in,
|
||||
int* values_out, const size_t num_key_value_pairs, cudaStream_t stream);
|
||||
|
||||
private:
|
||||
size_t num_key_value_pairs_;
|
||||
int num_experts_;
|
||||
int num_bits_;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
void initialize_moe_routing_kernelLauncher(const T* unpermuted_input, T* permuted_output,
|
||||
const int* expanded_dest_row_to_expanded_source_row,
|
||||
int* expanded_source_row_to_expanded_dest_row, int num_rows,
|
||||
int active_rows, int cols, int k, cudaStream_t stream);
|
||||
|
||||
template <typename T>
|
||||
void finalize_moe_routing_kernelLauncher(const T* expanded_permuted_rows, T* reduced_unpermuted_output, const T* bias,
|
||||
const T* scales, const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int num_rows, int cols,
|
||||
int k, cudaStream_t stream);
|
||||
|
||||
template <typename T>
|
||||
void finalize_moe_routing_kernelLauncher(const T* expanded_permuted_rows, T* reduced_unpermuted_output, const T* skip,
|
||||
const T* bias, const T* scales,
|
||||
const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int num_rows, int cols,
|
||||
int k, cudaStream_t stream);
|
||||
|
||||
template <typename T>
|
||||
void finalize_moe_routing_kernelLauncher(const T* expanded_permuted_rows, T* reduced_unpermuted_output, const T* skip_1,
|
||||
const T* skip_2, const T* bias, const T* scales,
|
||||
const int* expanded_source_row_to_expanded_dest_row,
|
||||
const int* expert_for_source_row, int num_rows, int cols,
|
||||
int k, cudaStream_t stream);
|
||||
|
||||
// Assumes inputs activations are row major. Weights need to be preprocessed by th_op/weight_quantize.cc .
|
||||
// Nested in a class to avoid multiple calls to cudaGetDeviceProperties as this call can be expensive.
|
||||
// Avoid making several duplicates of this class.
|
||||
template <typename T, /*The type used for activations/scales/compute*/
|
||||
typename WeightType, /* The type for the MoE weights */
|
||||
typename Enable = void>
|
||||
class CutlassMoeFCRunner {
|
||||
public:
|
||||
CutlassMoeFCRunner(int sm_version);
|
||||
|
||||
size_t getWorkspaceSize(int num_rows, int hidden_size, int inter_size, int num_experts, int k);
|
||||
|
||||
void run_moe_fc(const T* input_activations, const T* gating_output, const WeightType* fc1_expert_weights,
|
||||
const T* fc1_scales, const T* fc1_expert_biases, ActivationType fc1_activation_type,
|
||||
const WeightType* fc2_expert_weights, const T* fc2_scales, int num_rows, int hidden_size,
|
||||
int inter_size, int num_experts, int k, char* workspace_ptr, T* fc2_result,
|
||||
T* expert_scales, int* expanded_source_row_to_expanded_dest_row, int* expert_for_source_row,
|
||||
cudaStream_t stream);
|
||||
|
||||
void run_moe_fc(const T* input_activations, const T* gating_output, const WeightType* fc1_expert_weights,
|
||||
const T* fc1_scales, const T* fc1_expert_biases, ActivationType fc1_activation_type,
|
||||
const WeightType* fc2_expert_weights, const T* fc2_scales, int num_rows, int hidden_size,
|
||||
int inter_size, int num_experts, int k, char* workspace_ptr, T* fc2_result,
|
||||
const bool* finished, int active_rows, T* expert_scales,
|
||||
int* expanded_source_row_to_expanded_dest_row, int* expert_for_source_row, cudaStream_t stream);
|
||||
|
||||
void compute_total_rows_before_expert(const int* sorted_indices, int total_indices, int num_experts,
|
||||
int64_t* total_rows_before_expert, cudaStream_t stream);
|
||||
|
||||
private:
|
||||
void configure_ws_ptrs(char* ws_ptr, int num_rows, int hidden_size, int inter_size, int num_experts, int k);
|
||||
|
||||
private:
|
||||
CubKeyValueSorter sorter_;
|
||||
MoeGemmRunner<T, WeightType> moe_gemm_runner_;
|
||||
|
||||
// Pointers
|
||||
int* source_rows_;
|
||||
int* permuted_rows_;
|
||||
int* permuted_experts_;
|
||||
char* sorter_ws_;
|
||||
T* permuted_data_;
|
||||
T* softmax_out_;
|
||||
|
||||
int64_t* total_rows_before_expert_;
|
||||
|
||||
T* fc1_result_;
|
||||
};
|
||||
|
||||
template <typename WeightType>
|
||||
class CutlassMoeFCRunner<float, WeightType, typename std::enable_if_t<!std::is_same<float, WeightType>::value>> {
|
||||
public:
|
||||
CutlassMoeFCRunner(int sm_version);
|
||||
|
||||
size_t getWorkspaceSize(int num_rows, int hidden_size, int inter_size, int num_experts, int k) {
|
||||
return 0;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace ort_fastertransformer
|
||||
290
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_problem_visitor.h
Normal file
290
onnxruntime/contrib_ops/cuda/moe/ft_moe/moe_problem_visitor.h
Normal file
|
|
@ -0,0 +1,290 @@
|
|||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Base scheduler for grouped problems, using MoE
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/gemm/kernel/grouped_problem_visitor.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace gemm {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Visitor class to abstract away the algorithm for iterating over tiles
|
||||
template <typename ProblemSizeHelper, typename ThreadblockShape_>
|
||||
struct BaseMoeProblemVisitor {
|
||||
using ThreadblockShape = ThreadblockShape_;
|
||||
|
||||
struct ProblemInfo {
|
||||
static int32_t const kNoPrefetchEntry = -1;
|
||||
int32_t problem_idx;
|
||||
int32_t problem_start;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ProblemInfo() : problem_idx(kNoPrefetchEntry), problem_start(kNoPrefetchEntry) {}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
ProblemInfo(int32_t problem_idx_, int32_t problem_start_)
|
||||
: problem_idx(problem_idx_), problem_start(problem_start_) {}
|
||||
};
|
||||
|
||||
struct Params {
|
||||
int64_t const* last_row_for_problem;
|
||||
int64_t gemm_n;
|
||||
int64_t gemm_k;
|
||||
int32_t problem_count;
|
||||
void const* workspace;
|
||||
int32_t tile_count;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params()
|
||||
: last_row_for_problem(nullptr), gemm_n(0), gemm_k(0), problem_count(0), workspace(nullptr), tile_count(0) {}
|
||||
|
||||
/// Ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(int64_t const* last_row_for_problem, int64_t gemm_n, int64_t gemm_k, int32_t problem_count,
|
||||
void const* workspace = nullptr, int32_t tile_count = 0)
|
||||
: last_row_for_problem(last_row_for_problem),
|
||||
gemm_n(gemm_n),
|
||||
gemm_k(gemm_k),
|
||||
problem_count(problem_count),
|
||||
workspace(workspace),
|
||||
tile_count(tile_count) {}
|
||||
};
|
||||
|
||||
Params const& params;
|
||||
int32_t tile_idx;
|
||||
int32_t problem_tile_start;
|
||||
int32_t problem_idx;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
BaseMoeProblemVisitor(Params const& params_, int32_t block_idx)
|
||||
: params(params_), tile_idx(block_idx), problem_tile_start(0), problem_idx(0) {}
|
||||
|
||||
/// Get the grid shape
|
||||
CUTLASS_HOST_DEVICE
|
||||
static cutlass::gemm::GemmCoord grid_shape(const cutlass::gemm::GemmCoord& problem) {
|
||||
return cutlass::gemm::GemmCoord(((problem.m() - 1 + ThreadblockShape::kM) / ThreadblockShape::kM),
|
||||
((problem.n() - 1 + ThreadblockShape::kN) / ThreadblockShape::kN), 1);
|
||||
}
|
||||
|
||||
/// Gets the global tile index
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t tile_index() const { return tile_idx; }
|
||||
|
||||
/// Gets the index of the problem
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t problem_index() const { return problem_idx; }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t threadblock_idx() const { return tile_idx - problem_tile_start; }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void advance(int32_t grid_size) { tile_idx += grid_size; }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static void possibly_transpose_problem(cutlass::gemm::GemmCoord& problem) {
|
||||
ProblemSizeHelper::possibly_transpose_problem(problem);
|
||||
}
|
||||
|
||||
/// Returns the problem size for the current problem
|
||||
CUTLASS_HOST_DEVICE
|
||||
cutlass::gemm::GemmCoord problem_size() const { return problem_size(problem_idx); }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
cutlass::gemm::GemmCoord problem_size(int idx) const {
|
||||
const int64_t prev_problem_row = idx == 0 ? 0 : params.last_row_for_problem[idx - 1];
|
||||
const int64_t current_problem_row = params.last_row_for_problem[idx];
|
||||
const int64_t gemm_m = current_problem_row - prev_problem_row;
|
||||
GemmCoord problem(GemmCoord::Index(gemm_m), GemmCoord::Index(params.gemm_n), GemmCoord::Index(params.gemm_k));
|
||||
ProblemSizeHelper::possibly_transpose_problem(problem);
|
||||
return problem;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t tile_count(const cutlass::gemm::GemmCoord& grid) { return ProblemSizeHelper::tile_count(grid); }
|
||||
|
||||
static int32_t group_tile_count(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr, int32_t problem_count) {
|
||||
int32_t total_tiles = 0;
|
||||
for (int32_t i = 0; i < problem_count; ++i) {
|
||||
auto problem = host_problem_sizes_ptr[i];
|
||||
possibly_transpose_problem(problem);
|
||||
auto grid = grid_shape(problem);
|
||||
total_tiles += tile_count(grid);
|
||||
}
|
||||
|
||||
return total_tiles;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename ProblemSizeHelper, typename ThreadblockShape, GroupScheduleMode GroupScheduleMode_,
|
||||
int PrefetchTileCount, int ThreadCount>
|
||||
struct MoeProblemVisitor;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// ProblemVisitor that performs all scheduling on device
|
||||
//
|
||||
template <typename ProblemSizeHelper, typename ThreadblockShape, int PrefetchTileCount, int ThreadCount>
|
||||
struct MoeProblemVisitor<ProblemSizeHelper, ThreadblockShape, GroupScheduleMode::kDeviceOnly, PrefetchTileCount,
|
||||
ThreadCount> : public BaseMoeProblemVisitor<ProblemSizeHelper, ThreadblockShape> {
|
||||
using Base = BaseMoeProblemVisitor<ProblemSizeHelper, ThreadblockShape>;
|
||||
using Params = typename Base::Params;
|
||||
static int const kThreadCount = ThreadCount;
|
||||
static bool const kRequiresPrecomputation = false;
|
||||
static int const kThreadsPerWarp = 32;
|
||||
|
||||
struct SharedStorage {};
|
||||
|
||||
// Final tile of the problem loaded by this thread. Each thread will hold
|
||||
// a separate value.
|
||||
int32_t problem_ending_tile;
|
||||
|
||||
SharedStorage& shared_storage;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
MoeProblemVisitor(Params const& params_, SharedStorage& shared_storage_, int32_t block_idx)
|
||||
: Base(params_, block_idx), problem_ending_tile(0), shared_storage(shared_storage_) {
|
||||
this->problem_idx = -1 * kThreadsPerWarp;
|
||||
this->problem_tile_start = 0;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool next_tile() {
|
||||
// Check whether the tile to compute is within the range of the current problem.
|
||||
int32_t problem_tile_end = __shfl_sync(0xffffffff, problem_ending_tile, this->problem_idx % kThreadsPerWarp);
|
||||
if (this->tile_idx < problem_tile_end) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Check whether the tile to compute is within the current group of problems fetched by the warp.
|
||||
// The last tile for this group is the final tile of the problem held by the final thread in the warp.
|
||||
int32_t group_tile_end = __shfl_sync(0xffffffff, problem_ending_tile, kThreadsPerWarp - 1);
|
||||
|
||||
// Keep the starting problem for this group in `problem_idx`. This is done to reduce
|
||||
// register pressure. The starting problem for this group is simply the first problem
|
||||
// in the group most recently fetched by the warp.
|
||||
int32_t& group_problem_start = this->problem_idx;
|
||||
group_problem_start = (this->problem_idx / kThreadsPerWarp) * kThreadsPerWarp;
|
||||
|
||||
// Keep the starting tile for this group in `problem_tile_start`. This is done to reduce
|
||||
// register pressure.
|
||||
int32_t& group_tile_start = this->problem_tile_start;
|
||||
|
||||
// Each thread in the warp processes a separate problem to advance until
|
||||
// reaching a problem whose starting tile is less less than tile_idx.
|
||||
while (group_tile_end <= this->tile_idx) {
|
||||
group_problem_start += kThreadsPerWarp;
|
||||
if (group_problem_start > this->params.problem_count) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Since `group_tile_start` is a reference to `this->problem_tile_start`, this
|
||||
// also sets `this->problem_tile_start`. The fact that `this->problem_tile_start`
|
||||
// is also set here is used later in `next_tile`.
|
||||
group_tile_start = group_tile_end;
|
||||
|
||||
int lane_idx = threadIdx.x % kThreadsPerWarp;
|
||||
int32_t lane_problem = group_problem_start + lane_idx;
|
||||
|
||||
// Compute the number of tiles in the problem assigned to each thread.
|
||||
problem_ending_tile = 0;
|
||||
if (lane_problem < this->params.problem_count) {
|
||||
cutlass::gemm::GemmCoord problem = this->problem_size(lane_problem);
|
||||
cutlass::gemm::GemmCoord grid = this->grid_shape(problem);
|
||||
problem_ending_tile = this->tile_count(grid);
|
||||
}
|
||||
|
||||
// Compute a warp-wide inclusive prefix sum to compute the ending tile index of
|
||||
// each thread's problem.
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 1; i < kThreadsPerWarp; i <<= 1) {
|
||||
int32_t val = __shfl_up_sync(0xffffffff, problem_ending_tile, i);
|
||||
if (lane_idx >= i) {
|
||||
problem_ending_tile += val;
|
||||
}
|
||||
}
|
||||
|
||||
// The total tile count for this group is now in the final position of the prefix sum
|
||||
int32_t tiles_in_group = __shfl_sync(0xffffffff, problem_ending_tile, kThreadsPerWarp - 1);
|
||||
|
||||
problem_ending_tile += group_tile_start;
|
||||
group_tile_end += tiles_in_group;
|
||||
}
|
||||
|
||||
// The next problem to process is the first one that does not have ending tile position
|
||||
// that is greater than or equal to tile index.
|
||||
int32_t problem_idx_in_group = __popc(__ballot_sync(0xffffffff, problem_ending_tile <= this->tile_idx));
|
||||
|
||||
this->problem_idx = group_problem_start + problem_idx_in_group;
|
||||
|
||||
// The starting tile for this problem is the ending tile of the previous problem. In cases
|
||||
// where `problem_idx_in_group` is the first problem in the group, we do not need to reset
|
||||
// `problem_tile_start`, because it is set to the previous group's ending tile in the while
|
||||
// loop above.
|
||||
if (problem_idx_in_group > 0) {
|
||||
this->problem_tile_start = __shfl_sync(0xffffffff, problem_ending_tile, problem_idx_in_group - 1);
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
static size_t get_workspace_size(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr, int32_t problem_count,
|
||||
int32_t block_count) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
static void host_precompute(const cutlass::gemm::GemmCoord* host_problem_sizes_ptr, int32_t problem_count,
|
||||
int32_t block_count, void* host_workspace_ptr) {}
|
||||
};
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
} // namespace cutlass
|
||||
|
|
@ -0,0 +1,61 @@
|
|||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2022 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Defines new layouts needed for MoE
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/matrix_coord.h"
|
||||
#include "cutlass/pitch_linear_coord.h"
|
||||
|
||||
namespace cutlass {
|
||||
namespace layout {
|
||||
|
||||
template <int RowsPerTile, int ColumnsInterleaved>
|
||||
class ColumnMajorTileInterleave {
|
||||
static constexpr int kRowsPerTile = RowsPerTile;
|
||||
static constexpr int kColumnsInterleaved = ColumnsInterleaved;
|
||||
};
|
||||
|
||||
template <class T>
|
||||
struct IsColumnMajorTileInterleave {
|
||||
static constexpr bool value = false;
|
||||
};
|
||||
|
||||
template <int U, int V>
|
||||
struct IsColumnMajorTileInterleave<ColumnMajorTileInterleave<U, V>> {
|
||||
static constexpr bool value = true;
|
||||
};
|
||||
|
||||
} // namespace layout
|
||||
} // namespace cutlass
|
||||
197
onnxruntime/contrib_ops/cuda/moe/moe.cc
Normal file
197
onnxruntime/contrib_ops/cuda/moe/moe.cc
Normal file
|
|
@ -0,0 +1,197 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "core/common/safeint.h"
|
||||
#include "core/providers/cuda/cuda_common.h"
|
||||
#include "moe.h"
|
||||
|
||||
using namespace onnxruntime::cuda;
|
||||
using namespace ::onnxruntime::common;
|
||||
using namespace ONNX_NAMESPACE;
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
namespace cuda {
|
||||
|
||||
#define REGISTER_KERNEL_TYPED(T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
MoE, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()) \
|
||||
.MayInplace(0, 0) \
|
||||
.TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
MoE<T>);
|
||||
|
||||
REGISTER_KERNEL_TYPED(float)
|
||||
REGISTER_KERNEL_TYPED(MLFloat16)
|
||||
|
||||
using namespace ONNX_NAMESPACE;
|
||||
|
||||
template <typename T>
|
||||
Status MoE<T>::ComputeInternal(OpKernelContext* context) const {
|
||||
const Tensor* input = context->Input<Tensor>(0);
|
||||
const Tensor* router_probs = context->Input<Tensor>(1);
|
||||
const Tensor* fc1_experts_weights = context->Input<Tensor>(2);
|
||||
const Tensor* fc2_experts_weights = context->Input<Tensor>(3);
|
||||
const Tensor* fc1_experts_bias_optional = context->Input<Tensor>(4);
|
||||
const Tensor* fc2_experts_bias_optional = context->Input<Tensor>(5);
|
||||
|
||||
const auto& input_dims = input->Shape().GetDims();
|
||||
const auto& router_probs_dims = router_probs->Shape().GetDims();
|
||||
const auto& fc1_experts_weights_dims = fc1_experts_weights->Shape().GetDims();
|
||||
const auto& fc2_experts_weights_dims = fc2_experts_weights->Shape().GetDims();
|
||||
|
||||
const int64_t num_rows = input_dims.size() == 2 ? input_dims[0] : input_dims[0] * input_dims[1];
|
||||
const int64_t hidden_size = input_dims[input_dims.size() - 1];
|
||||
const int64_t num_experts = fc1_experts_weights_dims[0];
|
||||
const int64_t inter_size = fc1_experts_weights_dims[2];
|
||||
|
||||
// TODO: refactor to helper function.
|
||||
if (fc1_experts_weights_dims.size() != 3) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "fc1_experts_weights_dims must be 3D, got ",
|
||||
fc1_experts_weights_dims.size());
|
||||
}
|
||||
if (fc2_experts_weights_dims.size() != 3) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "fc2_experts_weights_dims must be 3D, got ",
|
||||
fc2_experts_weights_dims.size());
|
||||
}
|
||||
if (fc1_experts_weights_dims[1] != hidden_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc1_experts_weights_dims[1] must be equal to hidden_size, got ",
|
||||
fc1_experts_weights_dims[1], " and ", hidden_size);
|
||||
}
|
||||
if (fc2_experts_weights_dims[1] != inter_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc2_experts_weights_dims[1] must be equal to inter_size, got ", fc2_experts_weights_dims[1],
|
||||
" and ", inter_size);
|
||||
}
|
||||
if (fc1_experts_weights_dims[2] != inter_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc1_experts_weights_dims[2] must be equal to inter_size, got ", fc1_experts_weights_dims[2],
|
||||
" and ", inter_size);
|
||||
}
|
||||
if (fc2_experts_weights_dims[2] != hidden_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc2_experts_weights_dims[2] must be equal to hidden_size, got ",
|
||||
fc2_experts_weights_dims[2], " and ", hidden_size);
|
||||
}
|
||||
if (router_probs_dims.size() != 2) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "router_probs_dims must be 2D, got ",
|
||||
router_probs_dims.size());
|
||||
}
|
||||
if (router_probs_dims[0] != num_rows) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "router_probs_dims[0] must be equal to num_rows, got ",
|
||||
router_probs_dims[0], " and ", num_rows);
|
||||
}
|
||||
if (router_probs_dims[1] != num_experts) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "router_probs_dims[1] must be equal to num_experts, got ",
|
||||
router_probs_dims[1], " and ", num_experts);
|
||||
}
|
||||
if (fc1_experts_bias_optional != nullptr && fc2_experts_bias_optional == nullptr) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "fc1_experts_bias is set but fc2_experts_bias is not set");
|
||||
}
|
||||
if (fc1_experts_bias_optional == nullptr && fc2_experts_bias_optional != nullptr) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "fc1_experts_bias is not set but fc2_experts_bias is set");
|
||||
}
|
||||
if (fc1_experts_bias_optional != nullptr && fc2_experts_bias_optional != nullptr) {
|
||||
const auto& fc1_experts_bias_dims = fc1_experts_bias_optional->Shape().GetDims();
|
||||
const auto& fc2_experts_bias_dims = fc2_experts_bias_optional->Shape().GetDims();
|
||||
if (fc1_experts_bias_dims.size() != 2) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "fc1_experts_bias_dims must be 2D, got ",
|
||||
fc1_experts_bias_dims.size());
|
||||
}
|
||||
if (fc2_experts_bias_dims.size() != 2) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "fc2_experts_bias_dims must be 2D, got ",
|
||||
fc2_experts_bias_dims.size());
|
||||
}
|
||||
if (fc1_experts_bias_dims[0] != num_experts) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc1_experts_bias_dims[0] must be equal to num_experts, got ", fc1_experts_bias_dims[0],
|
||||
" and ", num_experts);
|
||||
}
|
||||
if (fc2_experts_bias_dims[0] != num_experts) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc2_experts_bias_dims[0] must be equal to num_experts, got ", fc2_experts_bias_dims[0],
|
||||
" and ", num_experts);
|
||||
}
|
||||
if (fc1_experts_bias_dims[1] != inter_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc1_experts_bias_dims[1] must be equal to inter_size, got ", fc1_experts_bias_dims[1],
|
||||
" and ", inter_size);
|
||||
}
|
||||
if (fc2_experts_bias_dims[1] != hidden_size) {
|
||||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
|
||||
"fc2_experts_bias_dims[1] must be equal to hidden_size, got ", fc2_experts_bias_dims[1],
|
||||
" and ", hidden_size);
|
||||
}
|
||||
}
|
||||
|
||||
typedef typename ToCudaType<T>::MappedType CudaT;
|
||||
auto stream = context->GetComputeStream();
|
||||
|
||||
auto& device_prop = GetDeviceProp();
|
||||
const int sm = device_prop.major * 10 + device_prop.minor;
|
||||
|
||||
ort_fastertransformer::CutlassMoeFCRunner<CudaT, CudaT> moe_runner(sm);
|
||||
|
||||
size_t ws_size =
|
||||
moe_runner.getWorkspaceSize(static_cast<int>(num_rows), static_cast<int>(hidden_size),
|
||||
static_cast<int>(inter_size), static_cast<int>(num_experts), static_cast<int>(k_));
|
||||
size_t fc2_output_size = k_ * num_rows * hidden_size * sizeof(CudaT);
|
||||
size_t expert_scales_size = k_ * num_rows * sizeof(CudaT);
|
||||
size_t expanded_source_row_to_expanded_dest_row_size = k_ * num_rows * sizeof(int);
|
||||
size_t expert_for_source_row_size = k_ * num_rows * sizeof(int);
|
||||
|
||||
AllocatorPtr allocator;
|
||||
ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&allocator));
|
||||
|
||||
// TODO: allocate one buffer and reuse it.
|
||||
IAllocatorUniquePtr<void> work_space = IAllocator::MakeUniquePtr<void>(allocator, ws_size, false, stream);
|
||||
IAllocatorUniquePtr<void> fc2_output = IAllocator::MakeUniquePtr<void>(allocator, fc2_output_size, false, stream);
|
||||
IAllocatorUniquePtr<void> expert_scales =
|
||||
IAllocator::MakeUniquePtr<void>(allocator, expert_scales_size, false, stream);
|
||||
IAllocatorUniquePtr<void> expanded_source_row_to_expanded_dest_row =
|
||||
IAllocator::MakeUniquePtr<void>(allocator, expanded_source_row_to_expanded_dest_row_size, false, stream);
|
||||
IAllocatorUniquePtr<void> expert_for_source_row =
|
||||
IAllocator::MakeUniquePtr<void>(allocator, expert_for_source_row_size, false, stream);
|
||||
|
||||
// fc1_scales and fc2_scales are used in quantized MoE
|
||||
const CudaT* fc1_scales_ptr = nullptr;
|
||||
const CudaT* fc2_scales_ptr = nullptr;
|
||||
|
||||
moe_runner.run_moe_fc(reinterpret_cast<const CudaT*>(input->template Data<T>()),
|
||||
reinterpret_cast<const CudaT*>(router_probs->template Data<T>()),
|
||||
reinterpret_cast<const CudaT*>(fc1_experts_weights->template Data<T>()),
|
||||
std::move(fc1_scales_ptr),
|
||||
fc1_experts_bias_optional == nullptr
|
||||
? nullptr
|
||||
: reinterpret_cast<const CudaT*>(fc1_experts_bias_optional->template Data<T>()),
|
||||
activation_type_, reinterpret_cast<const CudaT*>(fc2_experts_weights->template Data<T>()),
|
||||
std::move(fc2_scales_ptr), static_cast<int>(num_rows), static_cast<int>(hidden_size),
|
||||
static_cast<int>(inter_size), static_cast<int>(num_experts), static_cast<int>(k_),
|
||||
reinterpret_cast<char*>(work_space.get()), reinterpret_cast<CudaT*>(fc2_output.get()),
|
||||
reinterpret_cast<CudaT*>(expert_scales.get()),
|
||||
reinterpret_cast<int*>(expanded_source_row_to_expanded_dest_row.get()),
|
||||
reinterpret_cast<int*>(expert_for_source_row.get()), Stream(context));
|
||||
|
||||
Tensor* output = context->Output(0, input->Shape());
|
||||
|
||||
ort_fastertransformer::finalize_moe_routing_kernelLauncher(
|
||||
reinterpret_cast<CudaT*>(fc2_output.get()), reinterpret_cast<CudaT*>(output->template MutableData<T>()),
|
||||
fc2_experts_bias_optional == nullptr
|
||||
? nullptr
|
||||
: reinterpret_cast<const CudaT*>(fc2_experts_bias_optional->template Data<T>()),
|
||||
reinterpret_cast<CudaT*>(expert_scales.get()),
|
||||
reinterpret_cast<int*>(expanded_source_row_to_expanded_dest_row.get()),
|
||||
reinterpret_cast<int*>(expert_for_source_row.get()), static_cast<int>(num_rows), static_cast<int>(hidden_size),
|
||||
static_cast<int>(k_), Stream(context));
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
45
onnxruntime/contrib_ops/cuda/moe/moe.h
Normal file
45
onnxruntime/contrib_ops/cuda/moe/moe.h
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "contrib_ops/cuda/moe/ft_moe/moe_kernel.h"
|
||||
#include "core/common/common.h"
|
||||
#include "core/providers/cuda/cuda_kernel.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
namespace cuda {
|
||||
|
||||
using namespace onnxruntime::cuda;
|
||||
|
||||
template <typename T>
|
||||
class MoE final : public CudaKernel {
|
||||
public:
|
||||
explicit MoE(const OpKernelInfo& op_kernel_info) : CudaKernel(op_kernel_info) {
|
||||
ORT_ENFORCE(op_kernel_info.GetAttr<int64_t>("k", &k_).IsOK());
|
||||
|
||||
std::string activation_type_str;
|
||||
ORT_ENFORCE(op_kernel_info.GetAttr<std::string>("activation_type", &activation_type_str).IsOK());
|
||||
if (activation_type_str == "relu") {
|
||||
activation_type_ = ort_fastertransformer::ActivationType::Relu;
|
||||
} else if (activation_type_str == "gelu") {
|
||||
activation_type_ = ort_fastertransformer::ActivationType::Gelu;
|
||||
} else if (activation_type_str == "silu") {
|
||||
activation_type_ = ort_fastertransformer::ActivationType::Silu;
|
||||
} else if (activation_type_str == "identity") {
|
||||
activation_type_ = ort_fastertransformer::ActivationType::Identity;
|
||||
} else {
|
||||
ORT_THROW("Unsupported MoE activation type: ", activation_type_str);
|
||||
}
|
||||
}
|
||||
Status ComputeInternal(OpKernelContext* ctx) const override;
|
||||
|
||||
private:
|
||||
int64_t k_;
|
||||
ort_fastertransformer::ActivationType activation_type_;
|
||||
};
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -1375,6 +1375,27 @@ ONNX_MS_OPERATOR_SET_SCHEMA(Sampling, 1,
|
|||
GreedySearchShapeInference(ctx);
|
||||
}));
|
||||
|
||||
constexpr const char* MoE_ver1_doc = R"DOC(
|
||||
Mixture of experts. Examples: Switch transformer(https://arxiv.org/pdf/2101.03961.pdf) use top 1,
|
||||
GLaM(https://arxiv.org/abs/2112.06905) activates top 2 FFN, and Vision MOE(https://arxiv.org/pdf/2106.05974.pdf)
|
||||
usually uses top 32 experts.
|
||||
)DOC";
|
||||
|
||||
ONNX_MS_OPERATOR_SET_SCHEMA(MoE, 1,
|
||||
OpSchema()
|
||||
.SetDoc(MoE_ver1_doc)
|
||||
.Attr("activation_type", "Activation function to use. Choose from relu, gelu, silu and identity. Default is relu", AttributeProto::STRING, std::string("relu"))
|
||||
.Attr("k", "Number of top experts to select from expert pool", AttributeProto::INT, static_cast<int64_t>(1))
|
||||
.Input(0, "input", "2D input tensor with shape (num_rows, hidden_size) or 3D input tensor with shape (batch_size, sequence_length, hidden_size)", "T")
|
||||
.Input(1, "router_probs", "2D input tensor with shape (num_rows, num_experts)", "T")
|
||||
.Input(2, "fc1_experts_weights", "3D input tensor with shape (num_experts, hidden_size, inter_size)", "T")
|
||||
.Input(3, "fc2_experts_weights", "3D input tensor with shape (num_experts, inter_size, hidden_size)", "T")
|
||||
.Input(4, "fc1_experts_bias", "2D optional input tensor with shape (num_experts, inter_size)", "T", OpSchema::Optional)
|
||||
.Input(5, "fc2_experts_bias", "2D optional input tensor with shape (num_experts, hidden_size)", "T", OpSchema::Optional)
|
||||
.Output(0, "output", "2D input tensor with shape (num_rows, hidden_size) or 3D input tensor with shape (batch_size, sequence_length, hidden_size)", "T")
|
||||
.TypeConstraint("T", {"tensor(float)", "tensor(float16)"}, "Constrain input and output types to float or float16 tensors.")
|
||||
.TypeAndShapeInferenceFunction(ONNX_NAMESPACE::propagateShapeAndTypeFromFirstInput));
|
||||
|
||||
ONNX_MS_OPERATOR_SET_SCHEMA(SampleOp, 1,
|
||||
OpSchema()
|
||||
.Input(0, "X", "input", "T")
|
||||
|
|
|
|||
|
|
@ -83,6 +83,7 @@ class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulInteger16);
|
|||
class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulFpQ4);
|
||||
#endif
|
||||
class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MaxpoolWithMask);
|
||||
class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MoE);
|
||||
class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MultiHeadAttention);
|
||||
class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GroupQueryAttention);
|
||||
class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MurmurHash3);
|
||||
|
|
@ -189,6 +190,7 @@ class OpSet_Microsoft_ver1 {
|
|||
fn(GetOpSchema<ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MatMulFpQ4)>());
|
||||
#endif
|
||||
fn(GetOpSchema<ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MaxpoolWithMask)>());
|
||||
fn(GetOpSchema<ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MoE)>());
|
||||
fn(GetOpSchema<ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MultiHeadAttention)>());
|
||||
fn(GetOpSchema<ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GroupQueryAttention)>());
|
||||
fn(GetOpSchema<ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, MurmurHash3)>());
|
||||
|
|
|
|||
|
|
@ -157,6 +157,7 @@ class SymbolicShapeInference:
|
|||
"MemcpyFromHost": self._pass_on_shape_and_type,
|
||||
"MemcpyToHost": self._pass_on_shape_and_type,
|
||||
"Min": self._infer_symbolic_compute_ops,
|
||||
"MoE": self._pass_on_shape_and_type,
|
||||
"Mul": self._infer_symbolic_compute_ops,
|
||||
"NonMaxSuppression": self._infer_NonMaxSuppression,
|
||||
"NonZero": self._infer_NonZero,
|
||||
|
|
|
|||
423
onnxruntime/test/contrib_ops/moe_test.cc
Normal file
423
onnxruntime/test/contrib_ops/moe_test.cc
Normal file
|
|
@ -0,0 +1,423 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "gtest/gtest.h"
|
||||
#include "test/common/tensor_op_test_utils.h"
|
||||
#include "test/common/cuda_op_test_utils.h"
|
||||
#include "test/providers/provider_test_utils.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace test {
|
||||
|
||||
static void RunMoETest(
|
||||
const std::vector<float>& input,
|
||||
const std::vector<float>& router_probs,
|
||||
const std::vector<float>& fc1_experts_weights,
|
||||
const std::vector<float>& fc2_experts_weights,
|
||||
const std::vector<float>& fc1_experts_bias,
|
||||
const std::vector<float>& fc2_experts_bias,
|
||||
const std::vector<float>& output_data,
|
||||
int num_rows,
|
||||
int num_experts,
|
||||
int hidden_size,
|
||||
int inter_size,
|
||||
std::string activation_type,
|
||||
bool use_float16 = false) {
|
||||
int min_cuda_architecture = use_float16 ? 530 : 0;
|
||||
|
||||
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture);
|
||||
if (enable_cuda) {
|
||||
OpTester tester("MoE", 1, onnxruntime::kMSDomain);
|
||||
tester.AddAttribute<int64_t>("k", static_cast<int64_t>(1));
|
||||
tester.AddAttribute<std::string>("activation_type", activation_type);
|
||||
|
||||
std::vector<int64_t> input_dims = {num_rows, hidden_size};
|
||||
std::vector<int64_t> router_probs_dims = {num_rows, num_experts};
|
||||
std::vector<int64_t> fc1_experts_weights_dims = {num_experts, hidden_size, inter_size};
|
||||
std::vector<int64_t> fc2_experts_weights_dims = {num_experts, inter_size, hidden_size};
|
||||
std::vector<int64_t> fc1_experts_bias_dims = {num_experts, inter_size};
|
||||
std::vector<int64_t> fc2_experts_bias_dims = {num_experts, hidden_size};
|
||||
std::vector<int64_t> output_dims = {num_rows, hidden_size};
|
||||
|
||||
if (use_float16) {
|
||||
tester.AddInput<MLFloat16>("input", input_dims, ToFloat16(input));
|
||||
tester.AddInput<MLFloat16>("router_probs", router_probs_dims, ToFloat16(router_probs));
|
||||
tester.AddInput<MLFloat16>("fc1_experts_weights", fc1_experts_weights_dims, ToFloat16(fc1_experts_weights));
|
||||
tester.AddInput<MLFloat16>("fc2_experts_weights", fc2_experts_weights_dims, ToFloat16(fc2_experts_weights));
|
||||
tester.AddInput<MLFloat16>("fc1_experts_bias", fc1_experts_bias_dims, ToFloat16(fc1_experts_bias));
|
||||
tester.AddInput<MLFloat16>("fc2_experts_bias", fc2_experts_bias_dims, ToFloat16(fc2_experts_bias));
|
||||
tester.AddOutput<MLFloat16>("output", output_dims, ToFloat16(output_data));
|
||||
} else {
|
||||
tester.AddInput<float>("input", input_dims, input);
|
||||
tester.AddInput<float>("router_probs", router_probs_dims, router_probs);
|
||||
tester.AddInput<float>("fc1_experts_weights", fc1_experts_weights_dims, fc1_experts_weights);
|
||||
tester.AddInput<float>("fc2_experts_weights", fc2_experts_weights_dims, fc2_experts_weights);
|
||||
tester.AddInput<float>("fc1_experts_bias", fc1_experts_bias_dims, fc1_experts_bias);
|
||||
tester.AddInput<float>("fc2_experts_bias", fc2_experts_bias_dims, fc2_experts_bias);
|
||||
tester.AddOutput<float>("output", output_dims, output_data);
|
||||
}
|
||||
|
||||
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
|
||||
execution_providers.push_back(DefaultCudaExecutionProvider());
|
||||
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
|
||||
}
|
||||
}
|
||||
|
||||
TEST(MoETest, MoETest_Gelu) {
|
||||
int num_rows = 4;
|
||||
int num_experts = 4;
|
||||
int hidden_size = 8;
|
||||
int inter_size = 16;
|
||||
|
||||
const std::vector<float> input = {
|
||||
-1.1200173f, -0.45884353f, -1.2929888f, 1.0784022f, 0.116372705f, 0.26902613f, -1.8818876f, -0.5457026f,
|
||||
0.22222236f, -0.28868636f, 0.6692926f, 1.4944887f, 0.02431708f, -0.49781424f, 0.7378293f, 1.276276f,
|
||||
-0.15469065f, -0.28456813f, -0.6296439f, -0.24855971f, 0.80565417f, -1.1018785f, -0.74082595f, 0.82407707f,
|
||||
-0.95033455f, 0.659333f, -0.68629056f, -0.2916592f, 1.869919f, -1.1053563f, -0.14417848f, -0.34625578f};
|
||||
const std::vector<float> router_probs = {
|
||||
-0.84837115f, 0.100507565f, -0.10548311f, 0.40957215f, 1.0159845f, 0.26919764f, 0.021741152f, -0.34184334f,
|
||||
-0.71324956f, 0.29018253f, -0.18227568f, 0.31496462f, -0.48426327f, -1.006643f, -0.100081146f, -0.07692295f};
|
||||
const std::vector<float> fc1_experts_weights = {
|
||||
0.14731085f, 0.52229995f, 0.14753294f, 0.22475791f, 0.20864725f, 0.6708725f, 0.20204341f, 0.4890914f,
|
||||
0.52103406f, 0.8223115f, 0.122039974f, 0.15674388f, 0.20966923f, 0.8499667f, 0.3202675f, 0.92174435f,
|
||||
0.6808038f, 0.563313f, 0.496278f, 0.40115923f, 0.5627332f, 0.38582766f, 0.49648678f, 0.5637965f,
|
||||
0.10889745f, 0.23793429f, 0.90374637f, 0.09422666f, 0.4640969f, 0.99461937f, 0.6806185f, 0.5141565f,
|
||||
0.066695035f, 0.74768895f, 0.14385962f, 0.35806787f, 0.33224183f, 0.4259563f, 0.50546914f, 0.91240376f,
|
||||
0.5624194f, 0.9478464f, 0.8058562f, 0.18389302f, 0.72425205f, 0.14655197f, 0.28808743f, 0.64706135f,
|
||||
0.66509604f, 0.875114f, 0.33904207f, 0.50080043f, 0.7574118f, 0.016453922f, 0.8614903f, 0.08653879f,
|
||||
0.50689125f, 0.41499162f, 0.23666352f, 0.5660855f, 0.91345936f, 0.35384023f, 0.20315295f, 0.31508058f,
|
||||
0.0044258237f, 0.725697f, 0.25986814f, 0.16632986f, 0.21194929f, 0.787478f, 0.76478684f, 0.8837609f,
|
||||
0.68136156f, 0.33302015f, 0.36027592f, 0.647715f, 0.91101736f, 0.6359461f, 0.26342732f, 0.2649613f,
|
||||
0.02726549f, 0.608024f, 0.21940875f, 0.054212093f, 0.93843824f, 0.1752944f, 0.44311923f, 0.64324677f,
|
||||
0.51592916f, 0.16355914f, 0.09583914f, 0.8985412f, 0.58141935f, 0.91481227f, 0.3323797f, 0.6472777f,
|
||||
0.3856619f, 0.47776443f, 0.1954779f, 0.66910046f, 0.65808296f, 0.4896857f, 0.38754892f, 0.1917851f,
|
||||
0.8457724f, 0.12778795f, 0.70483273f, 0.33187324f, 0.258766f, 0.58982253f, 0.24027151f, 0.6152024f,
|
||||
0.5981904f, 0.12875527f, 0.5832493f, 0.7129646f, 0.6979155f, 0.43706065f, 0.09010619f, 0.42292297f,
|
||||
0.67365384f, 0.31756145f, 0.68979055f, 0.8329813f, 0.2389242f, 0.5049309f, 0.7067495f, 0.5391889f,
|
||||
0.54176575f, 0.5624327f, 0.10692614f, 0.5392941f, 0.8462349f, 0.9505569f, 0.79387546f, 0.5670015f,
|
||||
0.7335071f, 0.25676018f, 0.08565581f, 0.07003945f, 0.99880487f, 0.8173947f, 0.15438312f, 0.6956213f,
|
||||
0.8775838f, 0.9998074f, 0.93719745f, 0.8873769f, 0.38537037f, 0.32452917f, 0.9105244f, 0.7801898f,
|
||||
0.19911051f, 0.9495086f, 0.7415793f, 0.77256775f, 0.18661183f, 0.6434499f, 0.32471877f, 0.8906783f,
|
||||
0.4100297f, 0.69465625f, 0.5888109f, 0.7127341f, 0.33008623f, 0.7437857f, 0.15076452f, 0.6129275f,
|
||||
0.16170406f, 0.006731212f, 0.09847212f, 0.89473504f, 0.7705178f, 0.96910787f, 0.9005606f, 0.053477287f,
|
||||
0.15878445f, 0.4192087f, 0.17528385f, 0.84719825f, 0.121996105f, 0.25604928f, 0.016954303f, 0.21612722f,
|
||||
0.91123873f, 0.90938f, 0.85791886f, 0.88606364f, 0.94459325f, 0.3719685f, 0.72000104f, 0.9454652f,
|
||||
0.6654094f, 0.9998382f, 0.75933146f, 0.81082416f, 0.32500392f, 0.73991376f, 0.5574533f, 0.38059133f,
|
||||
0.21814507f, 0.21944171f, 0.11525959f, 0.83566517f, 0.8554656f, 0.44309366f, 0.210657f, 0.88645273f,
|
||||
0.81974447f, 0.537167f, 0.26393235f, 0.9595239f, 0.70447034f, 0.12042731f, 0.97854143f, 0.8796869f,
|
||||
0.31775457f, 0.78107727f, 0.21590549f, 0.42164284f, 0.9245506f, 0.52065957f, 0.14639091f, 0.33288354f,
|
||||
0.36427742f, 0.4035356f, 0.5478503f, 0.9624148f, 0.5267702f, 0.19128f, 0.52562714f, 0.7397436f,
|
||||
0.7480201f, 0.04303074f, 0.41052878f, 0.12842774f, 0.2866572f, 0.6801467f, 0.1449349f, 0.68586344f,
|
||||
0.92438906f, 0.5327942f, 0.16675615f, 0.32085752f, 0.60918206f, 0.11884099f, 0.74840516f, 0.04606521f,
|
||||
0.01935333f, 0.014169693f, 0.39856833f, 0.83621645f, 0.026760519f, 0.91559356f, 0.29998857f, 0.64644206f,
|
||||
0.52280146f, 0.049140453f, 0.9146645f, 0.7692217f, 0.99699783f, 0.7526061f, 0.1699655f, 0.9172919f,
|
||||
0.5268722f, 0.73710823f, 0.09908545f, 0.35618675f, 0.009061217f, 0.30525374f, 0.6078656f, 0.10741913f,
|
||||
0.6593821f, 0.7684034f, 0.56965464f, 0.16545832f, 0.11234015f, 0.3457417f, 0.7194791f, 0.9931982f,
|
||||
0.7875145f, 0.44369537f, 0.6753082f, 0.009468555f, 0.07294935f, 0.73330396f, 0.2167924f, 0.74054784f,
|
||||
0.14703393f, 0.25234455f, 0.08815551f, 0.76092035f, 0.44905245f, 0.88480055f, 0.8094361f, 0.7766713f,
|
||||
0.51607805f, 0.345411f, 0.39128417f, 0.5664503f, 0.74785477f, 0.14970505f, 0.91963893f, 0.44563496f,
|
||||
0.08102721f, 0.22947109f, 0.94240886f, 0.9572636f, 0.036860168f, 0.85264915f, 0.7505796f, 0.79595923f,
|
||||
0.9232646f, 0.23052484f, 0.6578879f, 0.7046166f, 0.35225332f, 0.66732657f, 0.3561433f, 0.80913067f,
|
||||
0.3612727f, 0.31360215f, 0.6258745f, 0.6773468f, 0.25571418f, 0.54419917f, 0.78976786f, 0.45025164f,
|
||||
0.65216696f, 0.3794065f, 0.6752498f, 0.1378029f, 0.2059856f, 0.24620473f, 0.95950544f, 0.36545795f,
|
||||
0.49863482f, 0.25775224f, 0.99914503f, 0.9883351f, 0.122906685f, 0.09466505f, 0.12100351f, 0.49758863f,
|
||||
0.37254804f, 0.17272717f, 0.32066393f, 0.59446543f, 0.23875463f, 0.61079127f, 0.38534206f, 0.25771832f,
|
||||
0.56869274f, 0.9111291f, 0.16196036f, 0.5232172f, 0.31561613f, 0.99065316f, 0.025618374f, 0.0206694f,
|
||||
0.9926925f, 0.18365502f, 0.5958617f, 0.45684695f, 0.3946715f, 0.3883261f, 0.8177203f, 0.5238985f,
|
||||
0.013192713f, 0.20481992f, 0.32954985f, 0.7516082f, 0.17643315f, 0.9714598f, 0.38863534f, 0.410219f,
|
||||
0.891779f, 0.75130385f, 0.92406017f, 0.7892222f, 0.34832305f, 0.1682638f, 0.46279848f, 0.9138188f,
|
||||
0.3321901f, 0.036315024f, 0.7049642f, 0.9867357f, 0.3576584f, 0.08598822f, 0.046470165f, 0.6252997f,
|
||||
0.46214014f, 0.24750638f, 0.60106593f, 0.6898794f, 0.8976595f, 0.8881911f, 0.42515814f, 0.059116423f,
|
||||
0.048188448f, 0.9668448f, 0.7210276f, 0.7179537f, 0.06738949f, 0.96300787f, 0.97367156f, 0.95143014f,
|
||||
0.07820749f, 0.3113383f, 0.1561181f, 0.9734828f, 0.28516f, 0.27172273f, 0.76195645f, 0.26870382f,
|
||||
0.25373894f, 0.45626426f, 0.45194024f, 0.11051077f, 0.91683406f, 0.27943915f, 0.67735744f, 0.9348918f,
|
||||
0.7521582f, 0.57078993f, 0.9254285f, 0.5672131f, 0.2686717f, 0.97299975f, 0.61834025f, 0.012159586f,
|
||||
0.3576542f, 0.15941626f, 0.9383765f, 0.41742706f, 0.044237554f, 0.46856833f, 0.81400645f, 0.6299002f,
|
||||
0.6581022f, 0.5464366f, 0.68640935f, 0.378174f, 0.3010999f, 0.032645762f, 0.12333155f, 0.71670127f,
|
||||
0.20394331f, 0.57173324f, 0.6595957f, 0.53540194f, 0.17582512f, 0.9781642f, 0.20925027f, 0.9112503f,
|
||||
0.10224587f, 0.37972575f, 0.7719844f, 0.29570967f, 0.9200215f, 0.15592176f, 0.080114245f, 0.27454042f,
|
||||
0.5808252f, 0.96037793f, 0.26129955f, 0.6788141f, 0.37464648f, 0.39156884f, 0.8676517f, 0.112507045f,
|
||||
0.55310667f, 0.9702046f, 0.4312939f, 0.88821906f, 0.3460216f, 0.9024811f, 0.016334832f, 0.42793816f,
|
||||
0.4121768f, 0.6620425f, 0.6961637f, 0.88390845f, 0.425507f, 0.48017246f, 0.8424056f, 0.36471343f,
|
||||
0.9383168f, 0.16709393f, 0.44589508f, 0.47314453f, 0.72310495f, 0.84183806f, 0.4207481f, 0.0857597f,
|
||||
0.7477461f, 0.6495659f, 0.70084965f, 0.19156617f, 0.8217978f, 0.9735775f, 0.5433857f, 0.032975793f,
|
||||
0.85099494f, 0.12927437f, 0.61493605f, 0.5726589f, 0.26598173f, 0.6740978f, 0.052783668f, 0.61387974f};
|
||||
const std::vector<float> fc2_experts_weights = {
|
||||
0.18302453f, 0.44593316f, 0.5643144f, 0.9259722f, 0.26143986f, 0.82031804f, 0.4364831f, 0.2625361f,
|
||||
0.06460017f, 0.04124081f, 0.98830533f, 0.37530023f, 0.5249744f, 0.63555616f, 0.8398661f, 0.92673707f,
|
||||
0.9055086f, 0.12955844f, 0.4198916f, 0.20413119f, 0.21432412f, 0.6186035f, 0.969324f, 0.099448025f,
|
||||
0.80260223f, 0.24076664f, 0.40261286f, 0.89688545f, 0.38691485f, 0.5455279f, 0.15048373f, 0.92562044f,
|
||||
0.43536508f, 0.13430476f, 0.64640516f, 0.14449131f, 0.10324633f, 0.5304596f, 0.8964218f, 0.358508f,
|
||||
0.73533344f, 0.9296606f, 0.83163047f, 0.23771948f, 0.44519007f, 0.34265757f, 0.09793854f, 0.5002066f,
|
||||
0.87621754f, 0.9212578f, 0.54665035f, 0.6135615f, 0.28353918f, 0.8774212f, 0.29194576f, 0.1526736f,
|
||||
0.57699674f, 0.7996927f, 0.04920423f, 0.95198375f, 0.67986554f, 0.14969361f, 0.39229625f, 0.93378997f,
|
||||
0.11638266f, 0.3538614f, 0.66399014f, 0.06195748f, 0.7740991f, 0.7602738f, 0.81010276f, 0.18122643f,
|
||||
0.9980005f, 0.20361924f, 0.99917024f, 0.020154774f, 0.054515004f, 0.80709815f, 0.55225646f, 0.52884465f,
|
||||
0.22312081f, 0.29026228f, 0.35380626f, 0.012922287f, 0.52598435f, 0.58842945f, 0.4995767f, 0.66146517f,
|
||||
0.9744255f, 0.632942f, 0.3169638f, 0.29422665f, 0.18009722f, 0.15339059f, 0.41947508f, 0.4115672f,
|
||||
0.72243124f, 0.2862816f, 0.89860183f, 0.14915991f, 0.5014211f, 0.94945997f, 0.99719256f, 0.21036887f,
|
||||
0.5890645f, 0.55906135f, 0.26557416f, 0.32725257f, 0.635427f, 0.1523174f, 0.58249784f, 0.71636236f,
|
||||
0.30296493f, 0.9153206f, 0.46709478f, 0.72685635f, 0.9951532f, 0.34716582f, 0.7717041f, 0.3569854f,
|
||||
0.4269635f, 0.41526443f, 0.4968937f, 0.3111158f, 0.61719346f, 0.5188402f, 0.8169449f, 0.39879733f,
|
||||
0.5501401f, 0.31400484f, 0.08127314f, 0.7023336f, 0.56397897f, 0.29975814f, 0.33094752f, 0.63076067f,
|
||||
0.40959156f, 0.82673794f, 0.52832156f, 0.68886834f, 0.7178481f, 0.37731683f, 0.71633244f, 0.86896664f,
|
||||
0.5230092f, 0.59784645f, 0.5181678f, 0.8461837f, 0.28890234f, 0.23421508f, 0.7178768f, 0.06484294f,
|
||||
0.5080162f, 0.27005446f, 0.8300168f, 0.034480453f, 0.8031663f, 0.9946784f, 0.60117006f, 0.46668667f,
|
||||
0.9921749f, 0.28632385f, 0.45993322f, 0.28104752f, 0.43097937f, 0.60866946f, 0.5667807f, 0.40556252f,
|
||||
7.969141e-05f, 0.52560204f, 0.48518902f, 0.5752184f, 0.8831251f, 0.9860047f, 0.20335877f, 0.46882278f,
|
||||
0.2996632f, 0.03917718f, 0.13617045f, 0.96928054f, 0.79153055f, 0.76857555f, 0.7778716f, 0.102760494f,
|
||||
0.5525096f, 0.9653573f, 0.22095704f, 0.94479716f, 0.63141924f, 0.8517718f, 0.28580618f, 0.73050886f,
|
||||
0.05675614f, 0.46825224f, 0.6667756f, 0.6499472f, 0.91840404f, 0.99132854f, 0.9548785f, 0.8356961f,
|
||||
0.851531f, 0.43548512f, 0.111976564f, 0.31438643f, 0.44386774f, 0.22980672f, 0.75558543f, 0.6755136f,
|
||||
0.58067596f, 0.62078035f, 0.93922615f, 0.6821157f, 0.061530292f, 0.13705963f, 0.7203748f, 0.5681396f,
|
||||
0.7438458f, 0.0006400347f, 0.038565338f, 0.8066132f, 0.81982285f, 0.047644496f, 0.68979263f, 0.109577894f,
|
||||
0.8786539f, 0.6568952f, 0.99439347f, 0.0070040226f, 0.018661916f, 0.838051f, 0.94391155f, 0.80634f,
|
||||
0.8324149f, 0.078864336f, 0.8619068f, 0.027926445f, 0.61170083f, 0.17248261f, 0.30140227f, 0.5885344f,
|
||||
0.30341f, 0.42088854f, 0.02608782f, 0.02856338f, 0.69368154f, 0.28836077f, 0.19580519f, 0.30270886f,
|
||||
0.09121573f, 0.100299895f, 0.79918617f, 0.75412107f, 0.56660175f, 0.22687018f, 0.6663505f, 0.5224626f,
|
||||
0.1426636f, 0.6075949f, 0.95527196f, 0.008196831f, 0.0028039217f, 0.5640625f, 0.87651116f, 0.19575512f,
|
||||
0.61006856f, 0.85149264f, 0.6541582f, 0.6082054f, 0.998863f, 0.82573634f, 0.21878648f, 0.54321826f,
|
||||
0.7554362f, 0.94095474f, 0.002533555f, 0.77075267f, 0.35483408f, 0.010389388f, 0.610987f, 0.22779316f,
|
||||
0.5708561f, 0.17537653f, 0.12373549f, 0.4575745f, 0.33203715f, 0.79243237f, 0.54310906f, 0.8902793f,
|
||||
0.5937015f, 0.33921933f, 0.8386668f, 0.52732253f, 0.59384584f, 0.3391887f, 0.5017944f, 0.40386343f,
|
||||
0.45749134f, 0.110060334f, 0.49692506f, 0.084977865f, 0.3924346f, 0.7897731f, 0.15232486f, 0.16297412f,
|
||||
0.37791175f, 0.36293298f, 0.5846437f, 0.5830078f, 0.75354826f, 0.15555972f, 0.4647144f, 0.7796456f,
|
||||
0.93248576f, 0.46352726f, 0.2106899f, 0.6437313f, 0.78473866f, 0.18762505f, 0.20985329f, 0.7209991f,
|
||||
0.464967f, 0.02775067f, 0.21170747f, 0.7027664f, 0.33041215f, 0.8451145f, 0.89526993f, 0.57273495f,
|
||||
0.46046263f, 0.34128642f, 0.47471708f, 0.59101045f, 0.11807448f, 0.38050216f, 0.08409953f, 0.80687743f,
|
||||
0.18158185f, 0.9567719f, 0.3711096f, 0.21356237f, 0.74022657f, 0.57453954f, 0.846228f, 0.70873487f,
|
||||
0.018330276f, 0.8162452f, 0.40584308f, 0.27901447f, 0.81752694f, 0.86466515f, 0.060534656f, 0.45478833f,
|
||||
0.9106033f, 0.6936434f, 0.92123467f, 0.32865065f, 0.22417879f, 0.9299548f, 0.70841146f, 0.97999126f,
|
||||
0.2911517f, 0.17896658f, 0.44139355f, 0.029210031f, 0.6959876f, 0.8687942f, 0.62002844f, 0.45059657f,
|
||||
0.74790317f, 0.18262434f, 0.98912156f, 0.0028281808f, 0.021027386f, 0.38184917f, 0.90842223f, 0.5500629f,
|
||||
0.69202286f, 0.13349658f, 0.6823429f, 0.44412827f, 0.7004118f, 0.8531213f, 0.7173401f, 0.4574679f,
|
||||
0.46920043f, 0.18640989f, 0.31914896f, 0.82491904f, 0.29950172f, 0.8105199f, 0.30173403f, 0.38355058f,
|
||||
0.5106411f, 0.04116726f, 0.49500751f, 0.44960213f, 0.45508182f, 0.4000479f, 0.89418864f, 0.8689936f,
|
||||
0.16112137f, 0.7322634f, 0.10780871f, 0.07433933f, 0.652841f, 0.50734824f, 0.26674682f, 0.017748117f,
|
||||
0.30643195f, 0.66699976f, 0.03719926f, 0.014267266f, 0.56343627f, 0.13979793f, 0.061959863f, 0.3073569f,
|
||||
0.41949958f, 0.045647383f, 0.16613615f, 0.5327839f, 0.028514147f, 0.4297228f, 0.17714864f, 0.15338135f,
|
||||
0.6965155f, 0.11515516f, 0.1210829f, 0.78514075f, 0.59348315f, 0.9553564f, 0.36635226f, 0.25849247f,
|
||||
0.45372677f, 0.5025297f, 0.88132215f, 0.0019600391f, 0.46439964f, 0.7211761f, 0.22465849f, 0.2459296f,
|
||||
0.7416339f, 0.020907402f, 0.6184779f, 0.112906754f, 0.7485309f, 0.072479784f, 0.8074024f, 0.026683688f,
|
||||
0.07971662f, 0.50736845f, 0.8939942f, 0.0718022f, 0.27697015f, 0.9391413f, 0.4161513f, 0.7071423f,
|
||||
0.019000888f, 0.34275955f, 0.24608392f, 0.9215306f, 0.70751995f, 0.13516217f, 0.5806135f, 0.49425328f,
|
||||
0.29456508f, 0.21446168f, 0.3340807f, 0.89411324f, 0.14157385f, 0.14382833f, 0.34574044f, 0.50869817f,
|
||||
0.63610595f, 0.51500404f, 0.37963718f, 0.19682491f, 0.41028368f, 0.29872334f, 0.9039644f, 0.013295233f,
|
||||
0.1810705f, 0.093204916f, 0.4086216f, 0.8896367f, 0.9382696f, 0.06472236f, 0.47833657f, 0.7934831f,
|
||||
0.7203987f, 0.9095519f, 0.4861309f, 0.16405362f, 0.83076525f, 0.3285427f, 0.7588931f, 0.37678176f,
|
||||
0.71254706f, 0.949713f, 0.96492773f, 0.044967473f, 0.16925985f, 0.2932666f, 0.18114948f, 0.97975004f,
|
||||
0.4558406f, 0.16832972f, 0.27750528f, 0.2238177f, 0.7039947f, 0.06387442f, 0.033798456f, 0.007119417f};
|
||||
const std::vector<float> fc1_experts_bias = {
|
||||
0.71526206f, 0.7472273f, 0.18946046f, 0.6239893f, 0.86909235f, 0.5726507f, 0.3942092f, 0.5369412f,
|
||||
0.44638616f, 0.7517496f, 0.16049433f, 0.75355124f, 0.7818118f, 0.19706267f, 0.9082818f, 0.9910924f,
|
||||
0.30288565f, 0.3599528f, 0.74917775f, 0.10828978f, 0.697729f, 0.61665237f, 0.81516486f, 0.0656966f,
|
||||
0.0846076f, 0.72456455f, 0.6801054f, 0.034616888f, 0.22117025f, 0.042510748f, 0.14178854f, 0.27440017f,
|
||||
0.91376925f, 0.40047455f, 0.7871756f, 0.97484046f, 0.7278661f, 0.052394807f, 0.75161135f, 0.6907173f,
|
||||
0.8875328f, 0.0067828894f, 0.807508f, 0.9092707f, 0.034817636f, 0.55231315f, 0.92683655f, 0.13634592f,
|
||||
0.66405964f, 0.7209387f, 0.63104504f, 0.9971379f, 0.9093898f, 0.9289774f, 0.4376766f, 0.9193563f,
|
||||
0.03404367f, 0.23018533f, 0.39305943f, 0.3514716f, 0.96184736f, 0.73583263f, 0.8219065f, 0.8401047f};
|
||||
const std::vector<float> fc2_experts_bias = {
|
||||
0.12649822f, 0.4420895f, 0.5730123f, 0.63004625f, 0.7571163f, 0.3010466f, 0.3492328f, 0.91837066f,
|
||||
0.36580783f, 0.15267932f, 0.8390199f, 0.83857775f, 0.34321654f, 0.40003997f, 0.13106f, 0.08245313f,
|
||||
0.68802476f, 0.28640372f, 0.89804775f, 0.09964341f, 0.43088746f, 0.5107959f, 0.75697356f, 0.90466535f,
|
||||
0.83860224f, 0.720098f, 0.2705031f, 0.14292616f, 0.052693605f, 0.5248023f, 0.9849401f, 0.40502876f};
|
||||
const std::vector<float> output = {
|
||||
0.2552814f, 0.17651685f, 0.0034551744f, -0.123282805f, 0.0073816925f, 0.004265253f, 0.16927283f, -0.05276826f,
|
||||
9.555821f, 7.6907287f, 10.626425f, 7.0543795f, 8.10093f, 10.3664465f, 10.925815f, 8.737018f,
|
||||
0.565234f, 0.17098689f, 0.10810414f, 0.43916586f, 0.3535297f, 0.45673048f, 0.3853893f, 0.18613164f,
|
||||
1.3354061f, 0.5049282f, 0.72775036f, 0.90331376f, 1.2945517f, 0.9123066f, 1.1995136f, 0.7708638f};
|
||||
|
||||
RunMoETest(input,
|
||||
router_probs,
|
||||
fc1_experts_weights,
|
||||
fc2_experts_weights,
|
||||
fc1_experts_bias,
|
||||
fc2_experts_bias,
|
||||
output,
|
||||
num_rows,
|
||||
num_experts,
|
||||
hidden_size,
|
||||
inter_size,
|
||||
"gelu");
|
||||
}
|
||||
|
||||
TEST(MoETest, MoETest_Relu) {
|
||||
int num_rows = 4;
|
||||
int num_experts = 4;
|
||||
int hidden_size = 8;
|
||||
int inter_size = 16;
|
||||
|
||||
const std::vector<float> input = {
|
||||
0.7670296f, -0.93721074f, -2.330477f, -0.78088343f, 0.8250065f, 1.2206652f, -0.06297584f, 1.1463639f,
|
||||
1.2215378f, -0.31372663f, -0.7234253f, -0.3627346f, 0.44249064f, 0.19418247f, -0.49998695f, -0.55005103f,
|
||||
0.023851749f, -1.5203826f, 0.52939993f, -0.39082858f, -1.9291036f, 0.034976702f, -0.48336256f, -1.226073f,
|
||||
-0.33963847f, 0.0073261578f, -0.0521804f, 1.16749f, 1.7302082f, 2.0561688f, -0.2347232f, -1.3456243f};
|
||||
const std::vector<float> router_probs = {
|
||||
-0.08146476f, -0.40439552f, 1.0100367f, -0.7724162f, -0.08113786f, -0.36328858f, 0.3688482f, -0.013465762f,
|
||||
-0.32420647f, -0.3815508f, 0.79585606f, 0.14430691f, -0.21869831f, 0.11483674f, -0.11992836f, 0.35216537f};
|
||||
const std::vector<float> fc1_experts_weights = {
|
||||
0.81960344f, 0.9296998f, 0.45050132f, 0.38805157f, 0.50729614f, 0.47014588f, 0.62020564f, 0.6401168f,
|
||||
0.045871615f, 0.31548113f, 0.92106473f, 0.6947775f, 0.4751312f, 0.19854712f, 0.19409746f, 0.052116573f,
|
||||
0.3370188f, 0.6688521f, 0.8188108f, 0.73084867f, 0.058027983f, 0.19931877f, 0.42109168f, 0.98367476f,
|
||||
0.57232875f, 0.37051463f, 0.7068576f, 0.30955923f, 0.17637217f, 0.8649436f, 0.2726491f, 0.39976662f,
|
||||
0.0025978684f, 0.8346353f, 0.8788173f, 0.6822241f, 0.1513629f, 0.0065300465f, 0.093910515f, 0.8728501f,
|
||||
0.7400529f, 0.9207522f, 0.76193494f, 0.6265461f, 0.49510366f, 0.11974698f, 0.07161391f, 0.032325685f,
|
||||
0.704681f, 0.254516f, 0.3993737f, 0.21224737f, 0.40888822f, 0.14808255f, 0.17329216f, 0.6658554f,
|
||||
0.3514018f, 0.8086716f, 0.33959562f, 0.13321638f, 0.41178054f, 0.2576263f, 0.3470292f, 0.024002194f,
|
||||
0.77974546f, 0.15189773f, 0.75130886f, 0.7268921f, 0.85721636f, 0.11647397f, 0.8595984f, 0.2636242f,
|
||||
0.6855346f, 0.96955734f, 0.42948407f, 0.49613327f, 0.38488472f, 0.08250773f, 0.73995143f, 0.003641069f,
|
||||
0.81039995f, 0.87411255f, 0.9728532f, 0.38206023f, 0.08917904f, 0.61241513f, 0.77621365f, 0.0023456216f,
|
||||
0.38650817f, 0.20027226f, 0.45626813f, 0.25389326f, 0.2956162f, 0.34127057f, 0.024847984f, 0.91025376f,
|
||||
0.9191656f, 0.42156547f, 0.44305897f, 0.29594004f, 0.04846859f, 0.013427794f, 0.6858292f, 0.22547692f,
|
||||
0.17856151f, 0.4609884f, 0.33349442f, 0.3382396f, 0.5160656f, 0.3939438f, 0.3278438f, 0.26059705f,
|
||||
0.0930863f, 0.9192536f, 0.29990643f, 0.63248974f, 0.32651705f, 0.54063064f, 0.9661502f, 0.73036134f,
|
||||
0.06670016f, 0.6984514f, 0.9746214f, 0.63154167f, 0.83521235f, 0.99294376f, 0.4233855f, 0.6037772f,
|
||||
0.15248245f, 0.39696145f, 0.8702919f, 0.7563229f, 0.18360549f, 0.099057496f, 0.15831816f, 0.00656116f,
|
||||
0.114180505f, 0.3763513f, 0.8374386f, 0.5836911f, 0.11969727f, 0.09888804f, 0.74873763f, 0.12807935f,
|
||||
0.43843627f, 0.739853f, 0.26859397f, 0.44548005f, 0.45647776f, 0.38170832f, 0.24648392f, 0.054280818f,
|
||||
0.0958215f, 0.23226917f, 0.98291886f, 0.25849265f, 0.16423601f, 0.6211971f, 0.63780516f, 0.77395487f,
|
||||
0.8800602f, 0.7784371f, 0.004249513f, 0.5443443f, 0.80287653f, 0.45378727f, 0.20536041f, 0.9766699f,
|
||||
0.31298608f, 0.21532774f, 0.04922247f, 0.52233416f, 0.72156656f, 0.6106814f, 0.59887487f, 0.12080628f,
|
||||
0.03305638f, 0.5088047f, 0.95591706f, 0.7884607f, 0.20888287f, 0.43509573f, 0.13140821f, 0.2587883f,
|
||||
0.5905492f, 0.77226925f, 0.91418463f, 0.04094696f, 0.8343076f, 0.14735395f, 0.6872336f, 0.92312264f,
|
||||
0.5070212f, 0.9549045f, 0.07397425f, 0.3090204f, 0.79162645f, 0.39106607f, 0.39764988f, 0.29160416f,
|
||||
0.84465307f, 0.7452516f, 0.66022503f, 0.21901816f, 0.09412521f, 0.5540803f, 0.6481394f, 0.26914406f,
|
||||
0.36010116f, 0.83768386f, 0.53982985f, 0.52255917f, 0.37694973f, 0.04720515f, 0.029871285f, 0.26099247f,
|
||||
0.2458393f, 0.6557768f, 0.35444462f, 0.30438894f, 0.9767149f, 0.67416143f, 0.85645115f, 0.25794363f,
|
||||
0.2957666f, 0.68377024f, 0.16686243f, 0.17314798f, 0.47585016f, 0.31711966f, 0.125171f, 0.7965795f,
|
||||
0.90208143f, 0.58111167f, 0.41294336f, 0.036863506f, 0.31788063f, 0.6272928f, 0.73576546f, 0.43679124f,
|
||||
0.30232358f, 0.77861303f, 0.10180014f, 0.816009f, 0.30602258f, 0.5076527f, 0.40119207f, 0.5606195f,
|
||||
0.3489008f, 0.8635635f, 0.48700142f, 0.89029974f, 0.98074025f, 0.25640452f, 0.13524544f, 0.901151f,
|
||||
0.89180696f, 0.11822635f, 0.46134835f, 0.006936848f, 0.09070045f, 0.59657127f, 0.6330173f, 0.6059905f,
|
||||
0.36391765f, 0.96128887f, 0.571489f, 0.2049576f, 0.4716931f, 0.6200726f, 0.67509633f, 0.14645958f,
|
||||
0.6873948f, 0.24455917f, 0.08452982f, 0.22689629f, 0.9822047f, 0.9274289f, 0.9477422f, 0.7935056f,
|
||||
0.87772477f, 0.43307513f, 0.22488606f, 0.7498283f, 0.24090862f, 0.16256708f, 0.34033298f, 0.6049296f,
|
||||
0.7573983f, 0.3057955f, 0.20571685f, 0.56744653f, 0.2052834f, 0.17446929f, 0.76062596f, 0.4160077f,
|
||||
0.9568925f, 0.9863913f, 0.64955276f, 0.67207885f, 0.61514187f, 0.50783044f, 0.46363378f, 0.50687206f,
|
||||
0.6867124f, 0.9648854f, 0.37042046f, 0.2886421f, 0.37891757f, 0.25843787f, 0.58501935f, 0.8732242f,
|
||||
0.8909887f, 0.72956276f, 0.13203424f, 0.23164761f, 0.3901443f, 0.40783793f, 0.54112387f, 0.041014254f,
|
||||
0.65562236f, 0.11856395f, 0.18362767f, 0.08430874f, 0.9356598f, 0.026530087f, 0.8771834f, 0.48319155f,
|
||||
0.4418506f, 0.81273925f, 0.4537862f, 0.81357706f, 0.8615075f, 0.06589496f, 0.692392f, 0.5943895f,
|
||||
0.60750586f, 0.5729957f, 0.6367655f, 0.2594666f, 0.43602943f, 0.97506f, 0.83592474f, 0.48121578f,
|
||||
0.029734552f, 0.5219139f, 0.15951324f, 0.90659577f, 0.19645631f, 0.4638992f, 0.38902867f, 0.5889769f,
|
||||
0.9705138f, 0.5475096f, 0.789582f, 0.8881108f, 0.9036556f, 0.32732427f, 0.38817167f, 0.7409689f,
|
||||
0.36356616f, 0.734132f, 0.39076614f, 0.16087383f, 0.70352167f, 0.576659f, 0.7229242f, 0.996743f,
|
||||
0.84136647f, 0.97399056f, 0.5267614f, 0.06989372f, 0.14923638f, 0.18941313f, 0.059375823f, 0.24937624f,
|
||||
0.039716125f, 0.038692355f, 0.20122272f, 0.0070830584f, 0.19309378f, 0.69065434f, 0.9170264f, 0.3512686f,
|
||||
0.3545606f, 0.76697665f, 0.25331455f, 0.26358372f, 0.80806476f, 0.064349174f, 0.5611374f, 0.941691f,
|
||||
0.58574325f, 0.6359719f, 0.20880443f, 0.49310172f, 0.5274922f, 0.62271714f, 0.694273f, 0.9344639f,
|
||||
0.11835027f, 0.51498765f, 0.25018185f, 0.10446805f, 0.45996118f, 0.059881568f, 0.8489496f, 0.5579074f,
|
||||
0.23052096f, 0.76128954f, 0.02678603f, 0.3066004f, 0.40259063f, 0.07512486f, 0.18205583f, 0.4183907f,
|
||||
0.8793823f, 0.9828271f, 0.8181312f, 0.20143801f, 0.17288941f, 0.9363466f, 0.6768587f, 0.51328385f,
|
||||
0.56766605f, 0.098151624f, 0.33305728f, 0.98130906f, 0.3766839f, 0.47491795f, 0.08483446f, 0.22029644f,
|
||||
0.4897902f, 0.18942028f, 0.4379952f, 0.7034796f, 0.0109113455f, 0.64850605f, 0.16939592f, 0.25597447f,
|
||||
0.69195485f, 0.8975601f, 0.36334568f, 0.29471546f, 0.04788208f, 0.24217117f, 0.062181532f, 0.38556474f,
|
||||
0.6020277f, 0.03156215f, 0.93655676f, 0.81369543f, 0.010527074f, 0.2611835f, 0.6630776f, 0.3972702f,
|
||||
0.44551176f, 0.27424216f, 0.9016098f, 0.22050089f, 0.9146384f, 0.53226113f, 0.6005109f, 0.8900659f,
|
||||
0.4176172f, 0.21532834f, 0.4191329f, 0.9055267f, 0.12900633f, 0.6134902f, 0.008604288f, 0.76215106f,
|
||||
0.68473387f, 0.5211961f, 0.71459657f, 0.50056237f, 0.7766764f, 0.10418975f, 0.42657375f, 0.7218073f,
|
||||
0.9979084f, 0.7546957f, 0.1364128f, 0.8845484f, 0.38850087f, 0.39324278f, 0.04554516f, 0.42129284f,
|
||||
0.8536634f, 0.5697224f, 0.20877302f, 0.65390605f, 0.3396778f, 0.956497f, 0.066022694f, 0.34206223f,
|
||||
0.017213225f, 0.3030849f, 0.6576238f, 0.9813073f, 0.58397317f, 0.99017924f, 0.59782606f, 0.788768f,
|
||||
0.9008311f, 0.91796166f, 0.22013813f, 0.959695f, 0.80288273f, 0.2662105f, 0.26139832f, 0.080626905f};
|
||||
const std::vector<float> fc2_experts_weights = {
|
||||
0.6255686f, 0.09472537f, 0.71121234f, 0.65789884f, 0.065598905f, 0.63625044f, 0.45933473f, 0.7284089f,
|
||||
0.7868948f, 0.0029274821f, 0.95854944f, 0.919321f, 0.6989418f, 0.043019474f, 0.32138962f, 0.35509557f,
|
||||
0.37150103f, 0.78196156f, 0.6817853f, 0.89608955f, 0.31273842f, 0.6682699f, 0.6778976f, 0.08370459f,
|
||||
0.014990091f, 0.24055547f, 0.84227383f, 0.029270172f, 0.0647831f, 0.7801003f, 0.7697645f, 0.91119635f,
|
||||
0.12253064f, 0.13405013f, 0.75649333f, 0.9348151f, 0.7991694f, 0.57832605f, 0.66478735f, 0.97456336f,
|
||||
0.17739785f, 0.2729941f, 0.8497335f, 0.15788019f, 0.22429371f, 0.86499554f, 0.65776104f, 0.661535f,
|
||||
0.2880798f, 0.49309975f, 0.9576164f, 0.19988996f, 0.5039311f, 0.73779976f, 0.15482187f, 0.98558843f,
|
||||
0.25019473f, 0.379932f, 0.36471486f, 0.17417055f, 0.009367704f, 0.7819258f, 0.63283706f, 0.031699598f,
|
||||
0.1781866f, 0.994184f, 0.6911175f, 0.7006223f, 0.20085096f, 0.28080195f, 0.42452294f, 0.40856004f,
|
||||
0.15737581f, 0.5411925f, 0.549694f, 0.4366895f, 0.5693159f, 0.3018247f, 0.63012594f, 0.6885702f,
|
||||
0.2366305f, 0.004210472f, 0.7617172f, 0.61926836f, 0.24570602f, 0.981851f, 0.273876f, 0.8378734f,
|
||||
0.75366426f, 0.080795944f, 0.82247066f, 0.040263534f, 0.22299266f, 0.41664255f, 0.16297674f, 0.98845494f,
|
||||
0.39971018f, 0.69859487f, 0.053544044f, 0.7878332f, 0.34460813f, 0.11966437f, 0.5731115f, 0.7422309f,
|
||||
0.93269855f, 0.19460368f, 0.25394785f, 0.59613144f, 0.6356306f, 0.6922361f, 0.7744376f, 0.38662314f,
|
||||
0.7777848f, 0.8686458f, 0.36938924f, 0.8557286f, 0.74428976f, 0.9410264f, 0.21586305f, 0.2530955f,
|
||||
0.35543054f, 0.52536315f, 0.8000995f, 0.21456867f, 0.750327f, 0.3208093f, 0.80205464f, 0.47626138f,
|
||||
0.061956525f, 0.22487706f, 0.13812399f, 0.74798125f, 0.1647259f, 0.45834088f, 0.6078779f, 0.22580266f,
|
||||
0.644235f, 0.011788309f, 0.14224577f, 0.0469383f, 0.34876132f, 0.3178513f, 0.5715967f, 0.40754277f,
|
||||
0.735041f, 0.9583977f, 0.67939556f, 0.30301625f, 0.031807184f, 0.68110096f, 0.25227106f, 0.75443816f,
|
||||
0.83424246f, 0.69286025f, 0.9691554f, 0.9748982f, 0.60586995f, 0.13568163f, 0.94672066f, 0.26275212f,
|
||||
0.2638232f, 0.9183893f, 0.88740516f, 0.65107566f, 0.5313419f, 0.07941705f, 0.44809794f, 0.9795632f,
|
||||
0.6273294f, 0.542809f, 0.3961745f, 0.32560885f, 0.79801136f, 0.53083426f, 0.8252871f, 0.4115007f,
|
||||
0.7184546f, 0.70638496f, 0.57973206f, 0.8141865f, 0.81332296f, 0.96346164f, 0.88438797f, 0.37215167f,
|
||||
0.0766899f, 0.5914087f, 0.49563587f, 0.3695873f, 0.41627264f, 0.5235164f, 0.86481494f, 0.6558706f,
|
||||
0.32245284f, 0.29438752f, 0.37618434f, 0.3067485f, 0.9496114f, 0.76482266f, 0.95148784f, 0.5015968f,
|
||||
0.60083544f, 0.67338234f, 0.026723444f, 0.5446483f, 0.466555f, 0.21967298f, 0.112026334f, 0.9426372f,
|
||||
0.906533f, 0.73173434f, 0.97712487f, 0.29709607f, 0.41363865f, 0.6893093f, 0.4173867f, 0.4018826f,
|
||||
0.086719275f, 0.63433063f, 0.1978364f, 0.5181831f, 0.9874878f, 0.34609234f, 0.34240413f, 0.8016564f,
|
||||
0.31617337f, 0.4570613f, 0.96686924f, 0.29501313f, 0.14229488f, 0.22017813f, 0.36137718f, 0.26275063f,
|
||||
0.24053413f, 0.70197225f, 0.58496886f, 0.33996922f, 0.11154431f, 0.34257007f, 0.28898042f, 0.33729053f,
|
||||
0.048938513f, 0.60771453f, 0.13263822f, 0.11060041f, 0.091483414f, 0.70869184f, 0.19898665f, 0.29362458f,
|
||||
0.8919203f, 0.7654821f, 0.7866956f, 0.02524674f, 0.1414501f, 0.3112445f, 0.9130488f, 0.5511502f,
|
||||
0.12605143f, 0.5031309f, 0.11166459f, 0.39045036f, 0.36251247f, 0.9328308f, 0.65486836f, 0.41281444f,
|
||||
0.5844644f, 0.35566723f, 0.6964502f, 0.6977819f, 0.63427305f, 0.30511153f, 0.92657536f, 0.42781502f,
|
||||
0.30534166f, 0.813157f, 0.90752834f, 0.9975799f, 0.64812917f, 0.32955307f, 0.753946f, 0.92897725f,
|
||||
0.009582937f, 0.43805653f, 0.15901726f, 0.5931799f, 0.7067924f, 0.39670604f, 0.45817143f, 0.7250554f,
|
||||
0.41596514f, 0.08011025f, 0.900068f, 0.24834275f, 0.44507074f, 0.5471632f, 0.46995157f, 0.029657006f,
|
||||
0.7294f, 0.27288425f, 0.2406702f, 0.6194577f, 0.23906898f, 0.26892018f, 0.33152503f, 0.3121612f,
|
||||
0.29118127f, 0.36515707f, 0.6299379f, 0.095391035f, 0.19735986f, 0.5072957f, 0.56953406f, 0.77614623f,
|
||||
0.14877802f, 0.65959847f, 0.7841949f, 0.7776301f, 0.03428924f, 0.3091979f, 0.07021719f, 0.18359429f,
|
||||
0.77849144f, 0.42534047f, 0.7123557f, 0.20649683f, 0.57597995f, 0.19757104f, 0.749946f, 0.2813105f,
|
||||
0.37462044f, 0.06618434f, 0.50165176f, 0.9747401f, 0.7426891f, 0.23322952f, 0.50672436f, 0.44517577f,
|
||||
0.09746289f, 0.89204556f, 0.50806034f, 0.6052985f, 0.2980855f, 0.26604044f, 0.5824448f, 0.68485546f,
|
||||
0.612149f, 0.25902748f, 0.9854489f, 0.4263978f, 0.19379246f, 0.26614368f, 0.9922104f, 0.5000241f,
|
||||
0.4321279f, 0.2919191f, 0.3689273f, 0.078885734f, 0.10265827f, 0.79264474f, 0.9277247f, 0.9771502f,
|
||||
0.13902885f, 0.77043164f, 0.19051671f, 0.7982801f, 0.86077714f, 0.8869355f, 0.86002564f, 0.81278664f,
|
||||
0.5097318f, 0.7297412f, 0.32111454f, 0.7177174f, 0.33929902f, 0.49160433f, 0.064810574f, 0.3692627f,
|
||||
0.23706353f, 0.3313396f, 0.18070674f, 0.05027789f, 0.53255826f, 0.8244896f, 0.9553747f, 0.7917771f,
|
||||
0.24083132f, 0.005495131f, 0.6896569f, 0.78015697f, 0.07074398f, 0.67929304f, 0.9227386f, 0.5302883f,
|
||||
0.19877058f, 0.90993816f, 0.71350795f, 0.8311006f, 0.16185725f, 0.79097277f, 0.15846318f, 0.99474716f,
|
||||
0.28815013f, 0.80128354f, 0.6001208f, 0.63250524f, 0.4233225f, 0.7053677f, 0.29161406f, 0.028710365f,
|
||||
0.30789846f, 0.8917693f, 0.36836517f, 0.6571592f, 0.3151368f, 0.8750746f, 0.7992451f, 0.6765068f,
|
||||
0.24441916f, 0.091435075f, 0.5188247f, 0.20667112f, 0.9110969f, 0.019512117f, 0.72343415f, 0.998457f,
|
||||
0.7504142f, 0.6704894f, 0.01892668f, 0.9809466f, 0.41447622f, 0.032795787f, 0.9935814f, 0.29653466f,
|
||||
0.4646262f, 0.95763975f, 0.15339965f, 0.14625502f, 0.58130866f, 0.43307304f, 0.6151709f, 0.08064735f,
|
||||
0.5149533f, 0.27762014f, 0.25419557f, 0.04218155f, 0.7651092f, 0.59631824f, 0.077278376f, 0.89677596f,
|
||||
0.6508104f, 0.5927816f, 0.2064318f, 0.57540226f, 0.9817701f, 0.84294224f, 0.11056489f, 0.9564106f,
|
||||
0.5387549f, 0.74048257f, 0.88833815f, 0.9262546f, 0.11023259f, 0.93783194f, 0.16041255f, 0.53748304f,
|
||||
0.1506182f, 0.39038336f, 0.47727865f, 0.44018233f, 0.42101204f, 0.53943527f, 0.99320936f, 0.79050577f,
|
||||
0.77973497f, 0.7001237f, 0.88709056f, 0.4769255f, 0.5397561f, 0.60289854f, 0.06393474f, 0.09722155f,
|
||||
0.5613007f, 0.30437487f, 0.49082512f, 0.3852706f, 0.5778314f, 0.8253078f, 0.33417904f, 0.9004303f,
|
||||
0.8947809f, 0.11625093f, 0.11388689f, 0.09546256f, 0.22598988f, 0.30536187f, 0.46236527f, 0.3784039f,
|
||||
0.24737573f, 0.3411532f, 0.31912774f, 0.9905191f, 0.31468558f, 0.14199954f, 0.7078488f, 0.47111923f,
|
||||
0.882782f, 0.8124163f, 0.9593644f, 0.13382024f, 0.8214317f, 0.9196194f, 0.25308424f, 0.95958996f};
|
||||
const std::vector<float> fc1_experts_bias = {
|
||||
0.8748215f, 0.5054756f, 0.74107623f, 0.32518923f, 0.0639081f, 0.62639004f, 0.64906263f, 0.17322052f,
|
||||
0.7424998f, 0.07288867f, 0.93031204f, 0.9841952f, 0.6361292f, 0.18628561f, 0.7433356f, 0.5852079f,
|
||||
0.6359594f, 0.66432667f, 0.88067776f, 0.28508204f, 0.38752747f, 0.63635296f, 0.55448055f, 0.9031888f,
|
||||
0.23738074f, 0.48179168f, 0.5934266f, 0.3672055f, 0.84085834f, 0.5546908f, 0.03788501f, 0.44583207f,
|
||||
0.27322155f, 0.5485856f, 0.44189203f, 0.00403291f, 0.40888733f, 0.45211035f, 0.35256076f, 0.9593902f,
|
||||
0.39090043f, 0.8212086f, 0.62385887f, 0.07793343f, 0.61749303f, 0.9143678f, 0.17294967f, 0.17681253f,
|
||||
0.9894245f, 0.901755f, 0.221053f, 0.8008725f, 0.43603396f, 0.007035315f, 0.5375667f, 0.661547f,
|
||||
0.35001957f, 0.67394173f, 0.072449565f, 0.84650797f, 0.92626715f, 0.77573335f, 0.58474565f, 0.66467446f};
|
||||
const std::vector<float> fc2_experts_bias = {
|
||||
0.13822609f, 0.3750633f, 0.45226622f, 0.22175694f, 0.13068998f, 0.8363088f, 0.8393226f, 0.045905888f,
|
||||
0.65910596f, 0.7034011f, 0.97498417f, 0.78927684f, 0.95966834f, 0.33630514f, 0.8501932f, 0.9067007f,
|
||||
0.027835965f, 0.09864664f, 0.6012027f, 0.7730189f, 0.25159347f, 0.55506724f, 0.49927413f, 0.62655383f,
|
||||
0.23132521f, 0.7820195f, 0.8325047f, 0.15307087f, 0.5048437f, 0.5013873f, 0.66055787f, 0.96579224f};
|
||||
const std::vector<float> output = {
|
||||
1.3775184f, 2.0985768f, 2.091839f, 2.9706357f, 1.9404914f, 1.9915576f, 2.3302228f, 2.3702593f,
|
||||
0.51896286f, 0.7936432f, 0.9944805f, 1.3225251f, 0.73894113f, 0.87975955f, 1.0468717f, 1.1585085f,
|
||||
0.012911659f, 0.045757107f, 0.27884653f, 0.3585817f, 0.116771236f, 0.25755364f, 0.23161705f, 0.2906256f,
|
||||
4.8571277f, 5.649453f, 5.485141f, 5.306299f, 4.767025f, 6.9010167f, 5.3520975f, 6.711155f};
|
||||
|
||||
RunMoETest(input,
|
||||
router_probs,
|
||||
fc1_experts_weights,
|
||||
fc2_experts_weights,
|
||||
fc1_experts_bias,
|
||||
fc2_experts_bias,
|
||||
output,
|
||||
num_rows,
|
||||
num_experts,
|
||||
hidden_size,
|
||||
inter_size,
|
||||
"relu");
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
431
onnxruntime/test/python/transformers/test_parity_moe.py
Normal file
431
onnxruntime/test/python/transformers/test_parity_moe.py
Normal file
|
|
@ -0,0 +1,431 @@
|
|||
# --------------------------------------------------------------------------
|
||||
# Copyright 2020 The HuggingFace Inc. team
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0
|
||||
# --------------------------------------------------------------------------
|
||||
# Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
# Licensed under the MIT License. See License.txt in the project root for
|
||||
# license information.
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
import time
|
||||
import unittest
|
||||
|
||||
import numpy
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from onnx import TensorProto, helper
|
||||
|
||||
import onnxruntime
|
||||
|
||||
torch.manual_seed(42)
|
||||
numpy.random.seed(42)
|
||||
|
||||
|
||||
ORT_DTYPE = TensorProto.FLOAT16
|
||||
NP_TYPE = numpy.float16 if ORT_DTYPE == TensorProto.FLOAT16 else numpy.float32
|
||||
THRESHOLD = 3e-2
|
||||
|
||||
|
||||
def value_string_of(numpy_array):
|
||||
arr = numpy_array.flatten()
|
||||
lines = ["f, ".join([str(v) for v in arr[i : min(i + 8, arr.size)]]) for i in range(0, arr.size, 8)]
|
||||
return "{\n " + "f,\n ".join(lines) + "f}"
|
||||
|
||||
|
||||
def print_tensor(name, numpy_array):
|
||||
print(f"const std::vector<float> {name} = {value_string_of(numpy_array)};")
|
||||
|
||||
|
||||
def create_moe_onnx_graph(
|
||||
num_rows,
|
||||
num_experts,
|
||||
hidden_size,
|
||||
inter_size,
|
||||
fc1_experts_weights,
|
||||
fc2_experts_weights,
|
||||
fc1_experts_bias,
|
||||
fc2_experts_bias,
|
||||
):
|
||||
nodes = [
|
||||
helper.make_node(
|
||||
"MoE",
|
||||
[
|
||||
"input",
|
||||
"router_probs",
|
||||
"fc1_experts_weights",
|
||||
"fc2_experts_weights",
|
||||
"fc1_experts_bias",
|
||||
"fc2_experts_bias",
|
||||
],
|
||||
["output"],
|
||||
"MoE_0",
|
||||
k=1,
|
||||
activation_type="gelu",
|
||||
domain="com.microsoft",
|
||||
),
|
||||
]
|
||||
|
||||
fc1_shape = [num_experts, hidden_size, inter_size]
|
||||
fc2_shape = [num_experts, inter_size, hidden_size]
|
||||
|
||||
torch_type = torch.float16 if ORT_DTYPE == TensorProto.FLOAT16 else torch.float32
|
||||
|
||||
initializers = [
|
||||
helper.make_tensor(
|
||||
"fc1_experts_weights",
|
||||
ORT_DTYPE,
|
||||
fc1_shape,
|
||||
fc1_experts_weights.to(torch_type).flatten().tolist(),
|
||||
raw=False,
|
||||
),
|
||||
helper.make_tensor(
|
||||
"fc2_experts_weights",
|
||||
ORT_DTYPE,
|
||||
fc2_shape,
|
||||
fc2_experts_weights.to(torch_type).flatten().tolist(),
|
||||
raw=False,
|
||||
),
|
||||
]
|
||||
|
||||
fc1_bias_shape = [num_experts, inter_size]
|
||||
fc2_bias_shape = [num_experts, hidden_size]
|
||||
initializers.extend(
|
||||
[
|
||||
helper.make_tensor(
|
||||
"fc1_experts_bias",
|
||||
ORT_DTYPE,
|
||||
fc1_bias_shape,
|
||||
fc1_experts_bias.to(torch_type).flatten().tolist(),
|
||||
raw=False,
|
||||
),
|
||||
helper.make_tensor(
|
||||
"fc2_experts_bias",
|
||||
ORT_DTYPE,
|
||||
fc2_bias_shape,
|
||||
fc2_experts_bias.to(torch_type).flatten().tolist(),
|
||||
raw=False,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
graph_inputs = [
|
||||
helper.make_tensor_value_info("input", ORT_DTYPE, [num_rows, hidden_size]),
|
||||
]
|
||||
|
||||
graph_inputs.append(
|
||||
helper.make_tensor_value_info(
|
||||
"router_probs",
|
||||
ORT_DTYPE,
|
||||
[num_rows, num_experts],
|
||||
)
|
||||
)
|
||||
|
||||
graph_outputs = [
|
||||
helper.make_tensor_value_info("output", ORT_DTYPE, [num_rows, hidden_size]),
|
||||
]
|
||||
|
||||
graph = helper.make_graph(
|
||||
nodes,
|
||||
"MoE_Graph",
|
||||
graph_inputs,
|
||||
graph_outputs,
|
||||
initializers,
|
||||
)
|
||||
|
||||
model = helper.make_model(graph)
|
||||
return model.SerializeToString()
|
||||
|
||||
|
||||
def get_activation_fn(activation):
|
||||
if activation == "relu":
|
||||
return nn.ReLU
|
||||
elif activation == "gelu":
|
||||
return nn.GELU
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class MoEGate(nn.Module):
|
||||
def __init__(self, num_experts, in_features):
|
||||
super().__init__()
|
||||
self.wg_reduction = torch.nn.Linear(in_features, 16, bias=False)
|
||||
|
||||
wg = torch.empty(num_experts, 16)
|
||||
torch.nn.init.orthogonal_(wg, gain=0.32)
|
||||
self.register_parameter("wg", torch.nn.Parameter(wg))
|
||||
|
||||
def forward(self, input):
|
||||
input = self.wg_reduction(input)
|
||||
with torch.no_grad():
|
||||
wg_norm = self.wg.norm(p=2.0, dim=1, keepdim=True)
|
||||
self.wg.mul_(1.5 / wg_norm)
|
||||
logits = self._cosine(input, self.wg)
|
||||
return logits
|
||||
|
||||
def _cosine(self, mat1, mat2, eps=1e-4):
|
||||
assert mat1.dim() == 2
|
||||
assert mat2.dim() == 2
|
||||
|
||||
mat2 = F.normalize(mat2.float(), p=2.0, dim=1, eps=eps)
|
||||
return mat1.float().matmul(mat2.transpose(0, 1)).type_as(mat1)
|
||||
|
||||
|
||||
class MoERuntimeExperts(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
num_experts,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
drop=0.0,
|
||||
bias=True,
|
||||
chunk_size=-1,
|
||||
):
|
||||
super().__init__()
|
||||
# assert bias is False, "Current bias is not supported"
|
||||
assert drop == 0.0, "Current drop is not supported"
|
||||
assert chunk_size == -1, "Current chunk is not supported"
|
||||
|
||||
self.weight1 = nn.Parameter(torch.rand(num_experts, in_features, hidden_features))
|
||||
self.weight2 = nn.Parameter(torch.rand(num_experts, hidden_features, out_features))
|
||||
|
||||
self.bias1 = nn.Parameter(torch.rand(num_experts, hidden_features)) if bias else None
|
||||
self.bias2 = nn.Parameter(torch.rand(num_experts, in_features)) if bias else None
|
||||
|
||||
self.act = act_layer()
|
||||
|
||||
def forward(self, x, indices_s):
|
||||
x = x.unsqueeze(1)
|
||||
x = self.bmm(x, self.weight1, indices_s)
|
||||
if self.bias1 is not None:
|
||||
x = x + self.bias1[indices_s].unsqueeze(1) # S x hidden_features
|
||||
x = self.act(x)
|
||||
x = self.bmm(x, self.weight2, indices_s)
|
||||
if self.bias2 is not None:
|
||||
x = x + self.bias2[indices_s].unsqueeze(1) # S x 1 x in_features
|
||||
return x
|
||||
|
||||
def bmm(self, x, weight, indices_s):
|
||||
x = torch.bmm(x, weight[indices_s]) # S x 1 x hidden_features
|
||||
return x
|
||||
|
||||
|
||||
class MoE(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
batch_size,
|
||||
num_rows,
|
||||
num_experts,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
eval_capacity=-1,
|
||||
activation="gelu",
|
||||
):
|
||||
super().__init__()
|
||||
self.num_experts = num_experts
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.eval_capacity = eval_capacity # -1 means we route all tokens
|
||||
|
||||
self.gate = MoEGate(num_experts=num_experts, in_features=in_features)
|
||||
self.moe_experts = MoERuntimeExperts(
|
||||
num_experts=num_experts,
|
||||
in_features=in_features,
|
||||
hidden_features=hidden_features,
|
||||
out_features=out_features,
|
||||
act_layer=get_activation_fn(activation),
|
||||
bias=True,
|
||||
)
|
||||
|
||||
self.moe_onnx_graph = create_moe_onnx_graph(
|
||||
batch_size * num_rows,
|
||||
num_experts,
|
||||
in_features,
|
||||
hidden_features,
|
||||
self.moe_experts.weight1,
|
||||
self.moe_experts.weight2,
|
||||
self.moe_experts.bias1,
|
||||
self.moe_experts.bias2,
|
||||
)
|
||||
|
||||
self.ort_sess = self.create_ort_session()
|
||||
|
||||
self.torch_input = torch.randn(batch_size, num_rows, in_features)
|
||||
|
||||
def create_ort_session(self):
|
||||
from onnxruntime import InferenceSession, SessionOptions
|
||||
|
||||
sess_options = SessionOptions()
|
||||
|
||||
cuda_providers = ["CUDAExecutionProvider"]
|
||||
if cuda_providers[0] not in onnxruntime.get_available_providers():
|
||||
return None
|
||||
|
||||
sess_options.log_severity_level = 2
|
||||
ort_session = InferenceSession(self.moe_onnx_graph, sess_options, providers=["CUDAExecutionProvider"])
|
||||
|
||||
return ort_session
|
||||
|
||||
def ort_run_with_iobinding(self, ort_inputs, repeat=1000):
|
||||
iobinding = self.ort_sess.io_binding()
|
||||
device_id = torch.cuda.current_device()
|
||||
|
||||
iobinding.bind_input(
|
||||
name="input",
|
||||
device_type="cuda",
|
||||
device_id=device_id,
|
||||
element_type=NP_TYPE,
|
||||
shape=ort_inputs["input"].shape,
|
||||
buffer_ptr=onnxruntime.OrtValue.ortvalue_from_numpy(ort_inputs["input"], "cuda", device_id).data_ptr(),
|
||||
)
|
||||
iobinding.bind_input(
|
||||
name="router_probs",
|
||||
device_type="cuda",
|
||||
device_id=device_id,
|
||||
element_type=NP_TYPE,
|
||||
shape=ort_inputs["router_probs"].shape,
|
||||
buffer_ptr=onnxruntime.OrtValue.ortvalue_from_numpy(
|
||||
ort_inputs["router_probs"], "cuda", device_id
|
||||
).data_ptr(),
|
||||
)
|
||||
|
||||
iobinding.synchronize_inputs()
|
||||
|
||||
iobinding.bind_output(
|
||||
name="output",
|
||||
device_type="cuda",
|
||||
device_id=device_id,
|
||||
element_type=NP_TYPE,
|
||||
shape=ort_inputs["input"].shape,
|
||||
buffer_ptr=onnxruntime.OrtValue.ortvalue_from_numpy(
|
||||
numpy.zeros(ort_inputs["input"].shape), "cuda", device_id
|
||||
).data_ptr(),
|
||||
)
|
||||
iobinding.synchronize_outputs()
|
||||
|
||||
s = time.time()
|
||||
for _ in range(repeat):
|
||||
self.ort_sess.run_with_iobinding(iobinding)
|
||||
e = time.time()
|
||||
print(f"MoE cuda kernel time: {(e - s) / repeat * 1000} ms")
|
||||
|
||||
def torch_forward(self):
|
||||
x = self.torch_input
|
||||
|
||||
b, t, c = x.shape
|
||||
x = x.reshape(-1, c)
|
||||
logits = self.gate(x)
|
||||
gates = torch.nn.functional.softmax(logits, dim=1)
|
||||
ret = torch.max(gates, dim=1)
|
||||
indices_s = ret.indices # dim: [bs], the index of the expert with highest softmax value
|
||||
scores = ret.values.unsqueeze(-1).unsqueeze(-1) # S
|
||||
x = self.moe_experts(x, indices_s)
|
||||
|
||||
x = x * scores
|
||||
x = x.reshape(b * t, c)
|
||||
|
||||
return x, torch.sum(x)
|
||||
|
||||
def onnx_forward(self, iobinding=False):
|
||||
x = self.torch_input
|
||||
|
||||
_, _, c = x.shape
|
||||
y = x.reshape(-1, c)
|
||||
logits = self.gate(y)
|
||||
|
||||
ort_inputs = {
|
||||
"input": numpy.ascontiguousarray(y.detach().numpy().astype(NP_TYPE)),
|
||||
"router_probs": numpy.ascontiguousarray(logits.detach().numpy().astype(NP_TYPE)),
|
||||
}
|
||||
|
||||
ort_output = None
|
||||
if self.ort_sess is not None:
|
||||
if not iobinding:
|
||||
ort_output = self.ort_sess.run(None, ort_inputs)
|
||||
else:
|
||||
self.ort_run_with_iobinding(ort_inputs)
|
||||
return None
|
||||
|
||||
# print_tensor("input", ort_inputs["input"])
|
||||
# print_tensor("router_probs", ort_inputs["router_probs"])
|
||||
# print_tensor("fc1_experts_weights", self.moe_experts.weight1.detach().numpy())
|
||||
# print_tensor("fc2_experts_weights", self.moe_experts.weight2.detach().numpy())
|
||||
# print_tensor("fc1_experts_bias", self.moe_experts.bias1.detach().numpy())
|
||||
# print_tensor("fc2_experts_bias", self.moe_experts.bias2.detach().numpy())
|
||||
# print_tensor("output", ort_output[0])
|
||||
|
||||
return ort_output
|
||||
|
||||
def parity_check(self):
|
||||
torch_out = self.torch_forward()
|
||||
ort_out = self.onnx_forward()
|
||||
if ort_out is not None:
|
||||
# print("max diff", numpy.max(numpy.abs(torch_out[0].detach().numpy() - ort_out[0])))
|
||||
assert numpy.allclose(torch_out[0].detach().numpy(), ort_out[0], rtol=THRESHOLD, atol=THRESHOLD)
|
||||
|
||||
def benchmark(self):
|
||||
self.onnx_forward(iobinding=True)
|
||||
|
||||
|
||||
class TestMoE(unittest.TestCase):
|
||||
def test_moe_small(self):
|
||||
rt = MoE(
|
||||
batch_size=2,
|
||||
num_rows=8,
|
||||
num_experts=4,
|
||||
in_features=16,
|
||||
hidden_features=32,
|
||||
out_features=16,
|
||||
)
|
||||
rt.parity_check()
|
||||
|
||||
@pytest.mark.slow
|
||||
def test_moe_large(self):
|
||||
for batch_size in [1, 8]:
|
||||
for num_rows in [16, 64]:
|
||||
for num_experts in [16, 64]:
|
||||
for in_features in [256]:
|
||||
for hidden_features in [512]:
|
||||
print(
|
||||
f"batch_size={batch_size}, num_rows={num_rows}, num_experts={num_experts}, in_features={in_features}, hidden_features={hidden_features}"
|
||||
)
|
||||
rt = MoE(
|
||||
batch_size=batch_size,
|
||||
num_rows=num_rows,
|
||||
num_experts=num_experts,
|
||||
in_features=in_features,
|
||||
hidden_features=hidden_features,
|
||||
out_features=in_features,
|
||||
)
|
||||
rt.parity_check()
|
||||
|
||||
@pytest.mark.slow
|
||||
def test_moe_benchmark(self):
|
||||
for batch_size in [32, 64]:
|
||||
for num_rows in [128, 512]:
|
||||
for num_experts in [64, 128]:
|
||||
for in_features in [256, 512]:
|
||||
for hidden_features in [1024, 2048]:
|
||||
print(
|
||||
f"batch_size={batch_size}, num_rows={num_rows}, num_experts={num_experts}, in_features={in_features}, hidden_features={hidden_features}"
|
||||
)
|
||||
rt = MoE(
|
||||
batch_size=batch_size,
|
||||
num_rows=num_rows,
|
||||
num_experts=num_experts,
|
||||
in_features=in_features,
|
||||
hidden_features=hidden_features,
|
||||
out_features=in_features,
|
||||
)
|
||||
rt.benchmark()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -11,7 +11,7 @@ steps:
|
|||
packageType: upack
|
||||
feed: '/7424c8e4-5c62-490e-95c4-79446f31017c'
|
||||
definition: '517c4f6f-5437-4392-a70d-4f15ec5be2f0'
|
||||
version: 1.0.117
|
||||
version: 1.0.118
|
||||
downloadPath: $(Build.BinariesDirectory)/deps
|
||||
|
||||
# The private ADO project
|
||||
|
|
@ -22,7 +22,7 @@ steps:
|
|||
packageType: upack
|
||||
feed: '/4c7631f5-24c0-4307-8822-1aa8f180c325'
|
||||
definition: 'fd9dd5ad-b73e-4678-890e-edcf680dbc1a'
|
||||
version: 1.0.117
|
||||
version: 1.0.118
|
||||
downloadPath: $(Build.BinariesDirectory)/deps
|
||||
|
||||
# You can add more ADO accounts at here.
|
||||
|
|
|
|||
Loading…
Reference in a new issue