diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index c6f0f4e406..96a3ae394b 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -29,6 +29,7 @@ Do not modify directly.* * com.microsoft.FusedMatMul * com.microsoft.GatherND * com.microsoft.Gelu + * com.microsoft.GemmFastGelu * com.microsoft.GridSample * com.microsoft.Inverse * com.microsoft.Irfft @@ -1512,6 +1513,40 @@ This version of the operator has been available since version 1 of the 'com.micr +### **com.microsoft.GemmFastGelu** + + It's a fusion of MatMul and FastGelu. + +#### Version + +This version of the operator has been available since version 1 of the 'com.microsoft' operator set. + +#### Inputs (2 - 3) + +
+
X : T
+
input tensor
+
W : T
+
input tensor
+
bias (optional) : T
+
bias tensor
+
+ +#### Outputs + +
+
Y : T
+
output tensor
+
+ +#### Type Constraints + +
+
T : tensor(float), tensor(float16), tensor(bfloat16)
+
Constrain input and output types to float or half tensors.
+
+ + ### **com.microsoft.GridSample** Given an `input` and a flow-field `grid`, computes the `output` using `input` values and pixel locations from `grid`. diff --git a/onnxruntime/contrib_ops/rocm/bert/gemm_fast_gelu.cc b/onnxruntime/contrib_ops/rocm/bert/gemm_fast_gelu.cc new file mode 100644 index 0000000000..995b4eaf8b --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/bert/gemm_fast_gelu.cc @@ -0,0 +1,87 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/rocm/bert/gemm_fast_gelu.h" + +#include "contrib_ops/rocm/bert/fast_gelu_impl.h" +#include "contrib_ops/rocm/bert/transformer_common.h" +#include "core/providers/cpu/math/matmul_helper.h" +#include "core/providers/rocm/math/matmul_impl.h" +#include "core/providers/rocm/rocm_common.h" +#include "core/providers/rocm/shared_inc/fpgeneric.h" + +namespace onnxruntime { +namespace contrib { +namespace rocm { + +using onnxruntime::rocm::MatMulImpl; +using onnxruntime::rocm::ToHipType; + +#define REGISTER_KERNEL_TYPED(T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + GemmFastGelu, \ + kMSDomain, \ + 1, \ + T, \ + kRocmExecutionProvider, \ + (*KernelDefBuilder::Create()) \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()), \ + GemmFastGelu); + +REGISTER_KERNEL_TYPED(float) +REGISTER_KERNEL_TYPED(MLFloat16) +REGISTER_KERNEL_TYPED(BFloat16) + +template +Status GemmFastGelu::ComputeInternal(OpKernelContext* ctx) const { + typedef typename ToHipType::MappedType HipT; + + const auto* X = ctx->Input(0); + const auto* W = ctx->Input(1); + const auto* bias = ctx->Input(2); + + bool transa = false; + bool transb = false; + bool trans_batch_a = false; + bool trans_batch_b = false; + + MatMulComputeHelper helper; + ORT_RETURN_IF_ERROR(helper.Compute(X->Shape(), W->Shape(), transa, transb, trans_batch_a, trans_batch_b, false)); + + auto gemm_buffer = GetScratchBuffer(helper.OutputShape().Size()); + Tensor* Y = ctx->Output(0, helper.OutputShape()); + + // Bail out early if the output is going to be empty + if (Y->Shape().Size() == 0) + return Status::OK(); + + const float alpha = 1.0f; + const float zero = 0.0f; + + if (MatMulImpl(this, helper, reinterpret_cast(X->template Data()), + reinterpret_cast(W->template Data()), + reinterpret_cast(gemm_buffer.get()), + X->Shape(), W->Shape(), + transa, transb, trans_batch_a, trans_batch_b, alpha, zero) != Status::OK()) { + return Status(common::ONNXRUNTIME, common::FAIL); + } + + int64_t fast_gelu_input_length = Y->Shape().Size(); + int64_t bias_length = (nullptr == bias) ? 0 : bias->Shape().Size(); + + if (!LaunchFastGeluKernel(Stream(), + static_cast(fast_gelu_input_length), + static_cast(bias_length), + reinterpret_cast(gemm_buffer.get()), + (nullptr != bias) ? reinterpret_cast(bias->template Data()) : nullptr, + reinterpret_cast(Y->template MutableData()), + false)) { + HIP_CALL(hipGetLastError()); + return Status(common::ONNXRUNTIME, common::FAIL); + } + return Status::OK(); +} + +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/bert/gemm_fast_gelu.h b/onnxruntime/contrib_ops/rocm/bert/gemm_fast_gelu.h new file mode 100644 index 0000000000..ae4f84fa5f --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/bert/gemm_fast_gelu.h @@ -0,0 +1,23 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once +#include "core/common/common.h" +#include "core/providers/rocm/rocm_kernel.h" + +namespace onnxruntime { +namespace contrib { +namespace rocm { + +using onnxruntime::rocm::RocmKernel; + +template +class GemmFastGelu final : public RocmKernel { + public: + GemmFastGelu(const OpKernelInfo& op_kernel_info) : RocmKernel(op_kernel_info) {} + Status ComputeInternal(OpKernelContext* ctx) const override; +}; + +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/rocm_contrib_kernels.cc b/onnxruntime/contrib_ops/rocm/rocm_contrib_kernels.cc index 7efd5d8524..8baadd1caa 100644 --- a/onnxruntime/contrib_ops/rocm/rocm_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/rocm/rocm_contrib_kernels.cc @@ -95,6 +95,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1 class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, BFloat16, TransposeMatMul); // backward compatibility class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, BFloat16, FusedMatMul); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, BFloat16_float_BFloat16, LayerNormalization); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GemmFastGelu); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GemmFastGelu); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, BFloat16, GemmFastGelu); #ifdef ENABLE_ATEN class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kPytorchAtenDomain, 1, ATen); @@ -198,6 +201,9 @@ Status RegisterRocmContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #ifdef ENABLE_ATEN BuildKernelCreateInfo, diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index f8494bc2b0..1235527e57 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -9,6 +9,14 @@ using namespace ::ONNX_NAMESPACE; +namespace ONNX_NAMESPACE { +void matmulShapeInference( + ONNX_NAMESPACE::InferenceContext& ctx, + int input1Idx, + int input2Idx); + +} // namespace ONNX_NAMESPACE + namespace onnxruntime { namespace contrib { @@ -285,5 +293,23 @@ ONNX_MS_OPERATOR_SET_SCHEMA(BifurcationDetector, 1, // and current tokens length. // tokens_length = cur_tokens_length + bifurcation_index + 1. })); + +constexpr const char* GemmFastGelu_ver1_doc = R"DOC( +It's a fusion of MatMul and FastGelu.)DOC"; + +ONNX_MS_OPERATOR_SET_SCHEMA(GemmFastGelu, 1, + OpSchema() + .SetDoc(GemmFastGelu_ver1_doc) + .Input(0, "X", "input tensor", "T") + .Input(1, "W", "input tensor", "T") + .Input(2, "bias", "bias tensor", "T", OpSchema::Optional) + .Output(0, "Y", "output tensor", "T") + .TypeConstraint("T", {"tensor(float)", "tensor(float16)", "tensor(bfloat16)"}, + "Constrain input and output types to float or half tensors.") + .TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) { + ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 0, 0); + ONNX_NAMESPACE::matmulShapeInference(ctx, 0, 1); + })); + } } // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/graph/contrib_ops/ms_opset.h b/onnxruntime/core/graph/contrib_ops/ms_opset.h index 8e9a698c8f..f907bb77b6 100644 --- a/onnxruntime/core/graph/contrib_ops/ms_opset.h +++ b/onnxruntime/core/graph/contrib_ops/ms_opset.h @@ -74,6 +74,7 @@ class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, TransposeMatMul); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, Trilu); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, Unique); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, WordConvEmbedding); +class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, GemmFastGelu); class OpSet_Microsoft_ver1 { public: @@ -142,6 +143,7 @@ class OpSet_Microsoft_ver1 { fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); + fn(GetOpSchema()); } }; } // namespace contrib diff --git a/onnxruntime/core/providers/rocm/math/matmul.cc b/onnxruntime/core/providers/rocm/math/matmul.cc index 198ff367b8..3003243eb3 100644 --- a/onnxruntime/core/providers/rocm/math/matmul.cc +++ b/onnxruntime/core/providers/rocm/math/matmul.cc @@ -2,9 +2,9 @@ // Licensed under the MIT License. #include "core/providers/rocm/math/matmul.h" + #include "core/providers/cpu/math/matmul_helper.h" -#include "core/providers/rocm/shared_inc/fpgeneric.h" -#include "core/providers/rocm/rocm_allocator.h" +#include "core/providers/rocm/math/matmul_impl.h" namespace onnxruntime { namespace rocm { @@ -22,7 +22,7 @@ namespace rocm { ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \ MatMul, \ kOnnxDomain, \ - 9, 12, \ + 9, 12, \ T, \ kRocmExecutionProvider, \ (*KernelDefBuilder::Create()) \ @@ -31,7 +31,7 @@ namespace rocm { ONNX_OPERATOR_TYPED_KERNEL_EX( \ MatMul, \ kOnnxDomain, \ - 13, \ + 13, \ T, \ kRocmExecutionProvider, \ (*KernelDefBuilder::Create()) \ @@ -43,54 +43,8 @@ REGISTER_KERNEL_TYPED(double) REGISTER_KERNEL_TYPED(MLFloat16) REGISTER_KERNEL_TYPED(BFloat16) -// StridedBatchedGemm can be used for the following GEMM computation -// C[pnm] = A[pnk]*B[km] or C[pnm] = A[pnk]*B[pkm] -static bool CanUseStridedBatchedGemm(const TensorShape& left_shape, const TensorShape& right_shape, - bool transa, bool transb, bool trans_batch_a, bool trans_batch_b, - int64_t& stride_A, int64_t& stride_B, int64_t& stride_C, int64_t& batch_count) { - size_t left_num_dims = left_shape.NumDimensions(); - size_t right_num_dims = right_shape.NumDimensions(); - - if (!(left_num_dims >= 3 && right_num_dims >= 2)) { - return false; - } - - size_t left_leading_axis = trans_batch_a ? 0 : left_num_dims - 2; - size_t right_leading_axis = trans_batch_b ? 0 : right_num_dims - 2; - int64_t left_p = left_shape.SizeToDimension(left_num_dims - 2); - if (trans_batch_a) { - left_p = left_p * left_shape[left_num_dims - 2] / left_shape[0]; - } - int64_t left_k = transa ? left_shape[left_leading_axis] : left_shape[left_num_dims - 1]; - - if (right_num_dims >= 3) { - int64_t right_p = right_shape.SizeToDimension(right_num_dims - 2); - if (trans_batch_b) { - right_p = right_p * right_shape[right_num_dims - 2] / right_shape[0]; - } - if (left_p != right_p) { - return false; - } - } - - int64_t right_k = transb ? right_shape[right_num_dims - 1] : right_shape[right_leading_axis]; - if (left_k != right_k) { - return false; - } - - int64_t n = transa ? left_shape[left_num_dims - 1] : left_shape[left_leading_axis]; - int64_t m = transb ? right_shape[right_leading_axis] : right_shape[right_num_dims - 1]; - stride_A = n * left_k / (trans_batch_a ? left_shape[0] : 1); - stride_B = right_num_dims == 2 ? 0 : right_k * m / (trans_batch_b ? right_shape[0] : 1); - stride_C = n * m; - batch_count = left_p; - return true; -} - template Status MatMul::ComputeInternal(OpKernelContext* ctx) const { - typedef typename ToHipType::MappedType HipT; - const Tensor* left_X = ctx->Input(0); const Tensor* right_X = ctx->Input(1); @@ -106,94 +60,22 @@ Status MatMul::ComputeInternal(OpKernelContext* ctx) const { } MatMulComputeHelper helper; - ORT_RETURN_IF_ERROR(helper.Compute(left_X->Shape(), right_X->Shape(), transa, transb, trans_batch_a_, trans_batch_b_, false)); + ORT_RETURN_IF_ERROR(helper.Compute(left_X->Shape(), right_X->Shape(), transa, + transb, trans_batch_a_, trans_batch_b_, + false)); Tensor* Y = ctx->Output(0, helper.OutputShape()); // Bail out early if the output is going to be empty - if (Y->Shape().Size() == 0) - return Status::OK(); + if (Y->Shape().Size() == 0) return Status::OK(); - const HipT alpha = ToHipType::FromFloat(alpha_); - const HipT zero = ToHipType::FromFloat(0.0f); - - rocblas_operation transA = transa ? rocblas_operation_transpose : rocblas_operation_none; - rocblas_operation transB = transb ? rocblas_operation_transpose : rocblas_operation_none; - const int lda = helper.Lda(transa); - const int ldb = helper.Ldb(transb); - const int ldc = helper.Ldc(); - int64_t stride_A, stride_B, stride_C, batch_count; - - if (helper.OutputOffsets().size() == 1) { - ROCBLAS_RETURN_IF_ERROR(rocblasGemmHelper( - Base::RocblasHandle(), - transB, - transA, - static_cast(helper.N()), - static_cast(helper.M()), - static_cast(helper.K()), - &alpha, - reinterpret_cast(right_X->template Data()), - ldb, - reinterpret_cast(left_X->template Data()), - lda, - &zero, - reinterpret_cast(Y->template MutableData()), - ldc)); - return Status::OK(); - } else if (CanUseStridedBatchedGemm(left_X->Shape(), right_X->Shape(), - transa, transb, trans_batch_a_, trans_batch_b_, stride_A, stride_B, stride_C, batch_count)) { - ROCBLAS_RETURN_IF_ERROR(rocblasGemmStridedBatchedHelper(Base::RocblasHandle(), - transB, - transA, - static_cast(helper.N()), - static_cast(helper.M()), - static_cast(helper.K()), - &alpha, - reinterpret_cast(right_X->template Data()), - ldb, - stride_B, - reinterpret_cast(left_X->template Data()), - lda, - stride_A, - &zero, - reinterpret_cast(Y->template MutableData()), - ldc, - stride_C, - static_cast(batch_count))); - return Status::OK(); + if (MatMulImpl(this, helper, reinterpret_cast(left_X->template Data()), + reinterpret_cast(right_X->template Data()), + reinterpret_cast(Y->template MutableData()), + left_X->Shape(), right_X->Shape(), + transa, transb, trans_batch_a_, trans_batch_b_, alpha_, 0.0f) != Status::OK()) { + return Status(common::ONNXRUNTIME, common::FAIL); } - - // Fill offsets when needed. - helper.FillOffsets(); - RocmAsyncBuffer left_arrays(this, helper.LeftOffsets().size()); - RocmAsyncBuffer right_arrays(this, helper.RightOffsets().size()); - RocmAsyncBuffer output_arrays(this, helper.OutputOffsets().size()); - MatMulComputeHelper::OffsetToArrays(reinterpret_cast(left_X->template Data()), helper.LeftOffsets(), left_arrays.CpuSpan()); - MatMulComputeHelper::OffsetToArrays(reinterpret_cast(right_X->template Data()), helper.RightOffsets(), right_arrays.CpuSpan()); - MatMulComputeHelper::OffsetToArrays(reinterpret_cast(Y->template MutableData()), helper.OutputOffsets(), output_arrays.CpuSpan()); - ORT_RETURN_IF_ERROR(left_arrays.CopyToGpu()); - ORT_RETURN_IF_ERROR(right_arrays.CopyToGpu()); - ORT_RETURN_IF_ERROR(output_arrays.CopyToGpu()); - - // note that onnxruntime OrtValue is row major, while rocblas is column major, - // so swap left/right operands - ROCBLAS_RETURN_IF_ERROR(rocblasGemmBatchedHelper( - Base::RocblasHandle(), - transB, - transA, - static_cast(helper.N()), - static_cast(helper.M()), - static_cast(helper.K()), - &alpha, - right_arrays.GpuPtr(), - ldb, - left_arrays.GpuPtr(), - lda, - &zero, - output_arrays.GpuPtr(), - ldc, - static_cast(helper.OutputOffsets().size()))); return Status::OK(); } diff --git a/onnxruntime/core/providers/rocm/math/matmul_impl.cc b/onnxruntime/core/providers/rocm/math/matmul_impl.cc new file mode 100644 index 0000000000..931dde56d8 --- /dev/null +++ b/onnxruntime/core/providers/rocm/math/matmul_impl.cc @@ -0,0 +1,190 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// Modifications: Remove cudaDeviceProp in LaunchFastGeluKernel. +// Copyright (c) Advanced Micro Devices, Inc. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/rocm/math/matmul_impl.h" + +#include "core/providers/rocm/rocm_allocator.h" +#include "core/providers/rocm/rocm_kernel.h" + +namespace onnxruntime { +namespace rocm { + +// StridedBatchedGemm can be used for the following GEMM computation +// C[pnm] = A[pnk]*B[km] or C[pnm] = A[pnk]*B[pkm] +static bool CanUseStridedBatchedGemm(const TensorShape& left_shape, + const TensorShape& right_shape, + bool transa, bool transb, + bool trans_batch_a, bool trans_batch_b, + int64_t& stride_A, int64_t& stride_B, + int64_t& stride_C, int64_t& batch_count) { + size_t left_num_dims = left_shape.NumDimensions(); + size_t right_num_dims = right_shape.NumDimensions(); + + if (!(left_num_dims >= 3 && right_num_dims >= 2)) { + return false; + } + + size_t left_leading_axis = trans_batch_a ? 0 : left_num_dims - 2; + size_t right_leading_axis = trans_batch_b ? 0 : right_num_dims - 2; + int64_t left_p = left_shape.SizeToDimension(left_num_dims - 2); + if (trans_batch_a) { + left_p = left_p * left_shape[left_num_dims - 2] / left_shape[0]; + } + int64_t left_k = + transa ? left_shape[left_leading_axis] : left_shape[left_num_dims - 1]; + + if (right_num_dims >= 3) { + int64_t right_p = right_shape.SizeToDimension(right_num_dims - 2); + if (trans_batch_b) { + right_p = right_p * right_shape[right_num_dims - 2] / right_shape[0]; + } + if (left_p != right_p) { + return false; + } + } + + int64_t right_k = transb ? right_shape[right_num_dims - 1] + : right_shape[right_leading_axis]; + if (left_k != right_k) { + return false; + } + + int64_t n = + transa ? left_shape[left_num_dims - 1] : left_shape[left_leading_axis]; + int64_t m = transb ? right_shape[right_leading_axis] + : right_shape[right_num_dims - 1]; + stride_A = n * left_k / (trans_batch_a ? left_shape[0] : 1); + stride_B = right_num_dims == 2 ? 0 : right_k * m / (trans_batch_b ? right_shape[0] : 1); + stride_C = n * m; + batch_count = left_p; + return true; +} + +template +Status MatMulImpl(const RocmKernel* op, MatMulComputeHelper& helper, + const T* left_x_data, const T* right_x_data, T* output_y_data, + const TensorShape& left_shape, const TensorShape& right_shape, + bool transa, bool transb, bool trans_batch_a, bool trans_batch_b, + const float t_alpha, const float t_zero) { + typedef typename ToHipType::MappedType HipT; + + const HipT alpha = ToHipType::FromFloat(t_alpha); + const HipT zero = ToHipType::FromFloat(t_zero); + + rocblas_operation transA = transa ? rocblas_operation_transpose : rocblas_operation_none; + rocblas_operation transB = transb ? rocblas_operation_transpose : rocblas_operation_none; + const int lda = helper.Lda(transa); + const int ldb = helper.Ldb(transb); + const int ldc = helper.Ldc(); + int64_t stride_A, stride_B, stride_C, batch_count; + + if (helper.OutputOffsets().size() == 1) { + ROCBLAS_RETURN_IF_ERROR(rocblasGemmHelper( + op->RocblasHandle(), transB, transA, static_cast(helper.N()), + static_cast(helper.M()), static_cast(helper.K()), &alpha, + reinterpret_cast(right_x_data), ldb, + reinterpret_cast(left_x_data), lda, &zero, + reinterpret_cast(output_y_data), ldc)); + return Status::OK(); + } else if (CanUseStridedBatchedGemm(left_shape, right_shape, + transa, transb, trans_batch_a, trans_batch_b, + stride_A, stride_B, stride_C, batch_count)) { + ROCBLAS_RETURN_IF_ERROR(rocblasGemmStridedBatchedHelper( + op->RocblasHandle(), transB, transA, static_cast(helper.N()), + static_cast(helper.M()), static_cast(helper.K()), &alpha, + reinterpret_cast(right_x_data), ldb, stride_B, + reinterpret_cast(left_x_data), lda, stride_A, &zero, + reinterpret_cast(output_y_data), ldc, stride_C, + static_cast(batch_count))); + return Status::OK(); + } + + // Fill offsets when needed. + helper.FillOffsets(); + RocmKernel::RocmAsyncBuffer left_arrays(op, helper.LeftOffsets().size()); + RocmKernel::RocmAsyncBuffer right_arrays(op, helper.RightOffsets().size()); + RocmKernel::RocmAsyncBuffer output_arrays(op, helper.OutputOffsets().size()); + MatMulComputeHelper::OffsetToArrays( + reinterpret_cast(left_x_data), + helper.LeftOffsets(), left_arrays.CpuSpan()); + MatMulComputeHelper::OffsetToArrays( + reinterpret_cast(right_x_data), + helper.RightOffsets(), right_arrays.CpuSpan()); + MatMulComputeHelper::OffsetToArrays( + reinterpret_cast(output_y_data), + helper.OutputOffsets(), output_arrays.CpuSpan()); + ORT_RETURN_IF_ERROR(left_arrays.CopyToGpu()); + ORT_RETURN_IF_ERROR(right_arrays.CopyToGpu()); + ORT_RETURN_IF_ERROR(output_arrays.CopyToGpu()); + + // note that onnxruntime OrtValue is row major, while rocblas is column major, + // so swap left/right operands + ROCBLAS_RETURN_IF_ERROR(rocblasGemmBatchedHelper( + op->RocblasHandle(), transB, transA, static_cast(helper.N()), + static_cast(helper.M()), static_cast(helper.K()), &alpha, + right_arrays.GpuPtr(), ldb, + left_arrays.GpuPtr(), lda, &zero, + output_arrays.GpuPtr(), ldc, + static_cast(helper.OutputOffsets().size()))); + return Status::OK(); +} + +template Status MatMulImpl(const RocmKernel* op, + MatMulComputeHelper& helper, + const float* left_x_data, + const float* right_x_data, + float* output_y_data, + const TensorShape& left_shape, + const TensorShape& right_shape, + bool transa, + bool transb, + bool trans_batch_a, + bool trans_batch_b, + const float t_alpha, + const float t_zero); +template Status MatMulImpl(const RocmKernel* op, + MatMulComputeHelper& helper, + const double* left_x_data, + const double* right_x_data, + double* output_y_data, + const TensorShape& left_shape, + const TensorShape& right_shape, + bool transa, + bool transb, + bool trans_batch_a, + bool trans_batch_b, + const float t_alpha, + const float t_zero); +template Status MatMulImpl(const RocmKernel* op, + MatMulComputeHelper& helper, + const MLFloat16* left_x_data, + const MLFloat16* right_x_data, + MLFloat16* output_y_data, + const TensorShape& left_shape, + const TensorShape& right_shape, + bool transa, + bool transb, + bool trans_batch_a, + bool trans_batch_b, + const float t_alpha, + const float t_zero); +template Status MatMulImpl(const RocmKernel* op, + MatMulComputeHelper& helper, + const BFloat16* left_x_data, + const BFloat16* right_x_data, + BFloat16* output_y_data, + const TensorShape& left_shape, + const TensorShape& right_shape, + bool transa, + bool transb, + bool trans_batch_a, + bool trans_batch_b, + const float t_alpha, + const float t_zero); + +} // namespace rocm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/math/matmul_impl.h b/onnxruntime/core/providers/rocm/math/matmul_impl.h new file mode 100644 index 0000000000..f22f3b8eac --- /dev/null +++ b/onnxruntime/core/providers/rocm/math/matmul_impl.h @@ -0,0 +1,25 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// Modifications: Remove cudaDeviceProp in LaunchFastGeluKernel. +// Copyright (c) Advanced Micro Devices, Inc. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/providers/rocm/shared_inc/fpgeneric.h" +#include "core/providers/cpu/math/matmul_helper.h" +#include "core/providers/rocm/rocm_kernel.h" + +namespace onnxruntime { +namespace rocm { + +template +Status MatMulImpl(const RocmKernel* op, MatMulComputeHelper& helper, + const T* left_x_data, const T* right_x_data, T* output_y_data, + const TensorShape& left_shape, const TensorShape& right_shape, + bool transa, bool transb, bool trans_batch_a, bool trans_batch_b, + const float t_alpha, const float t_zero); + +} // namespace rocm +} // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/gemm_fastgelu_op_test.cc b/onnxruntime/test/contrib_ops/gemm_fastgelu_op_test.cc new file mode 100644 index 0000000000..6852788d75 --- /dev/null +++ b/onnxruntime/test/contrib_ops/gemm_fastgelu_op_test.cc @@ -0,0 +1,236 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include +#include "core/platform/threadpool.h" +#include "core/util/math.h" +#include "core/util/thread_utils.h" +#include "test/common/cuda_op_test_utils.h" +#include "test/common/tensor_op_test_utils.h" +#include "test/providers/provider_test_utils.h" + +namespace onnxruntime { +namespace test { +namespace gemmfastgelu { + +#if defined(USE_ROCM) +static void RunGemmFastGeluGpuTest(const std::vector& input_data, const std::vector& weight_data, + const std::vector& bias_data, const std::vector& output_data, + const std::vector& input_dims, const std::vector& weight_dims, + const std::vector& bias_dims, const std::vector& output_dims, + bool has_bias, bool use_float16 = false) { + OpTester tester("GemmFastGelu", 1, onnxruntime::kMSDomain); + + if (use_float16) { + tester.AddInput("X", input_dims, ToFloat16(input_data)); + tester.AddInput("W", weight_dims, ToFloat16(weight_data)); + if (has_bias) { + tester.AddInput("bias", bias_dims, ToFloat16(bias_data)); + } + tester.AddOutput("Y", output_dims, ToFloat16(output_data)); + } else { + tester.AddInput("X", input_dims, input_data); + tester.AddInput("W", weight_dims, weight_data); + if (has_bias) { + tester.AddInput("bias", bias_dims, bias_data); + } + tester.AddOutput("Y", output_dims, output_data); + } + + std::vector> execution_providers; + execution_providers.push_back(DefaultRocmExecutionProvider()); + tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} +#endif + +TEST(GemmFastGeluTest, GemmFastGeluWithoutBiasFloat32) { + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int dense_size = 6; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.7f, -0.5f, 0.7f, 1.2f, + 0.3f, 0.1f, 0.8f, -1.6f, + 0.9f, -0.1f, 3.0f, 2.f, + 0.4f, -0.7f, -0.3f, 0.6f}; + + std::vector bias_data = {}; + + std::vector output_data = { + 3.4894f, 1.8455f, 0.0260f, 0.2229f, -0.1003f, 0.0902f, + -0.1323f, -0.0953f, 0.0778f, 0.2152f, 0.6715f, -0.0240f}; + + std::vector input_dims = {batch_size, sequence_length, hidden_size}; + std::vector weight_dims = {hidden_size, dense_size}; + std::vector bias_dims = {dense_size}; + std::vector output_dims = {batch_size, sequence_length, dense_size}; +#if defined(USE_ROCM) + RunGemmFastGeluGpuTest(input_data, weight_data, bias_data, output_data, + input_dims, weight_dims, bias_dims, output_dims, + false); +#endif +} + +TEST(GemmFastGeluTest, GemmFastGeluWithBiasFloat32) { + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int dense_size = 6; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.7f, -0.5f, 0.7f, 1.2f, + 0.3f, 0.1f, 0.8f, -1.6f, + 0.9f, -0.1f, 3.0f, 2.f, + 0.4f, -0.7f, -0.3f, 0.6f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, -0.6f, 0.4f}; + + std::vector output_data = { + 2.9862f, 2.4849f, 1.1177f, 2.4329f, -0.1681f, 0.3988f, + -0.0702f, -0.1633f, 1.2190f, 2.4225f, 0.1428f, 0.2229f}; + + std::vector input_dims = {batch_size, sequence_length, hidden_size}; + std::vector weight_dims = {hidden_size, dense_size}; + std::vector bias_dims = {dense_size}; + std::vector output_dims = {batch_size, sequence_length, dense_size}; +#if defined(USE_ROCM) + RunGemmFastGeluGpuTest(input_data, weight_data, bias_data, output_data, + input_dims, weight_dims, bias_dims, output_dims, + true); +#endif +} + +// CUDA and ROCm only for Float16 and BFloat16 type. +#if defined(USE_ROCM) +TEST(GemmFastGeluTest, GemmFastGeluWithoutBiasFloat16) { + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int dense_size = 6; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.7f, -0.5f, 0.7f, 1.2f, + 0.3f, 0.1f, 0.8f, -1.6f, + 0.9f, -0.1f, 3.0f, 2.f, + 0.4f, -0.7f, -0.3f, 0.6f}; + + std::vector bias_data = {}; + + std::vector output_data = { + 3.4902f, 1.8467f, 0.0259f, 0.2227f, -0.1005f, 0.0901f, + -0.1324f, -0.0955f, 0.0778f, 0.2156f, 0.6714f, -0.0241f}; + + std::vector input_dims = {batch_size, sequence_length, hidden_size}; + std::vector weight_dims = {hidden_size, dense_size}; + std::vector bias_dims = {dense_size}; + std::vector output_dims = {batch_size, sequence_length, dense_size}; + RunGemmFastGeluGpuTest(input_data, weight_data, bias_data, output_data, + input_dims, weight_dims, bias_dims, output_dims, + false); +} + +TEST(GemmFastGeluTest, GemmFastGeluWithBiasFloat16) { + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int dense_size = 6; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.7f, -0.5f, 0.7f, 1.2f, + 0.3f, 0.1f, 0.8f, -1.6f, + 0.9f, -0.1f, 3.0f, 2.f, + 0.4f, -0.7f, -0.3f, 0.6f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, -0.6f, 0.4f}; + + std::vector output_data = { + 2.9883f, 2.4844f, 1.1182f, 2.4316f, -0.1680f, 0.3984f, + -0.0701f, -0.1633f, 1.2178f, 2.4219f, 0.1426f, 0.2227f}; + + std::vector input_dims = {batch_size, sequence_length, hidden_size}; + std::vector weight_dims = {hidden_size, dense_size}; + std::vector bias_dims = {dense_size}; + std::vector output_dims = {batch_size, sequence_length, dense_size}; + RunGemmFastGeluGpuTest(input_data, weight_data, bias_data, output_data, + input_dims, weight_dims, bias_dims, output_dims, + true); +} + +TEST(GemmFastGeluTest, GemmFastGeluWithBias_BFloat16) { + OpTester tester("GemmFastGelu", 1, onnxruntime::kMSDomain); + + int batch_size = 1; + int sequence_length = 2; + int hidden_size = 4; + int dense_size = 6; + + std::vector input_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f}; + + std::vector weight_data = { + 0.8f, -0.5f, 0.0f, 1.f, + 0.5f, 0.2f, 0.3f, -0.6f, + 0.7f, -0.5f, 0.7f, 1.2f, + 0.3f, 0.1f, 0.8f, -1.6f, + 0.9f, -0.1f, 3.0f, 2.f, + 0.4f, -0.7f, -0.3f, 0.6f}; + + std::vector bias_data = { + -0.5f, 0.6f, 1.2f, 2.1f, -0.6f, 0.4f}; + + std::vector output_data = { + 2.9883f, 2.4844f, 1.1182f, 2.4316f, -0.1680f, 0.3984f, + -0.0701f, -0.1633f, 1.2178f, 2.4219f, 0.1426f, 0.2227f}; + + std::vector input_dims = {batch_size, sequence_length, hidden_size}; + std::vector weight_dims = {hidden_size, dense_size}; + std::vector bias_dims = {dense_size}; + std::vector output_dims = {batch_size, sequence_length, dense_size}; + + std::vector f_X = FloatsToBFloat16s(input_data); + std::vector f_W = FloatsToBFloat16s(weight_data); + std::vector f_B = FloatsToBFloat16s(bias_data); + std::vector f_Y = FloatsToBFloat16s(output_data); + + tester.AddInput("X", input_dims, f_X); + tester.AddInput("W", weight_dims, f_W); + tester.AddInput("bias", bias_dims, f_B); + tester.AddOutput("Y", output_dims, f_Y); + + std::vector> execution_providers; + execution_providers.push_back(DefaultRocmExecutionProvider()); + tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers); +} +#endif + +} // namespace gemmfastgelu +} // namespace test +} // namespace onnxruntime