From d03a80e1497d376f20f01d4c8468ce8d33b7a6dc Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Sat, 15 Oct 2022 00:04:50 +0000 Subject: [PATCH 1/6] rocblas alt impl during backward pass only --- .../core/optimizer/graph_transformer_utils.cc | 6 +++ .../core/optimizer/rocm_blas_alt_impl.cc | 52 +++++++++++++++++++ .../core/optimizer/rocm_blas_alt_impl.h | 19 +++++++ .../core/providers/rocm/backward_guard.cc | 30 +++++++++++ .../core/providers/rocm/backward_guard.h | 15 ++++++ onnxruntime/core/providers/rocm/math/gemm.cc | 4 ++ onnxruntime/core/providers/rocm/rocm_kernel.h | 15 +++++- .../providers/rocm/shared_inc/fpgeneric.h | 26 ++++++++-- 8 files changed, 161 insertions(+), 6 deletions(-) create mode 100644 onnxruntime/core/optimizer/rocm_blas_alt_impl.cc create mode 100644 onnxruntime/core/optimizer/rocm_blas_alt_impl.h create mode 100644 onnxruntime/core/providers/rocm/backward_guard.cc create mode 100644 onnxruntime/core/providers/rocm/backward_guard.h diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index 22a64633e450a..bf13fdfcd0741 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -10,6 +10,7 @@ #include "core/optimizer/nhwc_transformer.h" #include "core/optimizer/qdq_transformer/qdq_final_cleanup.h" #include "core/optimizer/qdq_transformer/selectors_actions/qdq_selector_action_transformer.h" +#include "core/optimizer/rocm_blas_alt_impl.h" #include "core/optimizer/selectors_actions/selector_action_transformer_apply_contexts.h" #include "core/session/onnxruntime_session_options_config_keys.h" #include "core/optimizer/conv_add_act_fusion.h" @@ -206,6 +207,11 @@ InlinedVector> GenerateTransformers( // shouldn't affect the end result - just easier to debug any issue if it's last. auto cpu_allocator = cpu_execution_provider.GetAllocator(0, OrtMemTypeDefault); transformers.emplace_back(std::make_unique(std::move(cpu_allocator))); + + std::cerr << __FILE__ << ":" << __LINE__ << " emplace_back RocmBlasAltImpl" << std::endl; + // TODO document + const InlinedHashSet rocm_ep = {onnxruntime::kRocmExecutionProvider}; + transformers.emplace_back(std::make_unique(rocm_ep)); } break; case TransformerLevel::Level2: { diff --git a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc new file mode 100644 index 0000000000000..9594358dd257c --- /dev/null +++ b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc @@ -0,0 +1,52 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +#include + +#include "core/optimizer/initializer.h" +#include "core/optimizer/rocm_blas_alt_impl.h" +#include "core/graph/graph_utils.h" + +#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + +using namespace ONNX_NAMESPACE; +using namespace ::onnxruntime::common; +namespace onnxruntime { + +Status RocmBlasAltImpl::ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const { + GraphViewer graph_viewer(graph); + const auto& node_topology_list = graph_viewer.GetNodesInTopologicalOrder(); + + bool is_backward_pass = false; + + for (auto node_index : node_topology_list) { + auto& node = *graph.GetNode(node_index); + + std::cerr << PRE << node << std::endl; + +#if 0 + if (node.OpType() == "YieldOp") { + is_backward_pass = true; + //std::cerr << PRE << "YieldOp found, before recurse" << std::endl; + ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); + //std::cerr << PRE << "YieldOp found, after recurse" << std::endl; + } + else +#else + is_backward_pass = true; +#endif + { + ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); + } + + //if (node.OpType() == "MatMul" || node.OpType() == "FusedMatMul" || node.OpType() == "Gemm") { + //std::cerr << PRE << "HIT, is_backward_pass " << is_backward_pass << std::endl; + if (is_backward_pass) { + node.AddAttribute(std::string("__altimpl"), static_cast(1)); + modified = true; + } + //} + } + + return Status::OK(); +} +} // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/rocm_blas_alt_impl.h b/onnxruntime/core/optimizer/rocm_blas_alt_impl.h new file mode 100644 index 0000000000000..11744d0dac32b --- /dev/null +++ b/onnxruntime/core/optimizer/rocm_blas_alt_impl.h @@ -0,0 +1,19 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/optimizer/graph_transformer.h" +#include "core/graph/graph_utils.h" + +namespace onnxruntime { + +class RocmBlasAltImpl : public GraphTransformer { + public: + RocmBlasAltImpl(const InlinedHashSet& compatible_execution_providers = {}) noexcept + : GraphTransformer("RocmBlasAltImpl", compatible_execution_providers) {} + + Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; +}; + +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/backward_guard.cc b/onnxruntime/core/providers/rocm/backward_guard.cc new file mode 100644 index 0000000000000..df05e8b82a50e --- /dev/null +++ b/onnxruntime/core/providers/rocm/backward_guard.cc @@ -0,0 +1,30 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +#include "core/providers/rocm/backward_guard.h" + +#include +#include +#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + +namespace onnxruntime { + +thread_local bool BackwardPassGuard::is_backward_pass_; + +BackwardPassGuard::BackwardPassGuard() { + //std::cerr << PRE << "BackwardPassGuard ctor pre " << is_backward_pass_ << std::endl; + is_backward_pass_ = true; + //std::cerr << PRE << "BackwardPassGuard ctor post " << is_backward_pass_ << std::endl; +} + +BackwardPassGuard::~BackwardPassGuard() { + //std::cerr << PRE << "BackwardPassGuard dtor pre " << is_backward_pass_ << std::endl; + is_backward_pass_ = false; + //std::cerr << PRE << "BackwardPassGuard dtor post " << is_backward_pass_ << std::endl; +} + +bool BackwardPassGuard::is_backward_pass() { + //std::cerr << PRE << "is_backward_pass_ " << is_backward_pass_ << std::endl; + return is_backward_pass_; +} + +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/backward_guard.h b/onnxruntime/core/providers/rocm/backward_guard.h new file mode 100644 index 0000000000000..e36785af374c4 --- /dev/null +++ b/onnxruntime/core/providers/rocm/backward_guard.h @@ -0,0 +1,15 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +#pragma once + +namespace onnxruntime { + +struct BackwardPassGuard { + BackwardPassGuard(); + ~BackwardPassGuard(); + static bool is_backward_pass(); +private: + static thread_local bool is_backward_pass_; +}; + +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/math/gemm.cc b/onnxruntime/core/providers/rocm/math/gemm.cc index b0a32b5accf81..191ed96b0613a 100644 --- a/onnxruntime/core/providers/rocm/math/gemm.cc +++ b/onnxruntime/core/providers/rocm/math/gemm.cc @@ -6,6 +6,8 @@ #include "core/providers/rocm/rocm_common.h" #include "core/providers/rocm/shared_inc/fpgeneric.h" +#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + namespace onnxruntime { namespace rocm { @@ -59,6 +61,8 @@ template Status Gemm::ComputeInternal(OpKernelContext* ctx) const { typedef typename ToHipType::MappedType HipT; + //auto altimpl = Info().GetAttrOrDefault("altimpl", 0); + //std::cerr << PRE << "wtf altimpl " << altimpl << std::endl; const auto* X = ctx->Input(0); const auto* W = ctx->Input(1); const auto* B = ctx->Input(2); diff --git a/onnxruntime/core/providers/rocm/rocm_kernel.h b/onnxruntime/core/providers/rocm/rocm_kernel.h index a05bd21505313..0c8522594c641 100644 --- a/onnxruntime/core/providers/rocm/rocm_kernel.h +++ b/onnxruntime/core/providers/rocm/rocm_kernel.h @@ -3,10 +3,13 @@ #pragma once +#include "core/providers/rocm/backward_guard.h" #include "core/providers/rocm/rocm_common.h" #include "core/providers/rocm/rocm_execution_provider.h" #include "core/providers/rocm/rocm_fwd.h" +#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + namespace onnxruntime { namespace rocm { @@ -22,7 +25,17 @@ class RocmKernel : public OpKernel { } Status Compute(OpKernelContext* p_op_kernel_context) const override { - auto s = ComputeInternal(p_op_kernel_context); + Status s; + auto altimpl = Info().GetAttrOrDefault("__altimpl", 0); + //std::cerr << PRE << "altimpl " << altimpl << std::endl; + if (altimpl) { + //std::cerr << PRE << "creating BackwardPassGuard" << std::endl; + BackwardPassGuard guard; + s = ComputeInternal(p_op_kernel_context); + } + else { + s = ComputeInternal(p_op_kernel_context); + } // use this to precisely locate the node where ROCM failure comes from // if (hipSuccess != hipDeviceSynchronize()) // __debugbreak(); diff --git a/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h b/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h index 030275c52d54a..5e116ee9a9818 100644 --- a/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h +++ b/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h @@ -3,10 +3,26 @@ #pragma once +#include "core/providers/rocm/backward_guard.h" #include "core/providers/rocm/rocm_common.h" +#define ORT_ROCBLAS_VERSION_DECIMAL (ROCBLAS_VERSION_MAJOR * 100 + ROCBLAS_VERSION_MINOR) +#if ORT_ROCBLAS_VERSION_DECIMAL >= 242 +#define FLAG rocblas_gemm_flags_fp16_alt_impl +#else +#define FLAG 0 +#endif + +#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + using namespace onnxruntime; +inline int get_flag() { + int result = BackwardPassGuard::is_backward_pass() ? FLAG : 0; + //std::cerr << PRE << "get_flag() " << result << std::endl; + return result; +} + // Generalize library calls to be use in template functions // gemm @@ -67,7 +83,7 @@ inline rocblas_status rocblasGemmHelper(rocblas_handle handle, C, rocblas_datatype_f16_r, ldc, C, rocblas_datatype_f16_r, ldc, rocblas_datatype_f32_r, - rocblas_gemm_algo_standard, 0, 0); + rocblas_gemm_algo_standard, 0, get_flag()); } inline rocblas_status rocblasGemmHelper(rocblas_handle handle, @@ -90,7 +106,7 @@ inline rocblas_status rocblasGemmHelper(rocblas_handle handle, C, rocblas_datatype_f16_r, ldc, C, rocblas_datatype_f16_r, ldc, rocblas_datatype_f32_r, - rocblas_gemm_algo_standard, 0, 0); + rocblas_gemm_algo_standard, 0, get_flag()); } inline rocblas_status rocblasGemmHelper(rocblas_handle handle, @@ -225,7 +241,7 @@ inline rocblas_status rocblasGemmBatchedHelper(rocblas_handle handle, (void**)Carray, rocblas_datatype_f16_r, ldc, batchCount, rocblas_datatype_f32_r, - rocblas_gemm_algo_standard, 0, 0); + rocblas_gemm_algo_standard, 0, get_flag()); } inline rocblas_status rocblasGemmBatchedHelper(rocblas_handle handle, @@ -329,7 +345,7 @@ inline rocblas_status rocblasGemmStridedBatchedHelper(rocblas_handle handle, C, rocblas_datatype_f16_r, ldc, strideC, batchCount, rocblas_datatype_f32_r, - rocblas_gemm_algo_standard, 0, 0); + rocblas_gemm_algo_standard, 0, get_flag()); } inline rocblas_status rocblasGemmStridedBatchedHelper(rocblas_handle handle, @@ -357,7 +373,7 @@ inline rocblas_status rocblasGemmStridedBatchedHelper(rocblas_handle handle, C, rocblas_datatype_f16_r, ldc, strideC, batchCount, rocblas_datatype_f32_r, - rocblas_gemm_algo_standard, 0, 0); + rocblas_gemm_algo_standard, 0, get_flag()); } inline rocblas_status rocblasGemmStridedBatchedHelper(rocblas_handle handle, From f7c2bbc68584ac2cec88d139729f331c03df540e Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Mon, 17 Oct 2022 23:43:19 +0000 Subject: [PATCH 2/6] forward attributes to new FusedMatMul nodes --- onnxruntime/core/optimizer/matmul_scale_fusion.cc | 11 +++++++++++ onnxruntime/core/optimizer/matmul_transpose_fusion.cc | 11 +++++++++++ onnxruntime/core/optimizer/rocm_blas_alt_impl.cc | 4 ++-- onnxruntime/core/providers/rocm/rocm_kernel.h | 7 ++++--- 4 files changed, 28 insertions(+), 5 deletions(-) diff --git a/onnxruntime/core/optimizer/matmul_scale_fusion.cc b/onnxruntime/core/optimizer/matmul_scale_fusion.cc index 2c43f5ab12030..25d4d5f1074f4 100644 --- a/onnxruntime/core/optimizer/matmul_scale_fusion.cc +++ b/onnxruntime/core/optimizer/matmul_scale_fusion.cc @@ -12,6 +12,9 @@ #include "core/graph/graph_viewer.h" #include "core/optimizer/utils.h" +//#include +//#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + namespace onnxruntime { namespace { @@ -254,6 +257,14 @@ Status ProcessNode( kMSDomain); matmul_scale_node.SetExecutionProviderType(node.GetExecutionProviderType()); +#ifdef USE_ROCM + // forward the __altimpl, if present + auto& attrs = node.GetAttributes(); + if (attrs.count("__altimpl")) { + //std::cerr << PRE << " forwarding __altimpl attr " << static_cast(attrs.at("__altimpl").i()) << std::endl; + matmul_scale_node.AddAttribute("__altimpl", static_cast(attrs.at("__altimpl").i())); + } +#endif { InlinedVector> nodes_to_remove{node}; diff --git a/onnxruntime/core/optimizer/matmul_transpose_fusion.cc b/onnxruntime/core/optimizer/matmul_transpose_fusion.cc index d47538640b78e..6ac9c4fa757b4 100644 --- a/onnxruntime/core/optimizer/matmul_transpose_fusion.cc +++ b/onnxruntime/core/optimizer/matmul_transpose_fusion.cc @@ -6,6 +6,9 @@ #include "core/graph/graph_utils.h" #include +//#include +//#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " + using namespace ONNX_NAMESPACE; using namespace ::onnxruntime::common; namespace onnxruntime { @@ -404,6 +407,14 @@ Status MatmulTransposeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_ matmul_node.AddAttribute("alpha", alpha); // Assign provider to this new node. Provider should be same as the provider for old node. matmul_node.SetExecutionProviderType(node.GetExecutionProviderType()); +#ifdef USE_ROCM + // forward the __altimpl, if present + auto& attrs = node.GetAttributes(); + if (attrs.count("__altimpl")) { + //std::cerr << PRE << " forwarding __altimpl attr " << static_cast(attrs.at("__altimpl").i()) << std::endl; + matmul_node.AddAttribute("__altimpl", static_cast(attrs.at("__altimpl").i())); + } +#endif graph_utils::FinalizeNodeFusion(graph, matmul_node, node); diff --git a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc index 9594358dd257c..9c2a0cfca3207 100644 --- a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc +++ b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc @@ -21,9 +21,9 @@ Status RocmBlasAltImpl::ApplyImpl(Graph& graph, bool& modified, int graph_level, for (auto node_index : node_topology_list) { auto& node = *graph.GetNode(node_index); - std::cerr << PRE << node << std::endl; + //std::cerr << PRE << node << std::endl; -#if 0 +#if 1 if (node.OpType() == "YieldOp") { is_backward_pass = true; //std::cerr << PRE << "YieldOp found, before recurse" << std::endl; diff --git a/onnxruntime/core/providers/rocm/rocm_kernel.h b/onnxruntime/core/providers/rocm/rocm_kernel.h index 0c8522594c641..21cc0bdca357f 100644 --- a/onnxruntime/core/providers/rocm/rocm_kernel.h +++ b/onnxruntime/core/providers/rocm/rocm_kernel.h @@ -8,7 +8,7 @@ #include "core/providers/rocm/rocm_execution_provider.h" #include "core/providers/rocm/rocm_fwd.h" -#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " +//#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " namespace onnxruntime { namespace rocm { @@ -26,10 +26,11 @@ class RocmKernel : public OpKernel { Status Compute(OpKernelContext* p_op_kernel_context) const override { Status s; + const std::string& op_name = Info().GetKernelDef().OpName(); auto altimpl = Info().GetAttrOrDefault("__altimpl", 0); - //std::cerr << PRE << "altimpl " << altimpl << std::endl; + //std::cerr << PRE << op_name << " altimpl " << altimpl << std::endl; if (altimpl) { - //std::cerr << PRE << "creating BackwardPassGuard" << std::endl; + //std::cerr << PRE << op_name << " creating BackwardPassGuard" << std::endl; BackwardPassGuard guard; s = ComputeInternal(p_op_kernel_context); } From 3e7238a840e11b0a902411e39111ef87ff904993 Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Tue, 18 Oct 2022 15:14:01 +0000 Subject: [PATCH 3/6] remove debugging print statements --- onnxruntime/core/optimizer/matmul_scale_fusion.cc | 4 ---- onnxruntime/core/optimizer/matmul_transpose_fusion.cc | 4 ---- onnxruntime/core/optimizer/rocm_blas_alt_impl.cc | 7 ------- onnxruntime/core/providers/rocm/backward_guard.cc | 9 --------- onnxruntime/core/providers/rocm/math/gemm.cc | 4 ---- onnxruntime/core/providers/rocm/rocm_kernel.h | 4 ---- onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h | 3 --- 7 files changed, 35 deletions(-) diff --git a/onnxruntime/core/optimizer/matmul_scale_fusion.cc b/onnxruntime/core/optimizer/matmul_scale_fusion.cc index 25d4d5f1074f4..814dfc1667102 100644 --- a/onnxruntime/core/optimizer/matmul_scale_fusion.cc +++ b/onnxruntime/core/optimizer/matmul_scale_fusion.cc @@ -12,9 +12,6 @@ #include "core/graph/graph_viewer.h" #include "core/optimizer/utils.h" -//#include -//#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - namespace onnxruntime { namespace { @@ -261,7 +258,6 @@ Status ProcessNode( // forward the __altimpl, if present auto& attrs = node.GetAttributes(); if (attrs.count("__altimpl")) { - //std::cerr << PRE << " forwarding __altimpl attr " << static_cast(attrs.at("__altimpl").i()) << std::endl; matmul_scale_node.AddAttribute("__altimpl", static_cast(attrs.at("__altimpl").i())); } #endif diff --git a/onnxruntime/core/optimizer/matmul_transpose_fusion.cc b/onnxruntime/core/optimizer/matmul_transpose_fusion.cc index 6ac9c4fa757b4..7d06ee76b6107 100644 --- a/onnxruntime/core/optimizer/matmul_transpose_fusion.cc +++ b/onnxruntime/core/optimizer/matmul_transpose_fusion.cc @@ -6,9 +6,6 @@ #include "core/graph/graph_utils.h" #include -//#include -//#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - using namespace ONNX_NAMESPACE; using namespace ::onnxruntime::common; namespace onnxruntime { @@ -411,7 +408,6 @@ Status MatmulTransposeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_ // forward the __altimpl, if present auto& attrs = node.GetAttributes(); if (attrs.count("__altimpl")) { - //std::cerr << PRE << " forwarding __altimpl attr " << static_cast(attrs.at("__altimpl").i()) << std::endl; matmul_node.AddAttribute("__altimpl", static_cast(attrs.at("__altimpl").i())); } #endif diff --git a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc index 9c2a0cfca3207..ce05e69955d2f 100644 --- a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc +++ b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc @@ -6,8 +6,6 @@ #include "core/optimizer/rocm_blas_alt_impl.h" #include "core/graph/graph_utils.h" -#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - using namespace ONNX_NAMESPACE; using namespace ::onnxruntime::common; namespace onnxruntime { @@ -21,14 +19,10 @@ Status RocmBlasAltImpl::ApplyImpl(Graph& graph, bool& modified, int graph_level, for (auto node_index : node_topology_list) { auto& node = *graph.GetNode(node_index); - //std::cerr << PRE << node << std::endl; - #if 1 if (node.OpType() == "YieldOp") { is_backward_pass = true; - //std::cerr << PRE << "YieldOp found, before recurse" << std::endl; ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); - //std::cerr << PRE << "YieldOp found, after recurse" << std::endl; } else #else @@ -39,7 +33,6 @@ Status RocmBlasAltImpl::ApplyImpl(Graph& graph, bool& modified, int graph_level, } //if (node.OpType() == "MatMul" || node.OpType() == "FusedMatMul" || node.OpType() == "Gemm") { - //std::cerr << PRE << "HIT, is_backward_pass " << is_backward_pass << std::endl; if (is_backward_pass) { node.AddAttribute(std::string("__altimpl"), static_cast(1)); modified = true; diff --git a/onnxruntime/core/providers/rocm/backward_guard.cc b/onnxruntime/core/providers/rocm/backward_guard.cc index df05e8b82a50e..1695da092bef0 100644 --- a/onnxruntime/core/providers/rocm/backward_guard.cc +++ b/onnxruntime/core/providers/rocm/backward_guard.cc @@ -2,28 +2,19 @@ // Licensed under the MIT License. #include "core/providers/rocm/backward_guard.h" -#include -#include -#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - namespace onnxruntime { thread_local bool BackwardPassGuard::is_backward_pass_; BackwardPassGuard::BackwardPassGuard() { - //std::cerr << PRE << "BackwardPassGuard ctor pre " << is_backward_pass_ << std::endl; is_backward_pass_ = true; - //std::cerr << PRE << "BackwardPassGuard ctor post " << is_backward_pass_ << std::endl; } BackwardPassGuard::~BackwardPassGuard() { - //std::cerr << PRE << "BackwardPassGuard dtor pre " << is_backward_pass_ << std::endl; is_backward_pass_ = false; - //std::cerr << PRE << "BackwardPassGuard dtor post " << is_backward_pass_ << std::endl; } bool BackwardPassGuard::is_backward_pass() { - //std::cerr << PRE << "is_backward_pass_ " << is_backward_pass_ << std::endl; return is_backward_pass_; } diff --git a/onnxruntime/core/providers/rocm/math/gemm.cc b/onnxruntime/core/providers/rocm/math/gemm.cc index 191ed96b0613a..b0a32b5accf81 100644 --- a/onnxruntime/core/providers/rocm/math/gemm.cc +++ b/onnxruntime/core/providers/rocm/math/gemm.cc @@ -6,8 +6,6 @@ #include "core/providers/rocm/rocm_common.h" #include "core/providers/rocm/shared_inc/fpgeneric.h" -#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - namespace onnxruntime { namespace rocm { @@ -61,8 +59,6 @@ template Status Gemm::ComputeInternal(OpKernelContext* ctx) const { typedef typename ToHipType::MappedType HipT; - //auto altimpl = Info().GetAttrOrDefault("altimpl", 0); - //std::cerr << PRE << "wtf altimpl " << altimpl << std::endl; const auto* X = ctx->Input(0); const auto* W = ctx->Input(1); const auto* B = ctx->Input(2); diff --git a/onnxruntime/core/providers/rocm/rocm_kernel.h b/onnxruntime/core/providers/rocm/rocm_kernel.h index 21cc0bdca357f..ff235285c0923 100644 --- a/onnxruntime/core/providers/rocm/rocm_kernel.h +++ b/onnxruntime/core/providers/rocm/rocm_kernel.h @@ -8,8 +8,6 @@ #include "core/providers/rocm/rocm_execution_provider.h" #include "core/providers/rocm/rocm_fwd.h" -//#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - namespace onnxruntime { namespace rocm { @@ -28,9 +26,7 @@ class RocmKernel : public OpKernel { Status s; const std::string& op_name = Info().GetKernelDef().OpName(); auto altimpl = Info().GetAttrOrDefault("__altimpl", 0); - //std::cerr << PRE << op_name << " altimpl " << altimpl << std::endl; if (altimpl) { - //std::cerr << PRE << op_name << " creating BackwardPassGuard" << std::endl; BackwardPassGuard guard; s = ComputeInternal(p_op_kernel_context); } diff --git a/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h b/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h index 5e116ee9a9818..657a11ccb8e4d 100644 --- a/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h +++ b/onnxruntime/core/providers/rocm/shared_inc/fpgeneric.h @@ -13,13 +13,10 @@ #define FLAG 0 #endif -#define PRE __FILE__ << ":" << __LINE__ << ":" << std::this_thread::get_id() << " " - using namespace onnxruntime; inline int get_flag() { int result = BackwardPassGuard::is_backward_pass() ? FLAG : 0; - //std::cerr << PRE << "get_flag() " << result << std::endl; return result; } From fea021d7bb0be504682741af510fd37a75f63946 Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Tue, 18 Oct 2022 16:34:02 +0000 Subject: [PATCH 4/6] use __backwardpass attribute; code cleanup --- .../core/optimizer/graph_transformer_utils.cc | 3 +-- .../core/optimizer/matmul_scale_fusion.cc | 6 +++--- .../core/optimizer/matmul_transpose_fusion.cc | 6 +++--- .../core/optimizer/rocm_blas_alt_impl.cc | 17 +++++------------ onnxruntime/core/providers/rocm/rocm_kernel.h | 4 ++-- 5 files changed, 14 insertions(+), 22 deletions(-) diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index bf13fdfcd0741..28d4406eba7e6 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -208,8 +208,7 @@ InlinedVector> GenerateTransformers( auto cpu_allocator = cpu_execution_provider.GetAllocator(0, OrtMemTypeDefault); transformers.emplace_back(std::make_unique(std::move(cpu_allocator))); - std::cerr << __FILE__ << ":" << __LINE__ << " emplace_back RocmBlasAltImpl" << std::endl; - // TODO document + // add __backwardpass attribute to nodes after YieldOp, ROCm-only const InlinedHashSet rocm_ep = {onnxruntime::kRocmExecutionProvider}; transformers.emplace_back(std::make_unique(rocm_ep)); } break; diff --git a/onnxruntime/core/optimizer/matmul_scale_fusion.cc b/onnxruntime/core/optimizer/matmul_scale_fusion.cc index 814dfc1667102..b944b5536d2da 100644 --- a/onnxruntime/core/optimizer/matmul_scale_fusion.cc +++ b/onnxruntime/core/optimizer/matmul_scale_fusion.cc @@ -255,10 +255,10 @@ Status ProcessNode( matmul_scale_node.SetExecutionProviderType(node.GetExecutionProviderType()); #ifdef USE_ROCM - // forward the __altimpl, if present + // forward the __backwardpass, if present auto& attrs = node.GetAttributes(); - if (attrs.count("__altimpl")) { - matmul_scale_node.AddAttribute("__altimpl", static_cast(attrs.at("__altimpl").i())); + if (attrs.count("__backwardpass")) { + matmul_scale_node.AddAttribute("__backwardpass", static_cast(attrs.at("__backwardpass").i())); } #endif diff --git a/onnxruntime/core/optimizer/matmul_transpose_fusion.cc b/onnxruntime/core/optimizer/matmul_transpose_fusion.cc index 7d06ee76b6107..642805c93bb7c 100644 --- a/onnxruntime/core/optimizer/matmul_transpose_fusion.cc +++ b/onnxruntime/core/optimizer/matmul_transpose_fusion.cc @@ -405,10 +405,10 @@ Status MatmulTransposeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_ // Assign provider to this new node. Provider should be same as the provider for old node. matmul_node.SetExecutionProviderType(node.GetExecutionProviderType()); #ifdef USE_ROCM - // forward the __altimpl, if present + // forward the __backwardpass, if present auto& attrs = node.GetAttributes(); - if (attrs.count("__altimpl")) { - matmul_node.AddAttribute("__altimpl", static_cast(attrs.at("__altimpl").i())); + if (attrs.count("__backwardpass")) { + matmul_node.AddAttribute("__backwardpass", static_cast(attrs.at("__backwardpass").i())); } #endif diff --git a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc index ce05e69955d2f..4330f009df732 100644 --- a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc +++ b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc @@ -19,25 +19,18 @@ Status RocmBlasAltImpl::ApplyImpl(Graph& graph, bool& modified, int graph_level, for (auto node_index : node_topology_list) { auto& node = *graph.GetNode(node_index); -#if 1 if (node.OpType() == "YieldOp") { is_backward_pass = true; ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); } - else -#else - is_backward_pass = true; -#endif - { + else { ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); } - //if (node.OpType() == "MatMul" || node.OpType() == "FusedMatMul" || node.OpType() == "Gemm") { - if (is_backward_pass) { - node.AddAttribute(std::string("__altimpl"), static_cast(1)); - modified = true; - } - //} + if (is_backward_pass) { + node.AddAttribute(std::string("__backwardpass"), static_cast(1)); + modified = true; + } } return Status::OK(); diff --git a/onnxruntime/core/providers/rocm/rocm_kernel.h b/onnxruntime/core/providers/rocm/rocm_kernel.h index ff235285c0923..af74c1f2fc20d 100644 --- a/onnxruntime/core/providers/rocm/rocm_kernel.h +++ b/onnxruntime/core/providers/rocm/rocm_kernel.h @@ -25,8 +25,8 @@ class RocmKernel : public OpKernel { Status Compute(OpKernelContext* p_op_kernel_context) const override { Status s; const std::string& op_name = Info().GetKernelDef().OpName(); - auto altimpl = Info().GetAttrOrDefault("__altimpl", 0); - if (altimpl) { + auto is_backward_pass = Info().GetAttrOrDefault("__backwardpass", 0); + if (is_backward_pass) { BackwardPassGuard guard; s = ComputeInternal(p_op_kernel_context); } From 64ca66c03c0a35934c36373256139af4f318d5f4 Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Tue, 18 Oct 2022 16:42:58 +0000 Subject: [PATCH 5/6] remove redundant code --- onnxruntime/core/optimizer/rocm_blas_alt_impl.cc | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc index 4330f009df732..decb25f565efe 100644 --- a/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc +++ b/onnxruntime/core/optimizer/rocm_blas_alt_impl.cc @@ -21,12 +21,10 @@ Status RocmBlasAltImpl::ApplyImpl(Graph& graph, bool& modified, int graph_level, if (node.OpType() == "YieldOp") { is_backward_pass = true; - ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); - } - else { - ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); } + ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger)); + if (is_backward_pass) { node.AddAttribute(std::string("__backwardpass"), static_cast(1)); modified = true; From 532a9c9913f57245e26b87887530f8b361fb4918 Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Tue, 18 Oct 2022 16:44:26 +0000 Subject: [PATCH 6/6] remove unused var --- onnxruntime/core/providers/rocm/rocm_kernel.h | 1 - 1 file changed, 1 deletion(-) diff --git a/onnxruntime/core/providers/rocm/rocm_kernel.h b/onnxruntime/core/providers/rocm/rocm_kernel.h index af74c1f2fc20d..57473eb74db23 100644 --- a/onnxruntime/core/providers/rocm/rocm_kernel.h +++ b/onnxruntime/core/providers/rocm/rocm_kernel.h @@ -24,7 +24,6 @@ class RocmKernel : public OpKernel { Status Compute(OpKernelContext* p_op_kernel_context) const override { Status s; - const std::string& op_name = Info().GetKernelDef().OpName(); auto is_backward_pass = Info().GetAttrOrDefault("__backwardpass", 0); if (is_backward_pass) { BackwardPassGuard guard;