diff --git a/onnxruntime/core/platform/env_var_utils.h b/onnxruntime/core/platform/env_var_utils.h index 3b77e67c25..72b464f433 100644 --- a/onnxruntime/core/platform/env_var_utils.h +++ b/onnxruntime/core/platform/env_var_utils.h @@ -83,7 +83,7 @@ std::optional ParseTestOnlyEnvironmentVariable(const std::string& name, std::string default_hint = "End users should opt for provider options or session options."; const std::string& logged_hint = hint.empty() ? default_hint : hint; - LOGS_DEFAULT(WARNING) << "Environment variable " << name << " is used. It is reserved for internal testing prupose. " + LOGS_DEFAULT(WARNING) << "Environment variable " << name << " is used. It is reserved for internal testing purpose. " << logged_hint; return env; diff --git a/onnxruntime/core/providers/rocm/tunable/gemm_rocblas.h b/onnxruntime/core/providers/rocm/tunable/gemm_rocblas.h index a732f9c0c2..d36bc41527 100644 --- a/onnxruntime/core/providers/rocm/tunable/gemm_rocblas.h +++ b/onnxruntime/core/providers/rocm/tunable/gemm_rocblas.h @@ -103,6 +103,25 @@ constexpr rocblas_datatype RocBlasComputeTypeFor(const BFloat16*) { return rocblas_datatype_f32_r; } +template +auto DoCastForHalfOrBfloat16(const T fp) { + return fp; +} + +template <> +auto DoCastForHalfOrBfloat16(const half fp) { + // alpha and beta should be the same as compute_type, in half case it is float. + float h = onnxruntime::math::halfToFloat(*reinterpret_cast(&fp)); + return h; +} + +template <> +auto DoCastForHalfOrBfloat16(const BFloat16 fp) { + // alpha and beta should be the same as compute_type, in bfloat16 case it is float. + float h = fp.ToFloat(); + return h; +} + template class IndexedRocBlasGemmOp { public: @@ -113,16 +132,18 @@ class IndexedRocBlasGemmOp { Status operator()(const GemmParams* params) { RocblasHandleStreamGuard guard(params->handle, params->stream); + auto h_a = DoCastForHalfOrBfloat16(params->alpha); + auto h_b = DoCastForHalfOrBfloat16(params->beta); return ROCBLAS_CALL( rocblas_gemm_ex( params->handle, params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, params->n, params->m, params->k, - &(params->alpha), + &h_a, params->b, RocBlasDataTypeFor(params->b), params->ldb, params->a, RocBlasDataTypeFor(params->a), params->lda, - &(params->beta), + &h_b, params->c, RocBlasDataTypeFor(params->c), params->ldc, params->c, RocBlasDataTypeFor(params->c), params->ldc, RocBlasComputeTypeFor(params->a), @@ -169,16 +190,18 @@ class RocBlasGemmTunableOp : public TunableOp> { private: std::vector GetSolutions(const GemmParams* params) { int num_solutions = 0; + auto h_a = DoCastForHalfOrBfloat16(params->alpha); + auto h_b = DoCastForHalfOrBfloat16(params->beta); // Get the number of candidate solutions ROCBLAS_CALL_THROW(rocblas_gemm_ex_get_solutions( params->handle, params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, params->n, params->m, params->k, - &(params->alpha), + &h_a, params->b, RocBlasDataTypeFor(params->b), params->ldb, params->a, RocBlasDataTypeFor(params->a), params->lda, - &(params->beta), + &h_b, params->c, RocBlasDataTypeFor(params->c), params->ldc, params->c, RocBlasDataTypeFor(params->c), params->ldc, RocBlasComputeTypeFor(params->a), @@ -194,10 +217,10 @@ class RocBlasGemmTunableOp : public TunableOp> { params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, params->n, params->m, params->k, - &(params->alpha), + &h_a, params->b, RocBlasDataTypeFor(params->b), params->ldb, params->a, RocBlasDataTypeFor(params->a), params->lda, - &(params->beta), + &h_b, params->c, RocBlasDataTypeFor(params->c), params->ldc, params->c, RocBlasDataTypeFor(params->c), params->ldc, RocBlasComputeTypeFor(params->a), @@ -210,6 +233,234 @@ class RocBlasGemmTunableOp : public TunableOp> { } }; +template +class IndexedRocBlasBatchedGemmOp { + public: + IndexedRocBlasBatchedGemmOp() + : index_(0) {} + IndexedRocBlasBatchedGemmOp(int index) + : index_(index) {} + + Status operator()(const BatchedGemmParams* params) { + RocblasHandleStreamGuard guard(params->handle, params->stream); + auto h_a = DoCastForHalfOrBfloat16(params->alpha); + auto h_b = DoCastForHalfOrBfloat16(params->beta); + return ROCBLAS_CALL( + rocblas_gemm_batched_ex( + params->handle, + params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->n, params->m, params->k, + &h_a, + params->bs, RocBlasDataTypeFor(*(params->bs)), params->ldb, + params->as, RocBlasDataTypeFor(*(params->as)), params->lda, + &h_b, + params->cs, RocBlasDataTypeFor(*(params->cs)), params->ldc, + params->cs, RocBlasDataTypeFor(*(params->cs)), params->ldc, + params->batch, + RocBlasComputeTypeFor(*(params->as)), + rocblas_gemm_algo_solution_index, + index_, + rocblas_gemm_flags_none)); + } + + Status IsSupported(const BatchedGemmParams*) { + return Status::OK(); + } + + private: + int index_; +}; + +template +class RocBlasBatchedGemmTunableOp : public TunableOp> { + public: + RocBlasBatchedGemmTunableOp() { + // Ensure that the default implementation is always present + this->RegisterOp(IndexedRocBlasBatchedGemmOp{0}); + } + + Status IsSupported(const BatchedGemmParams* params) { + ORT_UNUSED_PARAMETER(params); + return Status::OK(); + } + + protected: + virtual int FindFastest(const BatchedGemmParams* params) override { + auto solution_indices = this->GetSolutions(params); + std::vector>> candidates; + for (int solution_idx : solution_indices) { + candidates.emplace_back(IndexedRocBlasBatchedGemmOp{solution_idx}); + } + + auto id = this->FindFastestImpl(params, candidates); + // memoize the result + this->RegisterOp(std::move(candidates[id])); + return this->NumberOfOps() - 1; + } + + private: + std::vector GetSolutions(const BatchedGemmParams* params) { + int num_solutions = 0; + auto h_a = DoCastForHalfOrBfloat16(params->alpha); + auto h_b = DoCastForHalfOrBfloat16(params->beta); + // Get the number of candidate solutions + ROCBLAS_CALL_THROW(rocblas_gemm_batched_ex_get_solutions( + params->handle, + params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->n, params->m, params->k, + &h_a, + params->bs, RocBlasDataTypeFor(*(params->bs)), params->ldb, + params->as, RocBlasDataTypeFor(*(params->as)), params->lda, + &h_b, + params->cs, RocBlasDataTypeFor(*(params->cs)), params->ldc, + params->cs, RocBlasDataTypeFor(*(params->cs)), params->ldc, + params->batch, + RocBlasComputeTypeFor(*(params->as)), + rocblas_gemm_algo_solution_index, + rocblas_gemm_flags_none, + NULL, + &num_solutions)); + + // Get the actual candidate solutions + std::vector solutions(num_solutions); + ROCBLAS_CALL_THROW(rocblas_gemm_batched_ex_get_solutions( + params->handle, + params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->n, params->m, params->k, + &h_a, + params->bs, RocBlasDataTypeFor(*(params->bs)), params->ldb, + params->as, RocBlasDataTypeFor(*(params->as)), params->lda, + &h_b, + params->cs, RocBlasDataTypeFor(*(params->cs)), params->ldc, + params->cs, RocBlasDataTypeFor(*(params->cs)), params->ldc, + params->batch, + RocBlasComputeTypeFor(*(params->as)), + rocblas_gemm_algo_solution_index, + rocblas_gemm_flags_none, + solutions.data(), + &num_solutions)); + + return solutions; + } +}; + +template +class IndexedRocBlasStridedBatchedGemmOp { + public: + IndexedRocBlasStridedBatchedGemmOp() + : index_(0) {} + IndexedRocBlasStridedBatchedGemmOp(int index) + : index_(index) {} + + Status operator()(const StridedBatchedGemmParams* params) { + RocblasHandleStreamGuard guard(params->handle, params->stream); + auto h_a = DoCastForHalfOrBfloat16(params->alpha); + auto h_b = DoCastForHalfOrBfloat16(params->beta); + return ROCBLAS_CALL( + rocblas_gemm_strided_batched_ex( + params->handle, + params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->n, params->m, params->k, + &h_a, + params->b, RocBlasDataTypeFor(params->b), params->ldb, params->stride_b, + params->a, RocBlasDataTypeFor(params->a), params->lda, params->stride_a, + &h_b, + params->c, RocBlasDataTypeFor(params->c), params->ldc, params->stride_c, + params->c, RocBlasDataTypeFor(params->c), params->ldc, params->stride_c, + params->batch, + RocBlasComputeTypeFor(params->a), + rocblas_gemm_algo_solution_index, + index_, + rocblas_gemm_flags_none)); + } + + Status IsSupported(const StridedBatchedGemmParams*) { + return Status::OK(); + } + + private: + int index_; +}; + +template +class RocBlasStridedBatchedGemmTunableOp : public TunableOp> { + public: + RocBlasStridedBatchedGemmTunableOp() { + // Ensure that the default implementation is always present + this->RegisterOp(IndexedRocBlasStridedBatchedGemmOp{0}); + } + + Status IsSupported(const StridedBatchedGemmParams* params) { + ORT_UNUSED_PARAMETER(params); + return Status::OK(); + } + + protected: + virtual int FindFastest(const StridedBatchedGemmParams* params) override { + auto solution_indices = this->GetSolutions(params); + std::vector>> candidates; + for (int solution_idx : solution_indices) { + candidates.emplace_back(IndexedRocBlasStridedBatchedGemmOp{solution_idx}); + } + + auto id = this->FindFastestImpl(params, candidates); + // memoize the result + this->RegisterOp(std::move(candidates[id])); + return this->NumberOfOps() - 1; + } + + private: + std::vector GetSolutions(const StridedBatchedGemmParams* params) { + int num_solutions = 0; + auto h_a = DoCastForHalfOrBfloat16(params->alpha); + auto h_b = DoCastForHalfOrBfloat16(params->beta); + // Get the number of candidate solutions + ROCBLAS_CALL_THROW(rocblas_gemm_strided_batched_ex_get_solutions( + params->handle, + params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->n, params->m, params->k, + &h_a, + params->b, RocBlasDataTypeFor(params->b), params->ldb, params->stride_b, + params->a, RocBlasDataTypeFor(params->a), params->lda, params->stride_a, + &h_b, + params->c, RocBlasDataTypeFor(params->c), params->ldc, params->stride_c, + params->c, RocBlasDataTypeFor(params->c), params->ldc, params->stride_c, + params->batch, + RocBlasComputeTypeFor(params->a), + rocblas_gemm_algo_solution_index, + rocblas_gemm_flags_none, + NULL, + &num_solutions)); + + // Get the actual candidate solutions + std::vector solutions(num_solutions); + ROCBLAS_CALL_THROW(rocblas_gemm_strided_batched_ex_get_solutions( + params->handle, + params->opb == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->opa == BlasOp::N ? rocblas_operation_none : rocblas_operation_transpose, + params->n, params->m, params->k, + &h_a, + params->b, RocBlasDataTypeFor(params->b), params->ldb, params->stride_b, + params->a, RocBlasDataTypeFor(params->a), params->lda, params->stride_a, + &h_b, + params->c, RocBlasDataTypeFor(params->c), params->ldc, params->stride_c, + params->c, RocBlasDataTypeFor(params->c), params->ldc, params->stride_c, + params->batch, + RocBlasComputeTypeFor(params->a), + rocblas_gemm_algo_solution_index, + rocblas_gemm_flags_none, + solutions.data(), + &num_solutions)); + + return solutions; + } +}; + #endif /* #ifdef USE_ROCBLAS_EXTENSION_API */ template diff --git a/onnxruntime/core/providers/rocm/tunable/gemm_tunable.cuh b/onnxruntime/core/providers/rocm/tunable/gemm_tunable.cuh index 4813f3ad2b..ceda036fb1 100644 --- a/onnxruntime/core/providers/rocm/tunable/gemm_tunable.cuh +++ b/onnxruntime/core/providers/rocm/tunable/gemm_tunable.cuh @@ -86,6 +86,10 @@ class BatchedGemmTunableOp : public TunableOp> { public: BatchedGemmTunableOp() { this->RegisterOp(RocBlasBatchedGemmOp); + +#ifdef USE_ROCBLAS_EXTENSION_API + this->RegisterNestedTunableOp(&rocblas_batched_gemm_tunable_op_); +#endif /* #ifdef USE_ROCBLAS_EXTENSION_API */ } const BatchedGemmParams* PreTuning(const BatchedGemmParams* params) override { @@ -122,6 +126,11 @@ class BatchedGemmTunableOp : public TunableOp> { delete params; } } + + private: +#ifdef USE_ROCBLAS_EXTENSION_API + RocBlasBatchedGemmTunableOp rocblas_batched_gemm_tunable_op_; +#endif }; template @@ -129,6 +138,11 @@ class StridedBatchedGemmTunableOp : public TunableOp public: StridedBatchedGemmTunableOp() { this->RegisterOp(RocBlasStridedBatchedGemmOp); + +#ifdef USE_ROCBLAS_EXTENSION_API + this->RegisterNestedTunableOp(&rocblas_strided_batched_gemm_tunable_op_); +#endif /* #ifdef USE_ROCBLAS_EXTENSION_API */ + #ifdef USE_COMPOSABLE_KERNEL for (auto&& [_, op] : GetCKStridedBatchedGemmTypeStringAndOps()) { ORT_UNUSED_PARAMETER(_); @@ -155,6 +169,11 @@ class StridedBatchedGemmTunableOp : public TunableOp delete params; } } + + private: +#ifdef USE_ROCBLAS_EXTENSION_API + RocBlasStridedBatchedGemmTunableOp rocblas_strided_batched_gemm_tunable_op_; +#endif }; } // namespace internal