diff --git a/python/tvm/contrib/debugger/debug_executor.py b/python/tvm/contrib/debugger/debug_executor.py index 75932c0d5e34..785959ce8dd7 100644 --- a/python/tvm/contrib/debugger/debug_executor.py +++ b/python/tvm/contrib/debugger/debug_executor.py @@ -272,8 +272,10 @@ def debug_get_output(self, node, out=None): node_index = node else: raise RuntimeError("Require node index or name only.") - - self._debug_get_output(node_index, out) + if out: + self._debug_get_output(node_index, out) + return out + return self._debug_get_output(node_index) # pylint: disable=arguments-differ def run( diff --git a/src/runtime/graph_executor/debug/graph_executor_debug.cc b/src/runtime/graph_executor/debug/graph_executor_debug.cc index 0dbcbff46ff2..892a13b46bb4 100644 --- a/src/runtime/graph_executor/debug/graph_executor_debug.cc +++ b/src/runtime/graph_executor/debug/graph_executor_debug.cc @@ -197,10 +197,17 @@ PackedFunc GraphExecutorDebug::GetFunction(const String& name, // return member functions during query. if (name == "debug_get_output") { return PackedFunc([sptr_to_self, this](TVMArgs args, TVMRetValue* rv) { + int args0 = -1; if (String::CanConvertFrom(args[0])) { - this->DebugGetNodeOutput(this->GetNodeIndex(args[0]), args[1]); + args0 = this->GetNodeIndex(args[0]); } else { - this->DebugGetNodeOutput(args[0], args[1]); + args0 = args[0]; + } + + if (args.num_args == 2) { + this->DebugGetNodeOutput(args0, args[1]); + } else { + *rv = this->DebugGetNodeOutput(args0); } }); } else if (name == "execute_node") { @@ -325,6 +332,18 @@ void GraphExecutorDebug::DebugGetNodeOutput(int index, DLTensor* data_out) { data_entry_[eid].CopyTo(data_out); } +NDArray GraphExecutorDebug::DebugGetNodeOutput(int index) { + ICHECK_LT(static_cast(index), op_execs_.size()); + uint32_t eid = index; + + for (size_t i = 0; i < op_execs_.size(); ++i) { + if (op_execs_[i]) op_execs_[i](); + if (static_cast(i) == index) break; + } + + return data_entry_[eid]; +} + NDArray GraphExecutorDebug::GetNodeOutput(int node, int out_ind) { ICHECK_EQ(node, last_executed_node_); ICHECK_LT(entry_id(node, out_ind), data_entry_.size()); diff --git a/src/runtime/graph_executor/debug/graph_executor_debug.h b/src/runtime/graph_executor/debug/graph_executor_debug.h index 7c9d8f2cd176..382083056604 100644 --- a/src/runtime/graph_executor/debug/graph_executor_debug.h +++ b/src/runtime/graph_executor/debug/graph_executor_debug.h @@ -122,6 +122,18 @@ class GraphExecutorDebug : public GraphExecutor { */ void DebugGetNodeOutput(int index, DLTensor* data_out); + /*! + * \brief return output of index-th node. + * + * This method will do a partial run of the graph + * from begining up to the index-th node and return output of index-th node. + * This is costly operation and suggest to use only for debug porpose. + * + * \param index: The index of the node. + * + */ + NDArray DebugGetNodeOutput(int index); + /*! * \brief Profile execution time of the module. *